From 7686a4515f6ce0010d4cfd417d1674868f56ffa7 Mon Sep 17 00:00:00 2001 From: nightcityblade Date: Mon, 20 Jul 2026 11:18:23 +0800 Subject: [PATCH] feat: support tuple values in cyclical schedulers --- ignite/handlers/param_scheduler.py | 35 ++++++++++++++----- tests/ignite/handlers/test_param_scheduler.py | 24 +++++++++++++ 2 files changed, 51 insertions(+), 8 deletions(-) diff --git a/ignite/handlers/param_scheduler.py b/ignite/handlers/param_scheduler.py index a2765fa3b3e8..3d75293a5457 100644 --- a/ignite/handlers/param_scheduler.py +++ b/ignite/handlers/param_scheduler.py @@ -277,7 +277,7 @@ def simulate_values(cls, num_events: int, **scheduler_kwargs: Any) -> list[list[ values.append([i, scheduler.optimizer_param_groups[0][scheduler.param_name]]) return values - def _get_param(self) -> list[float] | float: + def _get_param(self) -> list[float] | tuple[float, ...] | float: # `ParamScheduler` does nothing special, only returning what child class returns. # Intermediate child classes edit this method return self.get_param() @@ -310,6 +310,8 @@ class CyclicalScheduler(ParamScheduler): Note: If the scheduler is bound to an 'ITERATION_*' event, 'cycle_size' should usually be the number of batches in an epoch. + Tuple parameters can be scheduled by passing tuples for both ``start_value`` + and ``end_value``. .. versionadded:: 0.4.5 @@ -321,8 +323,8 @@ def __init__( self, optimizer: Optimizer, param_name: str, - start_value: float, - end_value: float, + start_value: float | tuple[float, ...], + end_value: float | tuple[float, ...], cycle_size: int, cycle_mult: float = 1.0, start_value_mult: float = 1.0, @@ -332,8 +334,20 @@ def __init__( param_group_index: int | None = None, ): super().__init__(optimizer, param_name, save_history=save_history, param_group_index=param_group_index) - self.start_value = start_value - self.end_value = end_value + if isinstance(start_value, tuple): + if not isinstance(end_value, tuple): + raise TypeError("start_value and end_value should both be tuples") + if len(start_value) != len(end_value): + raise ValueError("start_value and end_value should have the same length") + tuple_values = True + elif isinstance(end_value, tuple): + raise TypeError("start_value and end_value should both be tuples") + else: + tuple_values = False + + self._tuple_values = tuple_values + self.start_value: Any = torch.tensor(start_value, dtype=torch.float64) if tuple_values else start_value + self.end_value: Any = torch.tensor(end_value, dtype=torch.float64) if tuple_values else end_value self.cycle_size = cycle_size self.cycle_mult = cycle_mult self.cycle = 0 @@ -370,15 +384,20 @@ def __call__(self, engine: Engine | None, name: str | None = None) -> None: return super(CyclicalScheduler, self).__call__(engine, name) - def _get_param(self) -> list[float] | float: + def _get_param(self) -> list[float] | tuple[float, ...] | float: """Applies warm-up if the scheduler is in the warm-up phase, otherwise returns what is returned by `self.get_param()` """ if self.event_index > self.cycle_size: warmup_progress = (self.event_index - self.cycle_size) / self.warmup_duration - return self.end_value + (self.start_value - self.end_value) * warmup_progress + value = self.end_value + (self.start_value - self.end_value) * warmup_progress + else: + value = self.get_param() - return self.get_param() + if self._tuple_values: + assert isinstance(value, torch.Tensor) + return tuple(value.tolist()) + return value class LinearCyclicalScheduler(CyclicalScheduler): diff --git a/tests/ignite/handlers/test_param_scheduler.py b/tests/ignite/handlers/test_param_scheduler.py index 6cc3b2893f6d..e13f6705eb50 100644 --- a/tests/ignite/handlers/test_param_scheduler.py +++ b/tests/ignite/handlers/test_param_scheduler.py @@ -37,6 +37,24 @@ def get_param(self): return [0] +def test_linear_scheduler_with_tuple_value(): + parameter = torch.nn.Parameter(torch.tensor(1.0)) + optimizer = torch.optim.Adam([parameter], betas=(0.9, 0.999)) + scheduler = LinearCyclicalScheduler(optimizer, "betas", (0.9, 0.999), (0.7, 0.999), cycle_size=4) + + values = [] + for _ in range(5): + scheduler(None) + values.append(optimizer.param_groups[0]["betas"]) + optimizer.zero_grad() + parameter.square().backward() + optimizer.step() + + assert all(isinstance(value, tuple) for value in values) + assert [value[0] for value in values] == pytest.approx([0.9, 0.8, 0.7, 0.8, 0.9]) + assert [value[1] for value in values] == pytest.approx([0.999] * 5) + + def test_param_scheduler_asserts(): t1 = torch.zeros([1], requires_grad=True) t2 = torch.zeros([1], requires_grad=True) @@ -64,6 +82,12 @@ def test_linear_scheduler_asserts(): tensor = torch.zeros([1], requires_grad=True) optimizer = torch.optim.SGD([tensor], lr=0.0) + with pytest.raises(TypeError, match="start_value and end_value should both be tuples"): + LinearCyclicalScheduler(optimizer, "lr", (1.0,), 0.0, cycle_size=2) + + with pytest.raises(ValueError, match="start_value and end_value should have the same length"): + LinearCyclicalScheduler(optimizer, "lr", (1.0,), (0.0, 0.5), cycle_size=2) + with pytest.raises(ValueError, match=r"Argument cycle_size should be positive and larger than 1"): LinearCyclicalScheduler(optimizer, "lr", 1, 0, cycle_size=0)