Skip to content

Optimizers

create_optimizer(model, optimizer_name, lr=0.001, weight_decay=0.0, wd_ban_list=('bias', 'LayerNorm.bias', 'LayerNorm.weight'), use_lookahead=False, use_orthograd=False, compile=False, compile_kwargs=None, **kwargs)

Create an optimizer for a model with optional wrappers and compilation.

Parameters:

Name Type Description Default
model Module

Model whose parameters to optimize.

required
optimizer_name str

Case insensitive name accepted by load_optimizer().

required
lr float | Tensor

Learning rate. Compiled updates use scalar tensors on the parameter device.

0.001
weight_decay float

Weight decay coefficient.

0.0
wd_ban_list Sequence[str]

Name patterns to exclude from weight decay. Matches parameter names and module class names.

('bias', 'LayerNorm.bias', 'LayerNorm.weight')
use_lookahead bool

Wrap the optimizer with Lookahead, unless it already includes Lookahead.

False
use_orthograd bool

Project gradients with OrthoGrad before each update.

False
compile bool

Compile optimizer updates with torch.compile. Supported foreach paths keep scalar bookkeeping eager.

False
compile_kwargs dict | None

Options for torch.compile. Dynamic tracing defaults to True.

None
**kwargs dict

Optimizer and wrapper options.

{}

Returns:

Name Type Description
Optimizer Optimizer

Configured optimizer instance.

Source code in pytorch_optimizer/optimizer/__init__.py
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
def create_optimizer(
    model: nn.Module,
    optimizer_name: str,
    lr: float | torch.Tensor = 1e-3,
    weight_decay: float = 0.0,
    wd_ban_list: Sequence[str] = ('bias', 'LayerNorm.bias', 'LayerNorm.weight'),
    use_lookahead: bool = False,
    use_orthograd: bool = False,
    compile: bool = False,  # noqa: A002
    compile_kwargs: dict | None = None,
    **kwargs,
) -> Optimizer:
    """Create an optimizer for a model with optional wrappers and compilation.

    Args:
        model: Model whose parameters to optimize.
        optimizer_name: Case insensitive name accepted by `load_optimizer()`.
        lr: Learning rate. Compiled updates use scalar tensors on the parameter device.
        weight_decay: Weight decay coefficient.
        wd_ban_list: Name patterns to exclude from weight decay. Matches parameter names and module class names.
        use_lookahead: Wrap the optimizer with Lookahead, unless it already includes Lookahead.
        use_orthograd: Project gradients with OrthoGrad before each update.
        compile: Compile optimizer updates with `torch.compile`. Supported foreach paths keep scalar bookkeeping eager.
        compile_kwargs: Options for `torch.compile`. Dynamic tracing defaults to `True`.
        **kwargs (dict): Optimizer and wrapper options.

    Returns:
        Optimizer: Configured optimizer instance.

    """
    optimizer_name = optimizer_name.lower()
    optimizer_type = load_optimizer(optimizer_name)

    use_compiled_foreach = (
        compile
        and issubclass(optimizer_type, BaseOptimizer)
        and optimizer_type._supports_compiled_foreach
        and kwargs.get('foreach') is not False
        and not use_orthograd
        and not use_lookahead
    )

    if use_compiled_foreach:
        kwargs.setdefault('foreach', True)

    if compile and not use_compiled_foreach and not isinstance(lr, torch.Tensor):
        lr = torch.tensor(lr, device=next(model.parameters()).device)

    if optimizer_name != 'lbfgs':
        kwargs['weight_decay'] = weight_decay

    parameters = (
        get_optimizer_parameters(model, weight_decay, wd_ban_list)
        if weight_decay > 0.0
        else [{'params': model.parameters(), 'weight_decay': weight_decay}]
    )

    optimizer_class = cast(Callable[..., Optimizer], optimizer_type)

    if optimizer_name == 'alig':
        optimizer = optimizer_class(parameters, max_lr=lr, **kwargs)
    elif optimizer_name in ('lomo', 'adalomo', 'adammini'):
        optimizer = optimizer_class(model, lr=lr, **kwargs)
    elif optimizer_name in ('muon', 'adamuon', 'adago', 'normuon'):
        warn(f'highly recommend you to manually create the {optimizer_name} manually.', UserWarning, stacklevel=1)

        optimizer = prepare_muon_parameters(model, optimizer_name, lr=lr, **kwargs)
    else:
        optimizer = optimizer_class(parameters, lr=lr, **kwargs)

    if use_orthograd:
        optimizer = OrthoGrad(optimizer, **kwargs)

    if use_lookahead:
        if optimizer_name in ('ranger', 'ranger21', 'ranger25'):
            warn(f'{optimizer} already has a Lookahead variant.', UserWarning, stacklevel=1)
        else:
            optimizer = Lookahead(
                optimizer,
                k=kwargs.get('k', 5),
                alpha=kwargs.get('alpha', 0.5),
                pullback_momentum=kwargs.get('pullback_momentum', 'none'),
            )

    if use_compiled_foreach and isinstance(optimizer, BaseOptimizer):
        optimizer._compile_foreach(compile_kwargs)
    elif compile:
        optimizer.step = MethodType(  # ty: ignore[invalid-assignment]
            torch.compile(optimizer.step.__func__, **{'dynamic': True, **(compile_kwargs or {})}),
            optimizer,
        )

    return optimizer

get_optimizer_parameters(model_or_parameter, weight_decay, wd_ban_list=('bias', 'LayerNorm.bias', 'LayerNorm.weight'))

Group trainable parameters by whether to apply weight decay.

With a model, patterns match parameter names and module class names. For example, LayerNorm excludes all parameters of LayerNorm modules. With named parameters, patterns match parameter names only.

Parameters:

Name Type Description Default
model_or_parameter Module | list

Model or list of (name, parameter) pairs.

required
weight_decay float

Weight decay coefficient for parameters outside the ban list.

required
wd_ban_list Sequence[str]

Substrings identifying parameters to exclude from weight decay.

('bias', 'LayerNorm.bias', 'LayerNorm.weight')

Returns:

Name Type Description
ParamsT ParamsT

Nonempty parameter groups with the requested weight decay or zero weight decay.

Source code in pytorch_optimizer/optimizer/__init__.py
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
def get_optimizer_parameters(
    model_or_parameter: nn.Module | list,
    weight_decay: float,
    wd_ban_list: Sequence[str] = ('bias', 'LayerNorm.bias', 'LayerNorm.weight'),
) -> ParamsT:
    """Group trainable parameters by whether to apply weight decay.

    With a model, patterns match parameter names and module class names. For example,
    `LayerNorm` excludes all parameters of LayerNorm modules. With named parameters,
    patterns match parameter names only.

    Args:
        model_or_parameter: Model or list of `(name, parameter)` pairs.
        weight_decay: Weight decay coefficient for parameters outside the ban list.
        wd_ban_list: Substrings identifying parameters to exclude from weight decay.

    Returns:
        ParamsT: Nonempty parameter groups with the requested weight decay or zero weight decay.

    """
    banned_parameter_ids: set[int] = set()

    if isinstance(model_or_parameter, nn.Module):
        for module_name, module in model_or_parameter.named_modules():
            for param_name, param in module.named_parameters(recurse=False):
                full_param_name: str = f'{module_name}.{param_name}' if module_name else param_name
                if any(
                    banned in pattern
                    for banned in wd_ban_list
                    for pattern in (full_param_name, module._get_name(), f'{module._get_name()}.{param_name}')
                ):
                    banned_parameter_ids.add(id(param))

        model_or_parameter = list(model_or_parameter.named_parameters())
    else:
        banned_parameter_ids.update(
            id(p) for n, p in model_or_parameter if any(pattern in n for pattern in wd_ban_list)
        )

    groups = [
        {
            'params': [
                p
                for n, p in model_or_parameter
                if p.requires_grad and id(p) not in banned_parameter_ids
            ],
            'weight_decay': weight_decay,
        },
        {
            'params': [
                p
                for n, p in model_or_parameter
                if p.requires_grad and id(p) in banned_parameter_ids
            ],
            'weight_decay': 0.0,
        },
    ]
    return [group for group in groups if group['params']]

A2Grad

Bases: BaseOptimizer

Adaptive accelerated stochastic gradient descent with three averaging variants.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float | None

Learning rate. No needed.

None
beta float

Coefficient controlling the adaptive gradient scale.

10.0
lips float

Lipschitz constant.

10.0
rho float

Represents the degree of weighting decrease, a constant smoothing factor between 0 and 1.

0.5
variant VARIANTS

Variant of A2Grad optimizer. One of 'uni', 'inc', or 'exp'.

'uni'
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/a2grad.py
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
class A2Grad(BaseOptimizer):
    """Adaptive accelerated stochastic gradient descent with three averaging variants.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate. No needed.
        beta: Coefficient controlling the adaptive gradient scale.
        lips: Lipschitz constant.
        rho: Represents the degree of weighting decrease, a constant smoothing factor between 0 and 1.
        variant: Variant of A2Grad optimizer. One of 'uni', 'inc', or 'exp'.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float | None = None,
        beta: float = 10.0,
        lips: float = 10.0,
        rho: float = 0.5,
        variant: VARIANTS = 'uni',
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_non_negative(lips, 'lips')
        self.validate_non_negative(rho, 'rho')
        self.validate_options(variant, 'variant', ['uni', 'inc', 'exp'])

        self.variant = variant
        self.maximize = maximize

        defaults: Defaults = {'beta': beta, 'lips': lips}
        if variant == 'exp':
            defaults.update({'rho': rho})

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'A2Grad'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['alpha_k'] = 1.0
                state['v_k'] = torch.zeros((1,), dtype=grad.dtype, device=grad.device)
                state['avg_grad'] = grad.clone()
                state['x_k'] = p.clone()
                if self.variant == 'exp':
                    state['v_kk'] = torch.zeros((1,), dtype=grad.dtype, device=grad.device)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            gamma_k: float = 2.0 * group['lips'] / (group['step'] + 1)
            alpha_k_1: float = 2.0 / (group['step'] + 3)

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                avg_grad, v_k, x_k = state['avg_grad'], state['v_k'], state['x_k']
                avg_grad.add_(grad - avg_grad, alpha=1.0 / group['step'])

                delta_k = grad.clone()
                delta_k.add_(avg_grad, alpha=-1.0)

                delta_k_sq = delta_k.pow(2).sum()

                if self.variant in ('uni', 'inc'):
                    if self.variant == 'inc':
                        v_k.mul_((group['step'] / (group['step'] + 1)) ** 2)
                    v_k.add_(delta_k_sq)
                else:
                    v_kk = state['v_kk']

                    v_kk.lerp_(delta_k_sq, weight=1.0 - group['rho'])
                    torch.max(v_kk, v_k, out=v_k)

                h_k = v_k.sqrt()
                if self.variant != 'uni':
                    h_k.mul_(math.sqrt(group['step'] + 1))

                coefficient = -1.0 / (gamma_k + group['beta'] * h_k.item())

                x_k.add_(grad, alpha=coefficient)

                p.lerp_(x_k, weight=alpha_k_1)
                p.add_(grad, alpha=(1.0 - alpha_k_1) * state['alpha_k'] * coefficient)

                state['alpha_k'] = alpha_k_1

        return loss

AccSGD

Bases: BaseOptimizer

Accelerated SGD with coupled short and long steps.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
kappa float

Ratio of long to short step.

1000.0
xi float

Statistical advantage parameter.

10.0
constant float

Any small constant under 1.

0.7
weight_decay float

Weight decay coefficient.

0.0
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/sgd.py
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
class AccSGD(BaseOptimizer):
    """Accelerated SGD with coupled short and long steps.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        kappa: Ratio of long to short step.
        xi: Statistical advantage parameter.
        constant: Any small constant under 1.
        weight_decay: Weight decay coefficient.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        kappa: float = 1000.0,
        xi: float = 10.0,
        constant: float = 0.7,
        weight_decay: float = 0.0,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_non_negative(kappa, 'kappa')
        self.validate_non_negative(xi, 'xi')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_boundary(constant, boundary=1.0, bound_type='upper')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'kappa': kappa,
            'xi': xi,
            'constant': constant,
            'weight_decay': weight_decay,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AccSGD'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['momentum_buffer'] = p.clone()

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            large_lr: float = group['lr'] * group['kappa'] / group['constant']
            alpha: float = 1.0 - (group['xi'] * (group['constant'] ** 2) / group['kappa'])
            beta: float = 1.0 - alpha
            zeta: float = group['constant'] / (group['constant'] + beta)

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=False,
                    fixed_decay=False,
                )

                buf = state['momentum_buffer']
                buf.mul_((1.0 / beta) - 1.0).add_(grad, alpha=-large_lr).add_(p).mul_(beta)

                p.add_(grad, alpha=-group['lr']).lerp_(buf, weight=1.0 - zeta)

        return loss

AdaBelief

Bases: BaseOptimizer

Adaptive updates based on gradient prediction error.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the gradient mean and squared gradient prediction error.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
rectify bool

Perform the rectified update similar to RAdam.

False
n_sma_threshold int

Minimum effective simple moving average length for rectification.

5
degenerated_to_sgd bool

Use an SGD update before the moving average reaches the rectification threshold.

True
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
eps float

Term added to the denominator to improve numerical stability.

1e-16
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adabelief.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
class AdaBelief(BaseOptimizer):
    """Adaptive updates based on gradient prediction error.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the gradient mean and squared gradient prediction error.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        rectify: Perform the rectified update similar to RAdam.
        n_sma_threshold: Minimum effective simple moving average length for rectification.
        degenerated_to_sgd: Use an SGD update before the moving average reaches the rectification threshold.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        rectify: bool = False,
        n_sma_threshold: int = 5,
        degenerated_to_sgd: bool = True,
        ams_bound: bool = False,
        foreach: bool | None = None,
        eps: float = 1e-16,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.n_sma_threshold = n_sma_threshold
        self.degenerated_to_sgd = degenerated_to_sgd
        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'rectify': rectify,
            'ams_bound': ams_bound,
            'foreach': foreach,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdaBelief'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_var'] = torch.zeros_like(p)

                if group['ams_bound']:
                    state['max_exp_avg_var'] = torch.zeros_like(p)

                if group.get('adanorm'):
                    state['exp_grad_adanorm'] = torch.zeros((1,), dtype=grad.dtype, device=grad.device)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        if group.get('foreach') is False:
            return False

        if group.get('adanorm') or group['rectify'] or group['ams_bound']:
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_vars: list[torch.Tensor],
        step_size: float,
    ) -> None:
        beta1, beta2 = group['betas']
        lr = group['lr']

        bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=lr,
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - beta1)

        grad_residuals = torch._foreach_sub(grads, exp_avgs)

        torch._foreach_mul_(exp_avg_vars, beta2)
        torch._foreach_addcmul_(exp_avg_vars, grad_residuals, grad_residuals, value=1.0 - beta2)
        torch._foreach_add_(exp_avg_vars, group['eps'])

        de_noms = torch._foreach_sqrt(exp_avg_vars)
        torch._foreach_div_(de_noms, bias_correction2_sq)
        torch._foreach_add_(de_noms, group['eps'])

        torch._foreach_addcdiv_(params, exp_avgs, de_noms, value=-step_size)

    def _step_per_param(self, group: ParamGroup, step_size: float, n_sma: float) -> None:
        beta1, beta2 = group['betas']

        bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            self.apply_weight_decay(
                p=p,
                grad=grad,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=group['fixed_decay'],
            )

            exp_avg, exp_avg_var = state['exp_avg'], state['exp_avg_var']

            p, grad, exp_avg, exp_avg_var = self.view_as_real(p, grad, exp_avg, exp_avg_var)

            s_grad = self.get_adanorm_gradient(
                grad=grad,
                adanorm=group.get('adanorm', False),
                exp_grad_norm=state.get('exp_grad_adanorm', None),
                r=group.get('adanorm_r', None),
            )

            exp_avg.lerp_(s_grad, weight=1.0 - beta1)

            grad_residual = grad - exp_avg
            exp_avg_var.mul_(beta2).addcmul_(grad_residual, grad_residual, value=1.0 - beta2).add_(group['eps'])

            de_nom = self.apply_ams_bound(
                ams_bound=group['ams_bound'],
                exp_avg_sq=exp_avg_var,
                max_exp_avg_sq=state.get('max_exp_avg_var', None),
                eps=0.0,
                exp_avg_sq_eps=0.0,
            )

            if not group['rectify']:
                de_nom.div_(bias_correction2_sq).add_(group['eps'])
                p.addcdiv_(exp_avg, de_nom, value=-step_size)
                continue

            de_nom.add_(group['eps'])

            if n_sma >= self.n_sma_threshold:
                p.addcdiv_(exp_avg, de_nom, value=-step_size)
            elif step_size > 0:
                p.add_(exp_avg, alpha=-step_size)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])

            step_size, n_sma = self.get_rectify_step_size(
                is_rectify=group['rectify'],
                step=group['step'],
                lr=group['lr'],
                beta2=beta2,
                n_sma_threshold=self.n_sma_threshold,
                degenerated_to_sgd=self.degenerated_to_sgd,
            )

            step_size = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=step_size,
                bias_correction1=bias_correction1,
            )

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_avg_var']
                )
                if params:
                    self._step_foreach(
                        group, params, grads, state_dict['exp_avg'], state_dict['exp_avg_var'], step_size
                    )
            else:
                self._step_per_param(group, step_size, n_sma)

        return loss

AdaBound

Bases: BaseOptimizer

Adam updates with learning rate bounds that converge to SGD.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
final_lr float

Final learning rate.

0.1
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
gamma float

Convergence speed of the bound functions.

0.001
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
Source code in pytorch_optimizer/optimizer/adabound.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
class AdaBound(BaseOptimizer):
    """Adam updates with learning rate bounds that converge to SGD.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        final_lr: Final learning rate.
        betas: Decay rates for the first and second moments.
        gamma: Convergence speed of the bound functions.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.

    """

    _supports_compiled_foreach = True

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        final_lr: float = 1e-1,
        betas: Betas = (0.9, 0.999),
        gamma: float = 1e-3,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        ams_bound: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        foreach: bool | None = None,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_positive(gamma, 'gamma')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'final_lr': final_lr,
            'gamma': gamma,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'ams_bound': ams_bound,
            'eps': eps,
            'foreach': foreach,
        }

        super().__init__(params, defaults)

        self.base_lrs: list[float] = [group['lr'] for group in self.param_groups]

    def __str__(self) -> str:
        return 'AdaBound'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

                if group['ams_bound']:
                    state['max_exp_avg_sq'] = torch.zeros_like(p)

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        state_dict: dict[str, list[torch.Tensor]],
        step_size: float | torch.Tensor,
        lower_bound: float | torch.Tensor,
        upper_bound: float | torch.Tensor,
    ) -> None:
        beta1, beta2 = group['betas']
        exp_avgs, exp_avg_sqs = state_dict['exp_avg'], state_dict['exp_avg_sq']

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=group['lr'],
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - beta1)

        torch._foreach_mul_(exp_avg_sqs, beta2)
        torch._foreach_addcmul_(exp_avg_sqs, grads, grads, value=1.0 - beta2)

        de_noms = self.apply_ams_bound_foreach(
            group['ams_bound'], exp_avg_sqs, state_dict.get('max_exp_avg_sq', []), group['eps']
        )

        updates = de_noms
        foreach_scalar_div_(updates, step_size)
        if isinstance(lower_bound, torch.Tensor):
            for update in updates:
                update.clamp_(min=lower_bound, max=upper_bound)
        else:
            torch._foreach_clamp_min_(updates, lower_bound)
            torch._foreach_clamp_max_(updates, upper_bound)

        torch._foreach_mul_(updates, exp_avgs)

        torch._foreach_sub_(params, updates)

    def _step_per_param(self, group: ParamGroup, step_size: float, lower_bound: float, upper_bound: float) -> None:
        beta1, beta2 = group['betas']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            self.apply_weight_decay(
                p=p,
                grad=grad,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=group['fixed_decay'],
            )

            exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']
            p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

            exp_avg.lerp_(grad, weight=1.0 - beta1)

            exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

            de_nom = self.apply_ams_bound(
                ams_bound=group['ams_bound'],
                exp_avg_sq=exp_avg_sq,
                max_exp_avg_sq=state.get('max_exp_avg_sq', None),
                eps=group['eps'],
            )

            update = de_nom
            foreach_scalar_div_([update], step_size)
            update.clamp_(min=lower_bound, max=upper_bound).mul_(exp_avg)

            p.sub_(update)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group_index, group in enumerate(self.param_groups):
            self.init_group(group)
            group['step'] += 1

            base_lr = self.base_lrs[group_index]
            if base_lr == 0.0 and group['lr'] > 0.0:
                base_lr = self.base_lrs[group_index] = group['lr']

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            final_lr: float = group['final_lr'] * group['lr'] / base_lr if base_lr > 0.0 else 0.0
            lower_bound: float = final_lr * (1 - 1 / (group['gamma'] * group['step'] + 1))
            upper_bound: float = final_lr * (1 + 1 / (group['gamma'] * group['step']))

            step_size = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=group['lr'] * bias_correction2_sq,
                bias_correction1=bias_correction1,
            )

            if self.can_use_foreach(group, group.get('foreach')):
                state_keys = ['exp_avg', 'exp_avg_sq']

                if group['ams_bound']:
                    state_keys.append('max_exp_avg_sq')

                params, grads, state_dict = self.collect_trainable_params(group, self.state, state_keys=state_keys)

                for tensors in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(
                        group, tensors['params'], tensors['grads'], tensors, step_size, lower_bound, upper_bound
                    )
            else:
                self._step_per_param(group, step_size, lower_bound, upper_bound)

        return loss

AdaDelta

Bases: BaseOptimizer

Adaptive updates based on accumulated squared gradients and updates.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

1.0
rho float

Coefficient used for computing a running average of squared gradients.

0.9
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-06
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adadelta.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
class AdaDelta(BaseOptimizer):
    """Adaptive updates based on accumulated squared gradients and updates.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        rho: Coefficient used for computing a running average of squared gradients.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1.0,
        rho: float = 0.9,
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        eps: float = 1e-6,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(rho, 'rho', 0.0, 1.0)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'rho': rho,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdaDelta'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['square_avg'] = torch.zeros_like(p)
                state['acc_delta'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            rho: float = group['rho']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                square_avg, acc_delta = state['square_avg'], state['acc_delta']

                p, grad, square_avg, acc_delta = self.view_as_real(p, grad, square_avg, acc_delta)

                square_avg.mul_(rho).addcmul_(grad, grad, value=1.0 - rho)

                std = square_avg.add(group['eps']).sqrt_()
                delta = acc_delta.add(group['eps']).sqrt_().div_(std).mul_(grad)

                acc_delta.mul_(rho).addcmul_(delta, delta, value=1.0 - rho)

                p.add_(delta, alpha=-group['lr'])

        return loss

AdaFactor

Bases: BaseOptimizer

Factored adaptive updates with the BigVision AdaFactor variant.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float | None

Learning rate.

0.001
betas tuple[None, float] | tuple[float, float] | tuple[float, float, float]

Update momentum decay and second moment decay cap. Set the first value to None to disable momentum.

(0.9, 0.999)
decay_rate float

Exponent controlling the step-dependent second moment decay.

-0.8
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
clip_threshold float

Maximum root mean square of the preconditioned update.

1.0
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
scale_parameter bool

If True, the learning rate is scaled by root mean square of parameter.

True
relative_step bool

If True, time dependent learning rate is computed instead of external learning rate.

True
warmup_init bool

Warm up the relative step size from 1e-6 * step.

False
eps1 float

Stability constant added to squared gradients.

1e-30
eps2 float

Lower bound for parameter RMS scaling.

0.001
momentum_dtype dtype

Data type for the optional momentum buffer.

bfloat16
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adafactor.py
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
class AdaFactor(BaseOptimizer):
    """Factored adaptive updates with the BigVision AdaFactor variant.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Update momentum decay and second moment decay cap. Set the first value to `None` to disable
            momentum.
        decay_rate: Exponent controlling the step-dependent second moment decay.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        clip_threshold: Maximum root mean square of the preconditioned update.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        scale_parameter: If True, the learning rate is scaled by root mean square of parameter.
        relative_step: If True, time dependent learning rate is computed instead of external learning rate.
        warmup_init: Warm up the relative step size from `1e-6 * step`.
        eps1: Stability constant added to squared gradients.
        eps2: Lower bound for parameter RMS scaling.
        momentum_dtype: Data type for the optional momentum buffer.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float | None = 1e-3,
        betas: tuple[None, float] | tuple[float, float] | tuple[float, float, float] = (0.9, 0.999),
        decay_rate: float = -0.8,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        clip_threshold: float = 1.0,
        ams_bound: bool = False,
        scale_parameter: bool = True,
        relative_step: bool = True,
        warmup_init: bool = False,
        eps1: float = 1e-30,
        eps2: float = 1e-3,
        momentum_dtype: torch.dtype = torch.bfloat16,
        foreach: bool | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps1, 'eps1')
        self.validate_non_negative(eps2, 'eps2')

        self.decay_rate = decay_rate
        self.clip_threshold = clip_threshold
        self.eps1: float = eps1 if momentum_dtype != torch.float16 else 1e-7
        self.eps2 = eps2
        self.momentum_dtype = momentum_dtype
        self.foreach = foreach
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'ams_bound': ams_bound,
            'scale_parameter': scale_parameter,
            'relative_step': relative_step,
            'warmup_init': warmup_init,
            'eps1': eps1,
            'eps2': eps2,
            'foreach': foreach,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdaFactor'

    def load_state_dict(self, state_dict: dict) -> None:
        super().load_state_dict(state_dict)
        for group, saved_group in zip(self.param_groups, state_dict['param_groups']):
            for p, key in zip(group['params'], saved_group['params']):
                saved_state = state_dict['state'].get(key, {})
                if 'exp_avg' in saved_state:
                    self.state[p]['exp_avg'] = saved_state['exp_avg'].to(device=p.device, dtype=self.momentum_dtype)

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        beta1: float = kwargs.get('beta1', 0.9)

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            grad_shape: tuple[int, ...] = grad.shape
            factored: bool = self.get_options(grad_shape)

            if len(state) == 0:
                state['RMS'] = 0.0

                if beta1 is not None:
                    state['exp_avg'] = torch.zeros_like(p, dtype=self.momentum_dtype)

                if factored:
                    state['exp_avg_sq_row'] = torch.zeros(grad_shape[:-1], dtype=grad.dtype, device=grad.device)
                    state['exp_avg_sq_col'] = torch.zeros(
                        grad_shape[:-2] + grad_shape[-1:], dtype=grad.dtype, device=grad.device
                    )
                else:
                    state['exp_avg_sq'] = torch.zeros_like(grad)

                if group['ams_bound']:
                    state['exp_avg_sq_hat'] = torch.zeros_like(grad)

    @staticmethod
    def get_relative_step_size(lr: float, step: int, relative_step: bool, warmup_init: bool) -> float:
        if not relative_step:
            return lr

        min_step: float = 1e-6 * step if warmup_init else 1e-2
        return min(min_step, 1.0 / math.sqrt(step))

    def get_lr(
        self,
        relative_step_size: torch.Tensor | float,
        rms: Sequence[torch.Tensor] | torch.Tensor | float,
        scale_parameter: bool,
    ) -> Sequence[torch.Tensor] | torch.Tensor | float:
        """Compute effective learning rates with optional parameter RMS scaling."""
        if not scale_parameter:
            return relative_step_size

        if not isinstance(rms, Sequence):
            return max(self.eps2, rms) * relative_step_size

        lrs = torch._foreach_maximum(rms, self.eps2)
        torch._foreach_mul_(lrs, relative_step_size)

        return lrs

    @staticmethod
    def get_options(shape: tuple[int, ...]) -> bool:
        """Return whether the gradient supports factored second moments."""
        return len(shape) >= 2

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        if group.get('foreach') is False:
            return False

        if group.get('cautious'):
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sq_rows: list[torch.Tensor],
        exp_avg_sq_cols: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        exp_avg_sq_hats: list[torch.Tensor],
        beta1: float,
        beta2_t: float,
        relative_step_size: float,
    ) -> None:
        bias_correction2: float = 1.0 - beta2_t

        if self.maximize:
            torch._foreach_neg_(grads)

        rms_values = self.get_rms(params)
        lrs = self.get_lr(relative_step_size, rms_values, group['scale_parameter'])

        updates = torch._foreach_pow(grads, 2)
        torch._foreach_add_(updates, self.eps1)

        factored_offsets, non_factored_offsets = [], []
        factored_updates, non_factored_updates = [], []
        for i, grad in enumerate(grads):
            if self.get_options(grad.shape):
                factored_updates.append(updates[i])
                factored_offsets.append(i)
            else:
                non_factored_updates.append(updates[i])
                non_factored_offsets.append(i)

        if factored_updates:
            row_means, col_means = [], []
            for factored_update in factored_updates:
                row_means.append(factored_update.mean(dim=-1))
                col_means.append(factored_update.mean(dim=-2))

            torch._foreach_lerp_(exp_avg_sq_rows, row_means, weight=bias_correction2)
            torch._foreach_lerp_(exp_avg_sq_cols, col_means, weight=bias_correction2)

            self.approximate_sq_grad(exp_avg_sq_rows, exp_avg_sq_cols, factored_updates)

        if non_factored_updates:
            torch._foreach_lerp_(exp_avg_sqs, non_factored_updates, weight=bias_correction2)

            non_factored_updates = foreach_rsqrt(exp_avg_sqs)

        updates = [None] * len(grads)

        for offset, update in zip(factored_offsets, factored_updates):
            updates[offset] = update

        for offset, update in zip(non_factored_offsets, non_factored_updates):
            updates[offset] = update

        if group['ams_bound']:
            inv_updates = torch._foreach_reciprocal(updates)
            torch._foreach_maximum_(exp_avg_sq_hats, inv_updates)

            updates = foreach_rsqrt(torch._foreach_div(exp_avg_sq_hats, bias_correction2))

        torch._foreach_mul_(updates, grads)

        rms_values = self.get_rms(updates)
        torch._foreach_div_(rms_values, self.clip_threshold)
        torch._foreach_clamp_min_(rms_values, 1.0)

        torch._foreach_div_(updates, rms_values)
        torch._foreach_mul_(updates, lrs)

        if beta1 is not None:
            is_dtype_different: bool = self.momentum_dtype != grads[0].dtype
            if is_dtype_different:
                updates = [update.to(self.momentum_dtype) for update in updates]

            torch._foreach_lerp_(exp_avgs, updates, weight=1.0 - beta1)

            if is_dtype_different:
                updates = [exp_avg.to(grads[0].dtype) for exp_avg in exp_avgs]
            else:
                torch._foreach_copy_(updates, exp_avgs)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=lrs,
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        torch._foreach_sub_(params, updates)

    def _step_per_param(self, group: ParamGroup, beta1: float, beta2_t: float, relative_step_size: float) -> None:
        bias_correction2: float = 1.0 - beta2_t
        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            factored: bool = self.get_options(grad.shape)

            state['RMS'] = self.get_rms(p)

            lr = self.get_lr(relative_step_size, state['RMS'], group['scale_parameter'])

            # NOTE(kozistr): adding `eps1` here instead of clipping max by eps1 later
            update = grad.square().add_(self.eps1)

            if factored:
                exp_avg_sq_row, exp_avg_sq_col = state['exp_avg_sq_row'], state['exp_avg_sq_col']

                exp_avg_sq_row.lerp_(update.mean(dim=-1), weight=bias_correction2)
                exp_avg_sq_col.lerp_(update.mean(dim=-2), weight=bias_correction2)

                self.approximate_sq_grad(exp_avg_sq_row, exp_avg_sq_col, update)
            else:
                exp_avg_sq = state['exp_avg_sq']
                exp_avg_sq.lerp_(update, weight=bias_correction2)
                torch.rsqrt(exp_avg_sq, out=update)

            if group['ams_bound']:
                exp_avg_sq_hat = state['exp_avg_sq_hat']
                torch.max(exp_avg_sq_hat, 1.0 / update, out=exp_avg_sq_hat)
                torch.rsqrt(exp_avg_sq_hat / bias_correction2, out=update)

            update.mul_(grad)

            factor = self.get_rms(update).div_(self.clip_threshold).clamp_min_(1.0)
            update.div_(factor).mul_(lr)

            if beta1 is not None:
                exp_avg = state['exp_avg']
                if self.momentum_dtype != grad.dtype:
                    exp_avg.lerp_(update.to(self.momentum_dtype), weight=1.0 - beta1)
                    update = exp_avg.to(grad.dtype)
                else:
                    exp_avg.lerp_(update, weight=1.0 - beta1)
                    update = exp_avg.clone()

                if group.get('cautious'):
                    self.apply_cautious(update, grad)

            self.apply_weight_decay(
                p=p,
                grad=None,
                lr=lr,
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=group['fixed_decay'],
            )

            p.add_(-update)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            beta1, beta2_cap = group['betas']

            self.init_group(group, beta1=beta1)
            group['step'] += 1

            beta2_t: float = min(beta2_cap, 1.0 - math.pow(group['step'], self.decay_rate))

            relative_step_size: float = self.get_relative_step_size(
                lr=group['lr'],
                step=group['step'],
                relative_step=group['relative_step'],
                warmup_init=group['warmup_init'],
            )

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(
                    group,
                    self.state,
                    state_keys=['exp_avg', 'exp_avg_sq_row', 'exp_avg_sq_col', 'exp_avg_sq', 'exp_avg_sq_hat'],
                )
                if params:
                    self._step_foreach(
                        group,
                        params,
                        grads,
                        state_dict['exp_avg'],
                        state_dict['exp_avg_sq_row'],
                        state_dict['exp_avg_sq_col'],
                        state_dict['exp_avg_sq'],
                        state_dict['exp_avg_sq_hat'],
                        beta1,
                        beta2_t,
                        relative_step_size,
                    )
            else:
                self._step_per_param(group, beta1, beta2_t, relative_step_size)

        return loss

get_lr(relative_step_size, rms, scale_parameter)

Compute effective learning rates with optional parameter RMS scaling.

Source code in pytorch_optimizer/optimizer/adafactor.py
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
def get_lr(
    self,
    relative_step_size: torch.Tensor | float,
    rms: Sequence[torch.Tensor] | torch.Tensor | float,
    scale_parameter: bool,
) -> Sequence[torch.Tensor] | torch.Tensor | float:
    """Compute effective learning rates with optional parameter RMS scaling."""
    if not scale_parameter:
        return relative_step_size

    if not isinstance(rms, Sequence):
        return max(self.eps2, rms) * relative_step_size

    lrs = torch._foreach_maximum(rms, self.eps2)
    torch._foreach_mul_(lrs, relative_step_size)

    return lrs

get_options(shape) staticmethod

Return whether the gradient supports factored second moments.

Source code in pytorch_optimizer/optimizer/adafactor.py
166
167
168
169
@staticmethod
def get_options(shape: tuple[int, ...]) -> bool:
    """Return whether the gradient supports factored second moments."""
    return len(shape) >= 2

AdaGC

Bases: BaseOptimizer

Adam with adaptive gradient clipping for stable training.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
beta float

Smoothing coefficient for the exponential moving average (EMA).

0.98
lambda_abs float

Absolute clipping threshold to prevent unstable updates from gradient explosions.

1.0
lambda_rel float

Relative clipping threshold to prevent unstable updates from gradient explosions.

1.05
warmup_steps int

Number of warmup steps.

100
weight_decay float

Weight decay coefficient.

0.1
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adagc.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
class AdaGC(BaseOptimizer):
    """Adam with adaptive gradient clipping for stable training.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        beta: Smoothing coefficient for the exponential moving average (EMA).
        lambda_abs: Absolute clipping threshold to prevent unstable updates from gradient explosions.
        lambda_rel: Relative clipping threshold to prevent unstable updates from gradient explosions.
        warmup_steps: Number of warmup steps.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        beta: float = 0.98,
        lambda_abs: float = 1.0,
        lambda_rel: float = 1.05,
        warmup_steps: int = 100,
        weight_decay: float = 1e-1,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_range(beta, 'beta', 0.0, 1.0, '[)')
        self.validate_positive(lambda_abs, 'lambda_abs')
        self.validate_positive(lambda_rel, 'lambda_rel')
        self.validate_non_negative(warmup_steps, 'warmup_steps')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'beta': beta,
            'lambda_abs': lambda_abs,
            'lambda_rel': lambda_rel,
            'warmup_steps': warmup_steps,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdaGC'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if 'exp_avg' not in state:
                state['exp_avg'] = torch.zeros_like(grad)
                state['exp_avg_sq'] = torch.zeros_like(grad)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']
                gamma = state.get('gamma')

                if gamma is None or group['step'] < group['warmup_steps']:
                    grad_norm = get_global_gradient_norm(self.param_groups).sqrt_().add_(group['eps'])

                    h_t = min(group['lambda_abs'] / grad_norm, 1.0)
                    g_hat = grad.mul(h_t)

                    g_hat_norm = g_hat.norm()

                    state['gamma'] = g_hat_norm if gamma is None else gamma.copy_(min(gamma, g_hat_norm))
                else:
                    h_t = (
                        group['lambda_rel'] * gamma.clamp_min(group['eps']) / grad.norm().clamp_min(group['eps'])
                    ).clamp_max_(1.0)
                    g_hat = grad.mul(h_t)

                    gamma.lerp_(g_hat.norm(), weight=1.0 - group['beta'])

                exp_avg.lerp_(g_hat, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(g_hat, g_hat, value=1.0 - beta2)

                update = (exp_avg / bias_correction1) / exp_avg_sq.sqrt().div_(bias_correction2_sq).add_(group['eps'])

                p.add_(update, alpha=-group['lr'])

        return loss

AdaGO

Bases: MuonBase

Orthogonal momentum updates with AdaGrad step size adaptation.

Set use_muon=True for hidden weight matrices and use_muon=False for AdamW groups, such as embeddings, classifier heads, biases, and gains. Pass higher dimensional weights directly. The orthogonal update uses a flattened matrix view.

Parameters:

Name Type Description Default
params ParamsT

Parameter group dictionaries with a use_muon flag for each group.

required
lr float

Learning rate.

0.05
momentum float

Momentum factor.

0.95
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
nesterov bool

Use Nesterov momentum.

False
gamma float

Gradient norm cap for the accumulator and adaptive step size.

10.0
v float

Initial value of the AdaGrad accumulator.

1e-06
eps float

Epsilon value. Lower bound eps > 0 on the stepsizes.

0.0005
ns_steps int

Number of Newton-Schulz iterations.

5
ns_coeffs NewtonSchulzWeights

Newton-Schulz coefficients or preset name.

'original'
use_adjusted_lr bool

Scale orthogonal updates using the Moonlight shape adjustment.

False
adamw_lr float

Learning rate for parameters in the AdamW groups.

0.0003
adamw_betas Betas

Decay rates for the first and second moments in the AdamW groups.

(0.9, 0.95)
adamw_wd float

Weight decay for parameters in the AdamW groups.

0.0
adamw_eps float

Numerical stability constant for the AdamW groups.

1e-10
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Batch tensor updates and compatible matrix shapes. False disables batching; None enables it.

False

Examples:

from pytorch_optimizer import AdaGO

hidden_weights = [p for p in model.body.parameters() if p.ndim >= 2]
hidden_gains_biases = [p for p in model.body.parameters() if p.ndim < 2]
non_hidden_params = [*model.head.parameters(), *model.embed.parameters()]

param_groups = [
    dict(params=hidden_weights, lr=0.02, weight_decay=0.01, use_muon=True),
    dict(
        params=hidden_gains_biases + non_hidden_params,
        lr=3e-4,
        betas=(0.9, 0.95),
        weight_decay=0.01,
        use_muon=False,
    ),
]

optimizer = AdaGO(param_groups)
Source code in pytorch_optimizer/optimizer/muon.py
 831
 832
 833
 834
 835
 836
 837
 838
 839
 840
 841
 842
 843
 844
 845
 846
 847
 848
 849
 850
 851
 852
 853
 854
 855
 856
 857
 858
 859
 860
 861
 862
 863
 864
 865
 866
 867
 868
 869
 870
 871
 872
 873
 874
 875
 876
 877
 878
 879
 880
 881
 882
 883
 884
 885
 886
 887
 888
 889
 890
 891
 892
 893
 894
 895
 896
 897
 898
 899
 900
 901
 902
 903
 904
 905
 906
 907
 908
 909
 910
 911
 912
 913
 914
 915
 916
 917
 918
 919
 920
 921
 922
 923
 924
 925
 926
 927
 928
 929
 930
 931
 932
 933
 934
 935
 936
 937
 938
 939
 940
 941
 942
 943
 944
 945
 946
 947
 948
 949
 950
 951
 952
 953
 954
 955
 956
 957
 958
 959
 960
 961
 962
 963
 964
 965
 966
 967
 968
 969
 970
 971
 972
 973
 974
 975
 976
 977
 978
 979
 980
 981
 982
 983
 984
 985
 986
 987
 988
 989
 990
 991
 992
 993
 994
 995
 996
 997
 998
 999
1000
1001
1002
1003
1004
1005
1006
1007
1008
1009
1010
1011
1012
1013
1014
1015
1016
1017
1018
1019
1020
1021
1022
1023
1024
1025
1026
1027
1028
1029
1030
1031
1032
1033
1034
1035
1036
1037
1038
1039
1040
1041
1042
1043
1044
1045
1046
1047
1048
1049
1050
1051
1052
1053
1054
1055
1056
1057
1058
1059
1060
1061
1062
1063
1064
1065
1066
1067
1068
1069
1070
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
1081
1082
1083
1084
1085
class AdaGO(MuonBase):
    """Orthogonal momentum updates with AdaGrad step size adaptation.

    Set `use_muon=True` for hidden weight matrices and `use_muon=False` for AdamW groups,
    such as embeddings, classifier heads, biases, and gains. Pass higher dimensional
    weights directly. The orthogonal update uses a flattened matrix view.

    Args:
        params: Parameter group dictionaries with a `use_muon` flag for each group.
        lr: Learning rate.
        momentum: Momentum factor.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        nesterov: Use Nesterov momentum.
        gamma: Gradient norm cap for the accumulator and adaptive step size.
        v: Initial value of the AdaGrad accumulator.
        eps: Epsilon value. Lower bound eps > 0 on the stepsizes.
        ns_steps: Number of Newton-Schulz iterations.
        ns_coeffs: Newton-Schulz coefficients or preset name.
        use_adjusted_lr: Scale orthogonal updates using the Moonlight shape adjustment.
        adamw_lr: Learning rate for parameters in the AdamW groups.
        adamw_betas: Decay rates for the first and second moments in the AdamW groups.
        adamw_wd: Weight decay for parameters in the AdamW groups.
        adamw_eps: Numerical stability constant for the AdamW groups.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Batch tensor updates and compatible matrix shapes. `False` disables batching; `None` enables it.

    Examples:
        ```python
        from pytorch_optimizer import AdaGO

        hidden_weights = [p for p in model.body.parameters() if p.ndim >= 2]
        hidden_gains_biases = [p for p in model.body.parameters() if p.ndim < 2]
        non_hidden_params = [*model.head.parameters(), *model.embed.parameters()]

        param_groups = [
            dict(params=hidden_weights, lr=0.02, weight_decay=0.01, use_muon=True),
            dict(
                params=hidden_gains_biases + non_hidden_params,
                lr=3e-4,
                betas=(0.9, 0.95),
                weight_decay=0.01,
                use_muon=False,
            ),
        ]

        optimizer = AdaGO(param_groups)
        ```

    """

    _muon_state_keys = ('momentum_buffer', 'v')

    def __init__(
        self,
        params: ParamsT,
        lr: float = 5e-2,
        momentum: float = 0.95,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        gamma: float = 10.0,
        eps: float = 5e-4,
        v: float = 1e-6,
        nesterov: bool = False,
        ns_steps: int = 5,
        ns_coeffs: NewtonSchulzWeights = 'original',
        use_adjusted_lr: bool = False,
        adamw_lr: float = 3e-4,
        adamw_betas: Betas = (0.9, 0.95),
        adamw_wd: float = 0.0,
        adamw_eps: float = 1e-10,
        maximize: bool = False,
        foreach: bool | None = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_learning_rate(adamw_lr)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(momentum, 'momentum', 0.0, 1.0, range_type='[)')
        self.validate_positive(ns_steps, 'ns_steps')
        self.validate_positive(gamma, 'gamma')
        self.validate_positive(eps, 'eps')
        self.validate_positive(v, 'v')
        self.validate_betas(adamw_betas)
        self.validate_non_negative(adamw_wd, 'adamw_wd')
        self.validate_non_negative(adamw_eps, 'adamw_eps')
        ns_coeffs = get_newton_schulz_weights(ns_coeffs)

        self.maximize = maximize
        self.foreach = foreach

        for group in params:
            group = cast(ParamGroup, group)
            if 'use_muon' not in group:
                raise ValueError('`use_muon` must be set.')

            if group['use_muon']:
                group['lr'] = group.get('lr', lr)
                group['momentum'] = group.get('momentum', momentum)
                group['nesterov'] = group.get('nesterov', nesterov)
                group['weight_decay'] = group.get('weight_decay', weight_decay)
                group['ns_steps'] = group.get('ns_steps', ns_steps)
                group['ns_coeffs'] = get_newton_schulz_weights(group.get('ns_coeffs', ns_coeffs))
                group['gamma'] = group.get('gamma', gamma)
                group['eps'] = group.get('eps', eps)
                group['v'] = group.get('v', v)
                group['use_adjusted_lr'] = group.get('use_adjusted_lr', use_adjusted_lr)
            else:
                group['lr'] = group.get('lr', adamw_lr)
                group['betas'] = group.get('betas', adamw_betas)
                group['eps'] = group.get('eps', adamw_eps)
                group['weight_decay'] = group.get('weight_decay', adamw_wd)

            group['weight_decouple'] = group.get('weight_decouple', weight_decouple)

        super().__init__(params, {'foreach': foreach, **kwargs})

    def __str__(self) -> str:
        return 'AdaGO'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                if group['use_muon']:
                    state['momentum_buffer'] = torch.zeros_like(p)
                    state['v'] = torch.tensor(group['v'], dtype=p.dtype, device=p.device)
                else:
                    state['exp_avg'] = torch.zeros_like(p)
                    state['exp_avg_sq'] = torch.zeros_like(p)

    def _step_muon_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        state_dict: dict[str, list[torch.Tensor]],
        bias_correction2: float | torch.Tensor,
    ) -> None:
        buffers, variances = state_dict['momentum_buffer'], state_dict['v']

        torch._foreach_lerp_(buffers, grads, weight=1.0 - group['momentum'])

        grad_norms = torch._foreach_norm(grads, ord=2)
        squared_norms = torch._foreach_mul(grad_norms, grad_norms)
        torch._foreach_clamp_max_(squared_norms, group['gamma'] ** 2)
        torch._foreach_add_(variances, squared_norms)

        if group['nesterov']:
            torch._foreach_lerp_(grads, buffers, weight=group['momentum'])

        updates = self._orthogonalize(group, grads if group['nesterov'] else buffers)
        updates = [update.reshape(p.shape) for p, update in zip(params, updates)]

        if group.get('cautious'):
            for update, grad in zip(updates, grads):
                self.apply_cautious(update, grad)

        # Nesterov modifies gradients before the adaptive step size is computed.
        step_sizes = torch._foreach_norm(grads, ord=2) if group['nesterov'] else grad_norms
        torch._foreach_clamp_max_(step_sizes, group['gamma'])

        lr = get_adjusted_lr(group['lr'], params[0].shape, use_adjusted_lr=group['use_adjusted_lr'])
        torch._foreach_mul_(step_sizes, lr)
        torch._foreach_div_(step_sizes, variances)
        torch._foreach_clamp_min_(step_sizes, group['eps'])

        torch._foreach_addcmul_(params, [update.to(params[0].dtype) for update in updates], step_sizes, value=-1.0)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if self.can_use_foreach(group, group.get('foreach', self.foreach)):
                self._step_foreach_group(group)
                continue

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=False,
                )

                if group['use_muon']:
                    buf, v = state['momentum_buffer'], state['v']
                    buf.lerp_(grad, weight=1.0 - group['momentum'])

                    grad_norm = grad.norm(p=2.0)
                    v.add_(grad_norm.square().clamp_max_(group['gamma'] ** 2))

                    update = grad.lerp_(buf, weight=group['momentum']) if group['nesterov'] else buf
                    if update.ndim > 2:
                        update = update.view(len(update), -1)

                    update = zero_power_via_newton_schulz_5(
                        update, num_steps=group['ns_steps'], weights=group['ns_coeffs']
                    )

                    if group.get('cautious'):
                        self.apply_cautious(update.reshape(p.shape), grad)

                    lr = get_adjusted_lr(group['lr'], p.size(), use_adjusted_lr=group['use_adjusted_lr'])

                    step_size = (lr * grad.norm(2).clamp_max_(group['gamma']) / v).clamp_min_(group['eps'])
                    p.addcmul_(update.reshape(p.shape).to(p.dtype), step_size, value=-1.0)
                else:
                    exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                    beta1, beta2 = group['betas']

                    bias_correction1: float = self.debias(beta1, group['step'])
                    bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

                    exp_avg.lerp_(grad, weight=1.0 - beta1)
                    exp_avg_sq.lerp_(grad.square(), weight=1.0 - beta2)

                    de_nom = exp_avg_sq.sqrt().div_(bias_correction2_sq).add_(group['eps'])

                    p.addcdiv_(exp_avg / bias_correction1, de_nom, value=-group['lr'])

        return loss

AdaHessian

Bases: BaseOptimizer

Adaptive second-order updates using Hutchinson Hessian estimates.

Use loss.backward(create_graph=True) for internal Hessian estimation, or supply external estimates through step(hessian=...).

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.1
betas Betas

Decay rates for the gradient mean and squared Hessian diagonal estimates.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
hessian_power float

Exponent applied to the root mean square Hessian diagonal estimate.

1.0
update_period int

Number of steps after which to apply the Hessian approximation.

1
num_samples int

Number of noise samples for each Hessian diagonal estimate.

1
hessian_distribution HutchinsonG

Type of distribution used to initialize the Hutchinson trace estimator.

'rademacher'
eps float

Term added to the denominator to improve numerical stability.

1e-16
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adahessian.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
class AdaHessian(BaseOptimizer):
    """Adaptive second-order updates using Hutchinson Hessian estimates.

    Use `loss.backward(create_graph=True)` for internal Hessian estimation, or supply
    external estimates through `step(hessian=...)`.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the gradient mean and squared Hessian diagonal estimates.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        hessian_power: Exponent applied to the root mean square Hessian diagonal estimate.
        update_period: Number of steps after which to apply the Hessian approximation.
        num_samples: Number of noise samples for each Hessian diagonal estimate.
        hessian_distribution: Type of distribution used to initialize the Hutchinson trace estimator.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-1,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        hessian_power: float = 1.0,
        update_period: int = 1,
        num_samples: int = 1,
        hessian_distribution: HutchinsonG = 'rademacher',
        eps: float = 1e-16,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')
        self.validate_range(hessian_power, 'Hessian Power', 0, 1, range_type='(]')

        self.update_period = update_period
        self.num_samples = num_samples
        self.distribution = hessian_distribution
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'hessian_power': hessian_power,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdaHessian'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if 'exp_avg' not in state:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_hessian_diag_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None, hessian: list[torch.Tensor] | None = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)

        step: int = self.param_groups[0]['step'] + 1
        update_hessian = hessian is not None or (step - 1) % self.update_period == 0

        if hessian is not None:
            self.set_hessian(self.param_groups, self.state, hessian)
        elif update_hessian:
            self.zero_hessian(self.param_groups, self.state)
            self.compute_hutchinson_hessian(
                param_groups=self.param_groups,
                state=self.state,
                num_samples=self.num_samples,
                distribution=self.distribution,
            )

        for group in self.param_groups:
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2: float = self.debias(beta2, group['step'])

            step_size: float = self.apply_adam_debias(group.get('adam_debias', False), group['lr'], bias_correction1)

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                exp_avg, exp_hessian_diag_sq = state['exp_avg'], state['exp_hessian_diag_sq']
                exp_avg.lerp_(grad, weight=1.0 - beta1)

                if 'hessian' in state and update_hessian:
                    exp_hessian_diag_sq.mul_(beta2).addcmul_(state['hessian'], state['hessian'], value=1.0 - beta2)

                de_nom = (exp_hessian_diag_sq / bias_correction2).pow_(group['hessian_power'] / 2).add_(group['eps'])

                p.addcdiv_(exp_avg, de_nom, value=-step_size)

        return loss

Adai

Bases: BaseOptimizer

SGD with gradient dependent momentum and optional stable weight decay.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Adaptive momentum scaling coefficient and squared gradient decay rate.

(0.1, 0.99)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
stable_weight_decay bool

Perform stable weight decay.

False
dampening float

Dampening factor for momentum.

1.0
eps float

Term added to the denominator to improve numerical stability.

0.001
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adai.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
class Adai(BaseOptimizer):
    """SGD with gradient dependent momentum and optional stable weight decay.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Adaptive momentum scaling coefficient and squared gradient decay rate.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        stable_weight_decay: Perform stable weight decay.
        dampening: Dampening factor for momentum.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.1, 0.99),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        stable_weight_decay: bool = False,
        dampening: float = 1.0,
        eps: float = 1e-3,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'stable_weight_decay': stable_weight_decay,
            'dampening': dampening,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Adai'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['beta1_prod'] = torch.ones_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        param_size: int = 0
        exp_avg_sq_hat_sum: float = 0.0

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            _, beta2 = group['betas']

            bias_correction2: float = self.debias(beta2, group['step'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                param_size += p.numel()

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                if group.get('use_gc'):
                    centralize_gradient(grad, gc_conv_only=False)

                if not group['stable_weight_decay'] and group['weight_decay'] > 0.0:
                    self.apply_weight_decay(
                        p=p,
                        grad=grad,
                        lr=group['lr'],
                        weight_decay=group['weight_decay'],
                        weight_decouple=group['weight_decouple'],
                        fixed_decay=group['fixed_decay'],
                    )

                exp_avg_sq = state['exp_avg_sq']
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                exp_avg_sq_hat_sum += exp_avg_sq.sum() / bias_correction2

        if param_size == 0:
            raise ZeroParameterSizeError

        exp_avg_sq_hat_mean = (exp_avg_sq_hat_sum / param_size).clamp_min_(self.defaults['eps'])

        for group in self.param_groups:
            beta0, beta2 = group['betas']

            beta0_dp: float = math.pow(beta0, 1.0 - group['dampening'])
            bias_correction2: float = self.debias(beta2, group['step'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                state = self.state[p]

                if group['stable_weight_decay'] and group['weight_decay'] > 0.0:
                    self.apply_weight_decay(
                        p=p,
                        grad=grad,
                        lr=group['lr'],
                        weight_decay=group['weight_decay'],
                        weight_decouple=group['weight_decouple'],
                        fixed_decay=group['fixed_decay'],
                    )

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                exp_avg_sq_hat = exp_avg_sq / bias_correction2
                beta1 = (
                    1.0
                    - (exp_avg_sq_hat / exp_avg_sq_hat_mean).pow_(1.0 / (3.0 - 2.0 * group['dampening'])).mul_(beta0)
                ).clamp_(0.0, 1.0 - group['eps'])
                beta3 = (1.0 - beta1).pow_(group['dampening'])

                beta1_prod = state['beta1_prod']
                beta1_prod.mul_(beta1)

                exp_avg.mul_(beta1).addcmul_(beta3, grad)
                exp_avg_hat = exp_avg.div(1.0 - beta1_prod).mul_(beta0_dp)

                p.add_(exp_avg_hat, alpha=-group['lr'])

        return loss

Adalite

Bases: BaseOptimizer

Adaptive updates with factored moments and tensor wise trust ratios.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.01
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
g_norm_min float

Lower bound for the gradient norm in the trust ratio.

1e-10
ratio_min float

Lower bound for the parameter-to-gradient norm ratio.

0.0001
tau float

Softmax temperature for row and column importance weights.

1.0
eps1 float

Stability constant for adaptive updates.

1e-06
eps2 float

Lower bound for factored moment normalization denominators.

1e-10
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adalite.py
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
class Adalite(BaseOptimizer):
    """Adaptive updates with factored moments and tensor wise trust ratios.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        g_norm_min: Lower bound for the gradient norm in the trust ratio.
        ratio_min: Lower bound for the parameter-to-gradient norm ratio.
        tau: Softmax temperature for row and column importance weights.
        eps1: Stability constant for adaptive updates.
        eps2: Lower bound for factored moment normalization denominators.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 1e-2,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        g_norm_min: float = 1e-10,
        ratio_min: float = 1e-4,
        tau: float = 1.0,
        eps1: float = 1e-6,
        eps2: float = 1e-10,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps1, 'eps1')
        self.validate_non_negative(eps2, 'eps2')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'g_norm_min': g_norm_min,
            'ratio_min': ratio_min,
            'tau': tau,
            'eps1': eps1,
            'eps2': eps2,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Adalite'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                if len(p.shape) < 2:
                    state['m_avg'] = torch.zeros_like(p)
                    state['v_avg'] = torch.zeros_like(p)
                else:
                    state['v_avg_0'] = torch.zeros_like(p.mean(dim=1))
                    state['v_avg_1'] = torch.zeros_like(p.mean(dim=0))

                    state['m_avg_c'] = torch.zeros_like(p.mean(dim=1)[:, None])
                    state['m_avg_r'] = torch.zeros_like(p.mean(dim=0)[None, :])
                    state['m_avg_u'] = torch.zeros_like(p.mean().unsqueeze(0).unsqueeze(0))

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                if sum(grad.shape) > 1:
                    trust_ratio = (p.norm() / grad.norm().clip(min=group['g_norm_min'])).clip(min=group['ratio_min'])
                    grad.mul_(trust_ratio)

                if len(grad.shape) < 2:
                    m = state['m_avg']
                    v = state['v_avg']
                else:
                    r, c = state['v_avg_0'][:, None], state['v_avg_1'][None, :]
                    v = (r * c) / r.sum().clamp(min=group['eps2'])
                    m = state['m_avg_c'] @ state['m_avg_u'] @ state['m_avg_r']

                m.lerp_(grad, 1.0 - beta1)
                v.lerp_((grad - m).square(), 1.0 - beta2)

                v_avg = v / (1.0 - beta2 ** group['step'])

                if len(grad.shape) == 2:
                    imp_c = softmax(v.mean(dim=1), dim=0)[:, None]
                    imp_r = softmax(v.mean(dim=0), dim=0)[None, :]
                    m.lerp_(grad, 1.0 - imp_c * imp_r)

                u = m.lerp(grad, 1.0 - beta1)

                if len(grad.shape) < 2:
                    state['m_avg'] = m
                    state['v_avg'] = v
                else:
                    state['v_avg_0'] = v.sum(dim=1)
                    state['v_avg_1'] = v.sum(dim=0) / v.sum().clamp(min=group['eps2'])

                    imp_c = softmax(v.mean(dim=1) / group['tau'], dim=-1)[:, None]
                    imp_r = softmax(v.mean(dim=0) / group['tau'], dim=-1)[None, :]

                    c = ((m * imp_r).sum(dim=1))[:, None]
                    r = ((m * imp_c).sum(dim=0))[None, :]

                    s = (c.T @ m @ r.T) / (c.T @ c @ r @ r.T).clamp(min=group['eps2'])

                    state['m_avg_c'] = c
                    state['m_avg_r'] = r
                    state['m_avg_u'] = s

                u.div_((v_avg + group['eps1']).sqrt())

                u = u.reshape(p.shape)
                u.add_(p, alpha=group['weight_decay'])

                p.add_(u, alpha=-group['lr'])

        return loss

AdaLOMO

Bases: BaseOptimizer

Factored adaptive updates fused into backward.

Parameters:

Name Type Description Default
model Module

PyTorch model.

required
lr float

Learning rate.

0.001
weight_decay float

Weight decay coefficient.

0.0
loss_scale float

Multiplier applied before backward and removed from gradients. 0 disables scaling.

2.0 ** 10
clip_threshold float

Maximum root mean square of the preconditioned update.

1.0
decay_rate float

Exponent controlling the step-dependent second moment decay.

-0.8
clip_grad_norm float | None

Clip gradient norm.

None
clip_grad_value float | None

Clip gradient value.

None
eps1 float

Stability constant added to squared gradients.

1e-30
eps2 float

Lower bound for parameter RMS scaling of the learning rate.

0.001
Source code in pytorch_optimizer/optimizer/lomo.py
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
class AdaLOMO(BaseOptimizer):
    """Factored adaptive updates fused into backward.

    Args:
        model: PyTorch model.
        lr: Learning rate.
        weight_decay: Weight decay coefficient.
        loss_scale: Multiplier applied before backward and removed from gradients. `0` disables scaling.
        clip_threshold: Maximum root mean square of the preconditioned update.
        decay_rate: Exponent controlling the step-dependent second moment decay.
        clip_grad_norm: Clip gradient norm.
        clip_grad_value: Clip gradient value.
        eps1: Stability constant added to squared gradients.
        eps2: Lower bound for parameter RMS scaling of the learning rate.

    """

    def __init__(
        self,
        model: nn.Module,
        lr: float = 1e-3,
        weight_decay: float = 0.0,
        loss_scale: float = 2.0 ** 10,
        clip_threshold: float = 1.0,
        decay_rate: float = -0.8,
        clip_grad_norm: float | None = None,
        clip_grad_value: float | None = None,
        eps1: float = 1e-30,
        eps2: float = 1e-3,
        **kwargs,
    ) -> None:  # fmt: skip
        self.validate_learning_rate(lr)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(loss_scale, 'loss_scale')
        self.validate_non_negative(clip_threshold, 'clip_threshold')
        self.validate_non_negative(clip_grad_norm, 'clip_grad_norm')
        self.validate_non_negative(clip_grad_value, 'clip_grad_value')
        self.validate_non_negative(eps1, 'eps1')
        self.validate_non_negative(eps2, 'eps2')

        self.model = model
        self.lr = lr
        self.weight_decay = weight_decay
        self.loss_scale = loss_scale
        self.clip_threshold = clip_threshold
        self.decay_rate = decay_rate
        self.clip_grad_norm = clip_grad_norm
        self.clip_grad_value = clip_grad_value
        self.eps1 = eps1
        self.eps2 = eps2

        self.num_steps: int = 0
        self.gather_norm: bool = False
        self.grad_norms: list[torch.Tensor] = []
        self.clip_coef: float | torch.Tensor | None = None

        self.local_rank: int = int(os.environ.get('LOCAL_RANK', '0'))
        self.zero3_enabled: bool = is_deepspeed_zero3_enabled()

        self.grad_func: Callable[[Any], Any] = self.fuse_update_zero3() if self.zero3_enabled else self.fuse_update()

        self.exp_avg_sq = {}
        self.exp_avg_sq_row = {}
        self.exp_avg_sq_col = {}

        self.initialize_states()

        defaults: Defaults = {
            'lr': lr,
            'weight_decay': weight_decay,
            'clip_grad_norm': clip_grad_norm,
            'clip_grad_value': clip_grad_value,
            'eps1': eps1,
            'eps2': eps2,
        }

        super().__init__(self.model.parameters(), defaults)

    def __str__(self) -> str:
        return 'AdaLOMO'

    def state_dict(self) -> dict:
        state = super().state_dict()
        state['num_steps'] = self.num_steps
        state['exp_avg_sq'] = self.exp_avg_sq
        state['exp_avg_sq_row'] = self.exp_avg_sq_row
        state['exp_avg_sq_col'] = self.exp_avg_sq_col
        return state

    def load_state_dict(self, state_dict: dict) -> None:
        super().load_state_dict(state_dict)
        self.num_steps = state_dict.get('num_steps', 0)
        with torch.no_grad():
            self.exp_avg_sq = {
                key: self.exp_avg_sq[key].copy_(value)
                for key, value in state_dict.get('exp_avg_sq', self.exp_avg_sq).items()
            }
            self.exp_avg_sq_row = {
                key: self.exp_avg_sq_row[key].copy_(value)
                for key, value in state_dict.get('exp_avg_sq_row', self.exp_avg_sq_row).items()
            }
            self.exp_avg_sq_col = {
                key: self.exp_avg_sq_col[key].copy_(value)
                for key, value in state_dict.get('exp_avg_sq_col', self.exp_avg_sq_col).items()
            }

    def initialize_states(self) -> None:
        for n, p in self.model.named_parameters():
            if self.zero3_enabled:  # pragma: no cover
                if len(p.ds_shape) == 1:
                    self.exp_avg_sq[n] = torch.zeros(p.ds_shape[0], dtype=torch.float32, device=p.device)
                else:
                    self.exp_avg_sq_row[n] = torch.zeros(p.ds_shape[0], dtype=torch.float32, device=p.device)
                    self.exp_avg_sq_col[n] = torch.zeros(p.ds_shape[1], dtype=torch.float32, device=p.device)
            elif len(p.shape) == 1:
                self.exp_avg_sq[n] = torch.zeros(p.shape[0], dtype=torch.float32, device=p.device)
            else:
                self.exp_avg_sq_row[n] = torch.zeros(p.shape[0], dtype=torch.float32, device=p.device)
                self.exp_avg_sq_col[n] = torch.zeros(p.shape[1], dtype=torch.float32, device=p.device)

            if p.requires_grad:
                p.register_hook(self.grad_func)

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

    def fuse_update(self) -> Callable[[Any], Any]:
        @torch.no_grad()
        def func(x: Any) -> Any:
            for n, p in self.model.named_parameters():
                if not p.requires_grad or p.grad is None:
                    continue

                grad_fp32 = p.grad.to(torch.float32)
                p.grad = None

                if self.loss_scale:
                    grad_fp32.div_(self.loss_scale)

                if self.gather_norm:
                    self.grad_norms.append(torch.norm(grad_fp32, 2.0))
                else:
                    if self.clip_grad_value is not None and self.clip_grad_value > 0.0:
                        grad_fp32.clamp_(min=-self.clip_grad_value, max=self.clip_grad_value)
                    if self.clip_grad_norm is not None and self.clip_grad_norm > 0.0 and self.clip_coef is not None:
                        grad_fp32.mul_(self.clip_coef)

                    beta2_t: float = 1.0 - math.pow(
                        self.num_steps, self.decay_rate if self.num_steps > 0 else -self.decay_rate
                    )

                    update = grad_fp32.pow(2).add_(self.eps1)

                    if len(p.shape) > 1:
                        self.exp_avg_sq_row[n].lerp_(update.mean(dim=-1), weight=1.0 - beta2_t)
                        self.exp_avg_sq_col[n].lerp_(update.mean(dim=-2), weight=1.0 - beta2_t)

                        self.approximate_sq_grad(self.exp_avg_sq_row[n], self.exp_avg_sq_col[n], update)
                        update.mul_(grad_fp32)
                    else:
                        self.exp_avg_sq[n].lerp_(update, weight=1.0 - beta2_t)
                        update = self.exp_avg_sq[n].rsqrt().mul_(grad_fp32)

                    factor = cast(torch.Tensor, self.get_rms(update)).div_(self.clip_threshold).clamp_min_(1.0)
                    update.div_(factor)

                    p_fp32 = p.to(torch.float32)
                    p_rms = torch.norm(p_fp32, 2.0) / math.sqrt(p.numel())

                    lr = self.lr * max(self.eps2, p_rms)

                    self.apply_weight_decay(
                        p_fp32,
                        grad_fp32,
                        lr,
                        self.weight_decay,
                        weight_decouple=True,
                        fixed_decay=False,
                    )

                    p_fp32.add_(update, alpha=-lr)
                    p.copy_(p_fp32)

            return x

        return func

    def fuse_update_zero3(self) -> Callable[[Any], Any]:  # pragma: no cover
        @torch.no_grad()
        def func(x: torch.Tensor) -> torch.Tensor:
            for n, p in self.model.named_parameters():
                if p.grad is None:
                    continue

                all_reduce(p.grad, op=ReduceOp.AVG, async_op=False)

                grad_fp32 = p.grad.to(torch.float32)
                p.grad = None

                if self.loss_scale:
                    grad_fp32.div_(self.loss_scale)

                start: int = 0
                end: int = grad_fp32.numel()
                if self.gather_norm:
                    self.grad_norms.append(torch.norm(grad_fp32, 2.0))
                else:
                    partition_size: int = p.ds_tensor.numel()
                    start = partition_size * self.local_rank
                    end = min(start + partition_size, grad_fp32.numel())

                if self.clip_grad_value is not None:
                    grad_fp32.clamp_(min=-self.clip_grad_value, max=self.clip_grad_value)
                if self.clip_grad_norm is not None and self.clip_grad_norm > 0 and self.clip_coef is not None:
                    grad_fp32.mul_(self.clip_coef)

                beta2_t: float = 1.0 - math.pow(
                    self.num_steps, self.decay_rate if self.num_steps > 0 else -self.decay_rate
                )

                update = grad_fp32.pow(2).add_(self.eps1)

                if len(p.ds_shape) > 1:
                    self.exp_avg_sq_row[n].mul_(beta2_t).add_(update.mean(dim=-1), alpha=1.0 - beta2_t)
                    self.exp_avg_sq_col[n].mul_(beta2_t).add_(update.mean(dim=-2), alpha=1.0 - beta2_t)

                    self.approximate_sq_grad(self.exp_avg_sq_row[n], self.exp_avg_sq_col[n], update)
                    update.mul_(grad_fp32)
                else:
                    self.exp_avg_sq[n].mul_(beta2_t).add_(update, alpha=1.0 - beta2_t)
                    update = self.exp_avg_sq[n].rsqrt().mul_(grad_fp32)

                factor = cast(torch.Tensor, self.get_rms(update)).div_(self.clip_threshold).clamp_min_(1.0)
                update.div_(factor)

                one_dim_update = update.view(-1)
                partitioned_update = one_dim_update.narrow(0, start, end - start)

                param_fp32 = p.ds_tensor.to(torch.float32)
                partitioned_p = param_fp32.narrow(0, 0, end - start)

                p_rms = torch.norm(partitioned_p, 2.0).pow_(2)
                all_reduce(p_rms, op=ReduceOp.SUM)
                p_rms.div_(p.ds_numel).sqrt_()

                lr = self.lr * max(self.eps2, p_rms)

                self.apply_weight_decay(
                    p=partitioned_p,
                    grad=grad_fp32,
                    lr=lr,
                    weight_decay=self.weight_decay,
                    weight_decouple=True,
                    fixed_decay=False,
                )

                partitioned_p.add_(partitioned_update, alpha=-lr)

                p.ds_tensor[:end - start] = partitioned_p  # fmt: skip

            return x

        return func

    def fused_backward(self, loss, lr: float) -> None:
        self.lr = lr

        if self.loss_scale:
            loss = loss * self.loss_scale

        self.num_steps += 1

        loss.backward()

        self.grad_func(0)

    def grad_norm(self, loss) -> None:
        self.gather_norm = True
        self.grad_norms = []

        if self.loss_scale:
            loss = loss * self.loss_scale

        loss.backward(retain_graph=True)

        self.grad_func(0)

        with torch.no_grad():
            grad_norms = torch.stack(self.grad_norms)

            total_norm = torch.norm(grad_norms, 2.0)
            self.clip_coef = torch.clamp(float(self.clip_grad_norm) / (total_norm + 1e-6), max=1.0)

        self.gather_norm = False

AdaMax

Bases: BaseOptimizer

Adam with an exponentially weighted infinity norm denominator.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the gradient mean and exponentially weighted infinity norm.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
Source code in pytorch_optimizer/optimizer/adamax.py
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
class AdaMax(BaseOptimizer):
    """Adam with an exponentially weighted infinity norm denominator.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the gradient mean and exponentially weighted infinity norm.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.

    """

    _supports_compiled_foreach = True

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        foreach: bool | None = None,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            'foreach': foreach,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdaMax'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_inf'] = torch.zeros_like(p)

                if group.get('adanorm'):
                    state['exp_grad_adanorm'] = torch.zeros((1,), dtype=grad.dtype, device=grad.device)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        return not group.get('adanorm') and self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_infs: list[torch.Tensor],
        step_size: float | torch.Tensor,
    ) -> None:
        beta1, beta2 = group['betas']

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=group['lr'],
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - beta1)

        torch._foreach_mul_(exp_infs, beta2)

        grad_abs = torch._foreach_abs(grads)
        torch._foreach_add_(grad_abs, group['eps'])
        torch._foreach_maximum_(exp_infs, grad_abs)

        foreach_addcdiv_(params, exp_avgs, exp_infs, value=-step_size)

    def _step_per_param(self, group: ParamGroup, step_size: float) -> None:
        beta1, beta2 = group['betas']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            exp_avg, exp_inf = state['exp_avg'], state['exp_inf']

            p, grad, exp_avg, exp_inf = self.view_as_real(p, grad, exp_avg, exp_inf)

            self.apply_weight_decay(
                p=p,
                grad=grad,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=group['fixed_decay'],
            )

            s_grad = self.get_adanorm_gradient(
                grad=grad,
                adanorm=group.get('adanorm', False),
                exp_grad_norm=state.get('exp_grad_adanorm', None),
                r=group.get('adanorm_r', None),
            )

            exp_avg.lerp_(s_grad, weight=1.0 - beta1)

            torch.maximum(exp_inf.mul_(beta2), grad.abs().add_(group['eps']), out=exp_inf)

            p.addcdiv_(exp_avg, exp_inf, value=-step_size)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, _ = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])

            step_size: float = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=group['lr'],
                bias_correction1=bias_correction1,
            )

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_inf']
                )

                for tensors in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(
                        group, tensors['params'], tensors['grads'], tensors['exp_avg'], tensors['exp_inf'], step_size
                    )
            else:
                self._step_per_param(group, step_size)

        return loss

AdamC

Bases: BaseOptimizer

Adam with weight decay scaled by the learning rate in normalization layers.

Set normalized=True for LayerNorm and BatchNorm layers.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adamc.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
class AdamC(BaseOptimizer):
    """Adam with weight decay scaled by the learning rate in normalization layers.

    Set `normalized=True` for LayerNorm and BatchNorm layers.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        ams_bound: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize
        self.max_lr: float = lr

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'ams_bound': ams_bound,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdamC'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

                if group['ams_bound']:
                    state['max_exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            wd_step_size: float = group['lr'] if not group.get('normalized') else (group['lr'] ** 2) / self.max_lr

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=wd_step_size,
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                de_nom = self.apply_ams_bound(
                    ams_bound=group['ams_bound'],
                    exp_avg_sq=exp_avg_sq,
                    max_exp_avg_sq=state.get('max_exp_avg_sq', None),
                    eps=group['eps'],
                )
                de_nom.div_(bias_correction2_sq)

                p.addcdiv_(exp_avg / bias_correction1, de_nom, value=-group['lr'])

        return loss

AdamG

Bases: BaseOptimizer

Parameter free adaptive updates with gradient dependent scaling.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

1.0
betas Betas

Decay rates for scaled gradient momentum, squared gradients, and the numerator scale.

(0.95, 0.999, 0.95)
p float

The p value in the numerator function s(x) = p * x^q.

0.2
q float

The q value in the numerator function s(x) = p * x^q.

0.24
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adamg.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
class AdamG(BaseOptimizer):
    """Parameter free adaptive updates with gradient dependent scaling.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for scaled gradient momentum, squared gradients, and the numerator scale.
        p: The p value in the numerator function `s(x) = p * x^q`.
        q: The q value in the numerator function `s(x) = p * x^q`.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1.0,
        betas: Betas = (0.95, 0.999, 0.95),
        p: float = 0.2,
        q: float = 0.24,
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_positive(p, 'p')
        self.validate_positive(q, 'q')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.p = p
        self.q = q
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdamG'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['m'] = torch.zeros_like(p)
                state['v'] = torch.zeros_like(p)
                state['r'] = torch.zeros_like(p)

    def s(self, p: torch.Tensor) -> torch.Tensor:
        """Compute the numerator scaling function `p * x ** q`."""
        return p.pow(self.q).mul_(self.p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2, beta3 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2: float = self.debias(beta2, group['step'])

            step_size: float = min(group['lr'], 1.0 / math.sqrt(group['step']))

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                m, v, r = state['m'], state['v'], state['r']

                p, grad, m, v, r = self.view_as_real(p, grad, m, v, r)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                v.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)
                r.lerp_(self.s(v), weight=1.0 - beta3)
                m.mul_(beta1).addcmul_(r, grad, value=1.0 - beta1)

                update = (m / bias_correction1) / (v / bias_correction2).sqrt_().add_(group['eps'])

                p.add_(update, alpha=-step_size)

        return loss

s(p)

Compute the numerator scaling function p * x ** q.

Source code in pytorch_optimizer/optimizer/adamg.py
86
87
88
def s(self, p: torch.Tensor) -> torch.Tensor:
    """Compute the numerator scaling function `p * x ** q`."""
    return p.pow(self.q).mul_(self.p)

AdamMini

Bases: BaseOptimizer

Adam with shared second moment estimates within parameter blocks.

Parameters:

Name Type Description Default
model Module

Model instance.

required
model_sharding bool

Set to True if you are using model parallelism with more than 1 GPU, including FSDP and zero_1, zero_2, zero_3 in DeepSpeed. Set to False otherwise.

False
lr float

Learning rate.

1.0
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.1
num_embeds int

Number of embedding dimensions. Could be unspecified if training non transformer models.

2048
num_heads int

Number of attention heads. Could be unspecified if training non transformer models.

32
num_query_groups int | None

Number of query groups in Group Query Attention (GQA). If not specified, defaults to num_heads. Could be unspecified for non transformer models.

None
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adam_mini.py
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
class AdamMini(BaseOptimizer):  # pragma: no cover
    """Adam with shared second moment estimates within parameter blocks.

    Args:
        model: Model instance.
        model_sharding: Set to True if you are using model parallelism with more than 1 GPU, including FSDP and
            zero_1, zero_2, zero_3 in DeepSpeed. Set to False otherwise.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        num_embeds: Number of embedding dimensions. Could be unspecified if training non transformer models.
        num_heads: Number of attention heads. Could be unspecified if training non transformer models.
        num_query_groups: Number of query groups in Group Query Attention (GQA). If not specified, defaults to
            num_heads. Could be unspecified for non transformer models.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        model: nn.Module,
        lr: float = 1.0,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.1,
        model_sharding: bool = False,
        num_embeds: int = 2048,
        num_heads: int = 32,
        num_query_groups: int | None = None,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_positive(num_embeds, 'num_embeds')
        self.validate_positive(num_heads, 'num_heads')
        self.validate_non_negative(eps, 'eps')

        self.num_query_groups: int = num_query_groups if num_query_groups is not None else num_heads
        self.validate_positive(self.num_query_groups, 'num_query_groups')
        self.validate_mod(num_embeds, num_heads)
        self.validate_mod(num_heads, self.num_query_groups)

        # Visible GPUs are not a process group. all_gather below requires
        # dist to be initialized; otherwise a single process that can see
        # several devices calls it and raises.
        if dist.is_available() and dist.is_initialized():
            self.world_size: int = dist.get_world_size()
        else:
            self.world_size = 1

        self.model = model
        self.model_sharding = model_sharding
        self.num_embeds = num_embeds
        self.num_heads = num_heads

        self.embed_blocks: set[str] = {'embed', 'embd', 'wte', 'lm_head.weight', 'output.weight'}
        self.qk_blocks: set[str] = {'k_proj.weight', 'q_proj.weight', 'wq.weight', 'wk.weight'}

        self.maximize = maximize

        groups = self.get_optimizer_groups(weight_decay)

        defaults: Defaults = {'lr': lr, 'betas': betas, 'eps': eps, **kwargs}

        super().__init__(groups, defaults)

    def __str__(self) -> str:
        return 'AdamMini'

    def load_state_dict(self, state_dict: dict) -> None:
        super().load_state_dict(state_dict)
        for group, saved_group in zip(self.param_groups, state_dict['param_groups']):
            for p, key in zip(group['params'], saved_group['params']):
                for name, value in state_dict['state'].get(key, {}).items():
                    if isinstance(value, torch.Tensor) and value.is_floating_point():
                        self.state[p][name] = value.to(device=p.device, dtype=torch.float32)

    def get_optimizer_groups(self, weight_decay: float):
        groups = []
        for name, param in self.model.named_parameters():
            if not param.requires_grad:
                continue

            group = {
                'name': name,
                'params': param,
                'weight_decay': 0.0 if ('norm' in name or 'ln_f' in name) else weight_decay,
            }

            if any(block in name for block in self.qk_blocks):
                group['parameter_per_head'] = self.num_embeds * self.num_embeds // self.num_heads

            groups.append(group)

        return groups

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

    @staticmethod
    def step_embed(
        p,
        grad,
        state,
        lr: float,
        beta1: float,
        beta2: float,
        bias_correction1: float,
        bias_correction2_sq: float,
        eps: float,
    ) -> None:
        if len(state) == 0:
            state['m'] = torch.zeros_like(p, dtype=torch.float32)
            state['v'] = torch.zeros_like(p, dtype=torch.float32)

        m, v = state['m'], state['v']

        m.lerp_(grad, weight=1.0 - beta1)
        v.mul_(beta2).addcmul_(grad, grad.conj(), value=1.0 - beta2)

        h = (v.sqrt() / bias_correction2_sq).add_(eps)

        # PyTorch 2.1 CPU addcdiv does not support mixed precision inputs.
        if p.device.type == 'cpu' and p.dtype != m.dtype:
            p.add_(m / h, alpha=-lr / bias_correction1)
        else:
            p.addcdiv_(m, h, value=-lr / bias_correction1)

    @staticmethod
    def step_attn_proj(
        p,
        grad,
        state,
        parameter_per_head: int,
        lr: float,
        beta1: float,
        beta2: float,
        bias_correction1: float,
        bias_correction2_sq: float,
        eps: float,
    ) -> None:
        if len(state) == 0:
            state['m'] = torch.zeros_like(p, dtype=torch.float32).view(-1, parameter_per_head)
            state['head'] = state['m'].shape[0]
            state['v_mean'] = torch.zeros(state['head'], device=state['m'].device)

        m, v = state['m'], state['v_mean']

        head: int = state['head']
        grad = grad.view(head, parameter_per_head)

        m.lerp_(grad, weight=1.0 - beta1)

        tmp_lr = torch.mean(grad * grad, dim=1).to(m.device)
        v.lerp_(tmp_lr.to(dtype=v.dtype), weight=1.0 - beta2)

        h = (v.sqrt() / bias_correction2_sq).add_(eps)

        update = (1 / (h * bias_correction1)).view(head, 1).mul(m)

        if p.dim() > 1:
            d0, d1 = p.size()
            update = update.view(d0, d1)
        else:
            update = update.view(-1)

        p.add_(update, alpha=-lr)

    @staticmethod
    def step_attn(
        p,
        grad,
        state,
        num_query_groups: int,
        q_per_kv: int,
        lr: float,
        beta1: float,
        beta2: float,
        bias_correction1: float,
        bias_correction2_sq: float,
        eps: float,
    ) -> None:
        if len(state) == 0:
            state['m'] = torch.zeros_like(p, dtype=torch.float32).view(num_query_groups, q_per_kv + 2, -1)
            state['v_mean'] = torch.zeros(num_query_groups, q_per_kv + 2, device=state['m'].device)

        m, v = state['m'], state['v_mean']

        grad = grad.view(num_query_groups, q_per_kv + 2, -1)

        m.lerp_(grad, weight=1.0 - beta1)

        tmp_lr = torch.mean(grad * grad, dim=2).to(m.device)
        v.lerp_(tmp_lr.to(dtype=v.dtype), weight=1.0 - beta2)

        h = (v.sqrt() / bias_correction2_sq).add_(eps)

        update = m / (h * bias_correction1).unsqueeze(-1)

        if p.dim() > 1:
            d0, d1 = p.size()
            update = update.view(d0, d1)
        else:
            update = update.view(-1)

        p.add_(update, alpha=-lr)

    def step_lefts(
        self,
        p,
        grad,
        state,
        lr: float,
        beta1: float,
        beta2: float,
        bias_correction1: float,
        bias_correction2_sq: float,
        eps: float,
    ) -> None:
        if len(state) == 0:
            dim = torch.tensor(p.numel(), device=p.device, dtype=torch.float32)

            reduced: bool = False
            if self.model_sharding and self.world_size > 1:
                tensor_list = [torch.zeros_like(dim) for _ in range(self.world_size)]
                dist.all_gather(tensor_list, dim)

                s, dim = 0, 0
                for d in tensor_list:
                    if d > 0:
                        s += 1
                    dim += d

                if s >= 2:
                    reduced = True

            state['m'] = torch.zeros_like(p, dtype=torch.float32)
            state['v_mean'] = torch.tensor(0.0, device=state['m'].device)
            state['dimension'] = dim
            state['reduced'] = reduced

        tmp_lr = torch.sum(grad * grad)

        if state['reduced']:
            dist.all_reduce(tmp_lr, op=dist.ReduceOp.SUM)

        tmp_lr.div_(state['dimension'])

        m, v = state['m'], state['v_mean']

        m.lerp_(grad, weight=1.0 - beta1)
        v.lerp_(tmp_lr.to(dtype=v.dtype), weight=1.0 - beta2)

        h = (v.sqrt() / bias_correction2_sq).add_(eps)

        stepsize = (1 / bias_correction1) / h

        update = m * stepsize

        p.add_(update, alpha=-lr)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            name = group['name']

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2: float = self.debias(beta2, group['step'])
            bias_correction2_sq: float = math.sqrt(bias_correction2)

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad
                if grad.is_sparse:
                    raise NoSparseGradientError(str(self))

                if torch.is_complex(p):
                    raise NoComplexParameterError(str(self))

                grad = grad.to(torch.float32)

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=True,
                    fixed_decay=False,
                )

                if any(block in name for block in self.embed_blocks):
                    self.step_embed(
                        p, grad, state, group['lr'], beta1, beta2, bias_correction1, bias_correction2_sq, group['eps']
                    )
                elif any(block in name for block in self.qk_blocks):
                    self.step_attn_proj(
                        p,
                        grad,
                        state,
                        group['parameter_per_head'],
                        group['lr'],
                        beta1,
                        beta2,
                        bias_correction1,
                        bias_correction2_sq,
                        group['eps'],
                    )
                elif 'attn.attn.weight' in name or 'attn.qkv.weight' in name:
                    self.step_attn(
                        p,
                        grad,
                        state,
                        self.num_query_groups,
                        self.num_heads // self.num_query_groups,
                        group['lr'],
                        beta1,
                        beta2,
                        bias_correction1,
                        bias_correction2_sq,
                        group['eps'],
                    )
                else:
                    self.step_lefts(
                        p,
                        grad,
                        state,
                        group['lr'],
                        beta1,
                        beta2,
                        bias_correction1,
                        bias_correction2_sq,
                        group['eps'],
                    )

        return loss

AdaMod

Bases: BaseOptimizer

Adam with bounds from a moving average of adaptive learning rates.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the gradient mean, squared gradients, and adaptive learning rates.

(0.9, 0.99, 0.9999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
Source code in pytorch_optimizer/optimizer/adamod.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
class AdaMod(BaseOptimizer):
    """Adam with bounds from a moving average of adaptive learning rates.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the gradient mean, squared gradients, and adaptive learning rates.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.

    """

    _supports_compiled_foreach = True

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.99, 0.9999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        foreach: bool | None = None,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            'foreach': foreach,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdaMod'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['exp_avg_lr'] = torch.zeros_like(p)

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        exp_avg_lrs: list[torch.Tensor],
        step_size: float | torch.Tensor,
    ) -> None:
        beta1, beta2, beta3 = group['betas']

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=group['lr'],
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - beta1)

        torch._foreach_mul_(exp_avg_sqs, beta2)
        torch._foreach_addcmul_(exp_avg_sqs, grads, grads, value=1.0 - beta2)

        de_noms = torch._foreach_sqrt(exp_avg_sqs)
        torch._foreach_add_(de_noms, group['eps'])

        updates = de_noms
        foreach_scalar_div_(updates, step_size)
        torch._foreach_lerp_(exp_avg_lrs, updates, weight=1.0 - beta3)
        torch._foreach_minimum_(updates, exp_avg_lrs)
        torch._foreach_mul_(updates, exp_avgs)

        torch._foreach_sub_(params, updates)

    def _step_per_param(self, group: ParamGroup, step_size: float) -> None:
        beta1, beta2, beta3 = group['betas']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            exp_avg, exp_avg_sq, exp_avg_lr = state['exp_avg'], state['exp_avg_sq'], state['exp_avg_lr']

            p, grad, exp_avg, exp_avg_sq, exp_avg_lr = self.view_as_real(p, grad, exp_avg, exp_avg_sq, exp_avg_lr)

            self.apply_weight_decay(
                p=p,
                grad=grad,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=group['fixed_decay'],
            )

            exp_avg.lerp_(grad, weight=1.0 - beta1)

            exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

            de_nom = exp_avg_sq.sqrt().add_(group['eps'])

            update = de_nom
            foreach_scalar_div_([update], step_size)

            exp_avg_lr.lerp_(update, weight=1.0 - beta3)

            torch.min(update, exp_avg_lr, out=update)
            update.mul_(exp_avg)

            p.sub_(update)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2, _ = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=group['lr'] * bias_correction2_sq,
                bias_correction1=bias_correction1,
            )

            if self.can_use_foreach(group, group.get('foreach')):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_avg_sq', 'exp_avg_lr']
                )

                for tensors in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(
                        group,
                        tensors['params'],
                        tensors['grads'],
                        tensors['exp_avg'],
                        tensors['exp_avg_sq'],
                        tensors['exp_avg_lr'],
                        step_size,
                    )
            else:
                self._step_per_param(group, step_size)

        return loss

AdamP

Bases: BaseOptimizer

Adam with projected updates for scale invariant weights.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
delta float

Threshold that determines whether a set of parameters is scale invariant or not.

0.1
wd_ratio float

Relative weight decay applied on scale invariant parameters compared to that applied on scale-variant parameters.

0.1
nesterov bool

Use Nesterov momentum.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adamp.py
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
class AdamP(BaseOptimizer):
    """Adam with projected updates for scale invariant weights.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        delta: Threshold that determines whether a set of parameters is scale invariant or not.
        wd_ratio: Relative weight decay applied on scale invariant parameters compared to that applied on
            scale-variant parameters.
        nesterov: Use Nesterov momentum.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        delta: float = 0.1,
        wd_ratio: float = 0.1,
        nesterov: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(wd_ratio, 'wd_ratio', 0.0, 1.0)
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'delta': delta,
            'wd_ratio': wd_ratio,
            'nesterov': nesterov,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdamP'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

                if group.get('adanorm'):
                    state['exp_grad_adanorm'] = torch.zeros((1,), dtype=grad.dtype, device=grad.device)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size: float = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=group['lr'],
                bias_correction1=bias_correction1,
            )

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

                if group.get('use_gc'):
                    centralize_gradient(grad, gc_conv_only=False)

                s_grad = self.get_adanorm_gradient(
                    grad=grad,
                    adanorm=group.get('adanorm', False),
                    exp_grad_norm=state.get('exp_grad_adanorm', None),
                    r=group.get('adanorm_r', None),
                )

                exp_avg.lerp_(s_grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                inv_de_nom = exp_avg_sq.sqrt().add_(group['eps']).reciprocal_().mul_(bias_correction2_sq)

                perturb = exp_avg.clone()

                if group.get('cautious'):
                    self.apply_cautious(perturb, grad)

                if group['nesterov']:
                    perturb.lerp_(grad, weight=1.0 - beta1).mul_(inv_de_nom)
                else:
                    perturb.mul_(inv_de_nom)

                wd_ratio: float = 1.0
                if len(p.shape) > 1:
                    perturb, wd_ratio = projection(
                        p,
                        grad,
                        perturb,
                        group['delta'],
                        group['wd_ratio'],
                        group['eps'],
                    )

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                    ratio=wd_ratio,
                )

                p.add_(perturb, alpha=-step_size)

        return loss

AdamS

Bases: BaseOptimizer

Adam with stable weight decay.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0001
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adams.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
class AdamS(BaseOptimizer):
    """Adam with stable weight decay.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 1e-4,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        ams_bound: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'ams_bound': ams_bound,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdamS'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

                if group['ams_bound']:
                    state['max_exp_avg_sq'] = torch.zeros_like(p)

                if group.get('adanorm'):
                    state['exp_grad_adanorm'] = torch.zeros((1,), dtype=p.dtype, device=p.device)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        param_size: int = 0
        exp_avg_sq_hat_sum: float = 0.0

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction2: float = self.debias(beta2, group['step'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                param_size += p.numel()

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                if not group['weight_decouple']:
                    self.apply_weight_decay(
                        p=p,
                        grad=grad,
                        lr=group['lr'],
                        weight_decay=group['weight_decay'],
                        weight_decouple=False,
                        fixed_decay=group['fixed_decay'],
                    )

                s_grad = self.get_adanorm_gradient(
                    grad=grad,
                    adanorm=group.get('adanorm', False),
                    exp_grad_norm=state.get('exp_grad_adanorm', None),
                    r=group.get('adanorm_r', None),
                )

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']
                exp_avg.lerp_(s_grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                if group['ams_bound']:
                    max_exp_avg_sq = state['max_exp_avg_sq']
                    torch.max(max_exp_avg_sq, exp_avg_sq, out=max_exp_avg_sq)
                    exp_avg_sq_hat = max_exp_avg_sq
                else:
                    exp_avg_sq_hat = exp_avg_sq

                exp_avg_sq_hat_sum += exp_avg_sq_hat.sum() / bias_correction2

        if param_size == 0:
            raise ZeroParameterSizeError

        exp_avg_sq_hat_mean: float = math.sqrt(exp_avg_sq_hat_sum / param_size) + self.defaults['eps']

        for group in self.param_groups:
            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2: float = self.debias(beta2, group['step'])

            step_size: float = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=group['lr'],
                bias_correction1=bias_correction1,
            )

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=None,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                    ratio=1.0 / exp_avg_sq_hat_mean,
                )

                exp_avg_sq_hat = state['max_exp_avg_sq'] if group['ams_bound'] else state['exp_avg_sq']
                de_nom = (exp_avg_sq_hat / bias_correction2).sqrt().add_(group['eps'])

                p.addcdiv_(state['exp_avg'], de_nom, value=-step_size)

        return loss

AdaMuon

Bases: MuonBase

Adaptive momentum updates with Newton-Schulz matrix orthogonalization.

Set use_muon=True for hidden weight matrices and use_muon=False for AdamW groups, such as embeddings, classifier heads, biases, and gains. Pass higher dimensional weights directly. The orthogonal update uses a flattened matrix view.

The default shape scaling gives the adaptive update an RMS of 0.2 before multiplication by the learning rate. use_adjusted_lr=True selects Moonlight scaling instead.

Parameters:

Name Type Description Default
params ParamsT

Parameter group dictionaries with a use_muon flag for each group.

required
lr float

Learning rate.

0.02
betas Betas

Decay rates for gradient momentum and squared orthogonalized updates.

(0.9, 0.95)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
ns_steps int

Number of Newton-Schulz iterations.

5
ns_coeffs NewtonSchulzWeights

Newton-Schulz coefficients or preset name.

'original'
use_adjusted_lr bool

Scale orthogonal updates using the Moonlight shape adjustment.

False
adamw_lr float

Learning rate for parameters in the AdamW groups.

0.0003
adamw_betas Betas

Decay rates for the first and second moments in the AdamW groups.

(0.9, 0.999)
adamw_wd float

Weight decay for parameters in the AdamW groups.

0.0
eps float

Term added to the denominator to improve numerical stability.

1e-10
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Batch tensor updates and compatible matrix shapes. False disables batching; None enables it.

False

Examples:

from pytorch_optimizer import AdaMuon

hidden_weights = [p for p in model.body.parameters() if p.ndim >= 2]
hidden_gains_biases = [p for p in model.body.parameters() if p.ndim < 2]
non_hidden_params = [*model.head.parameters(), *model.embed.parameters()]

param_groups = [
    dict(params=hidden_weights, lr=0.02, weight_decay=0.01, use_muon=True),
    dict(
        params=hidden_gains_biases + non_hidden_params,
        lr=3e-4,
        betas=(0.9, 0.95),
        weight_decay=0.01,
        use_muon=False,
    ),
]

optimizer = AdaMuon(param_groups)
Source code in pytorch_optimizer/optimizer/muon.py
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
class AdaMuon(MuonBase):
    """Adaptive momentum updates with Newton-Schulz matrix orthogonalization.

    Set `use_muon=True` for hidden weight matrices and `use_muon=False` for AdamW groups,
    such as embeddings, classifier heads, biases, and gains. Pass higher dimensional
    weights directly. The orthogonal update uses a flattened matrix view.

    The default shape scaling gives the adaptive update an RMS of 0.2 before multiplication
    by the learning rate. `use_adjusted_lr=True` selects Moonlight scaling instead.

    Args:
        params: Parameter group dictionaries with a `use_muon` flag for each group.
        lr: Learning rate.
        betas: Decay rates for gradient momentum and squared orthogonalized updates.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        ns_steps: Number of Newton-Schulz iterations.
        ns_coeffs: Newton-Schulz coefficients or preset name.
        use_adjusted_lr: Scale orthogonal updates using the Moonlight shape adjustment.
        adamw_lr: Learning rate for parameters in the AdamW groups.
        adamw_betas: Decay rates for the first and second moments in the AdamW groups.
        adamw_wd: Weight decay for parameters in the AdamW groups.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Batch tensor updates and compatible matrix shapes. `False` disables batching; `None` enables it.

    Examples:
        ```python
        from pytorch_optimizer import AdaMuon

        hidden_weights = [p for p in model.body.parameters() if p.ndim >= 2]
        hidden_gains_biases = [p for p in model.body.parameters() if p.ndim < 2]
        non_hidden_params = [*model.head.parameters(), *model.embed.parameters()]

        param_groups = [
            dict(params=hidden_weights, lr=0.02, weight_decay=0.01, use_muon=True),
            dict(
                params=hidden_gains_biases + non_hidden_params,
                lr=3e-4,
                betas=(0.9, 0.95),
                weight_decay=0.01,
                use_muon=False,
            ),
        ]

        optimizer = AdaMuon(param_groups)
        ```

    """

    _muon_state_keys = ('m', 'v')

    def __init__(
        self,
        params: ParamsT,
        lr: float = 2e-2,
        betas: Betas = (0.9, 0.95),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        ns_steps: int = 5,
        ns_coeffs: NewtonSchulzWeights = 'original',
        use_adjusted_lr: bool = False,
        adamw_lr: float = 3e-4,
        adamw_betas: Betas = (0.9, 0.999),
        adamw_wd: float = 0.0,
        eps: float = 1e-10,
        maximize: bool = False,
        foreach: bool | None = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_learning_rate(adamw_lr)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_positive(ns_steps, 'ns_steps')
        self.validate_betas(betas)
        self.validate_betas(adamw_betas)
        self.validate_non_negative(adamw_wd, 'adamw_wd')
        self.validate_non_negative(eps, 'eps')
        ns_coeffs = get_newton_schulz_weights(ns_coeffs)

        self.maximize = maximize
        self.foreach = foreach

        for group in params:
            group = cast(ParamGroup, group)
            if 'use_muon' not in group:
                raise ValueError('`use_muon` must be set.')

            if group['use_muon']:
                group['lr'] = group.get('lr', lr)
                group['betas'] = group.get('betas', betas)
                group['weight_decay'] = group.get('weight_decay', weight_decay)
                group['ns_steps'] = group.get('ns_steps', ns_steps)
                group['ns_coeffs'] = get_newton_schulz_weights(group.get('ns_coeffs', ns_coeffs))
                group['use_adjusted_lr'] = group.get('use_adjusted_lr', use_adjusted_lr)
            else:
                group['lr'] = group.get('lr', adamw_lr)
                group['betas'] = group.get('betas', adamw_betas)
                group['weight_decay'] = group.get('weight_decay', adamw_wd)

            group['weight_decouple'] = group.get('weight_decouple', weight_decouple)
            group['eps'] = group.get('eps', eps)

        super().__init__(params, {'foreach': foreach, **kwargs})

    def __str__(self) -> str:
        return 'AdaMuon'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                if group['use_muon']:
                    state['m'] = torch.zeros_like(p)
                    state['v'] = torch.zeros_like(p.flatten())
                else:
                    state['exp_avg'] = torch.zeros_like(p)
                    state['exp_avg_sq'] = torch.zeros_like(p)

    def _step_muon_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        state_dict: dict[str, list[torch.Tensor]],
        bias_correction2: float | torch.Tensor,
    ) -> None:
        beta1, beta2 = group['betas']
        moments, variances = state_dict['m'], state_dict['v']

        torch._foreach_lerp_(moments, grads, weight=1.0 - beta1)

        updates = [update.flatten() for update in self._orthogonalize(group, moments)]

        torch._foreach_mul_(variances, beta2)
        torch._foreach_addcmul_(variances, updates, updates, value=1.0 - beta2)

        de_noms = torch._foreach_sqrt(torch._foreach_div(variances, bias_correction2))
        torch._foreach_add_(de_noms, group['eps'])
        torch._foreach_div_(updates, de_noms)

        norms = [update.norm().add_(group['eps']) for update in updates]
        rows = params[0].size(0)
        torch._foreach_mul_(updates, math.sqrt(min(rows, params[0].numel() // rows)))
        torch._foreach_div_(updates, norms)

        updates = [update.reshape(p.shape) for p, update in zip(params, updates)]
        lr = get_adjusted_lr(group['lr'], params[0].shape, use_adjusted_lr=group['use_adjusted_lr'])
        foreach_add_(params, updates, alpha=-lr)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2: float = self.debias(beta2, group['step'])

            if self.can_use_foreach(group, group.get('foreach', self.foreach)):
                self._step_foreach_group(group)
                continue

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=False,
                )

                if group['use_muon']:
                    m = state['m']
                    m.lerp_(grad, weight=1.0 - beta1)

                    update = m

                    if update.ndim > 2:
                        update = update.view(len(update), -1)

                    update = zero_power_via_newton_schulz_5(
                        update, num_steps=group['ns_steps'], weights=group['ns_coeffs']
                    ).flatten()

                    v = state['v']
                    v.mul_(beta2).addcmul_(update, update, value=1.0 - beta2)

                    update.div_((v / bias_correction2).sqrt_().add_(group['eps']))
                    update = update.reshape(p.size())

                    scale = math.sqrt(min(p.size(0), p.numel() // p.size(0)))
                    update.mul_(scale / update.norm().add_(group['eps']))

                    lr = get_adjusted_lr(group['lr'], p.size(), use_adjusted_lr=group['use_adjusted_lr'])

                    p.add_(update, alpha=-lr)
                else:
                    exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                    exp_avg.lerp_(grad, weight=1.0 - beta1)
                    exp_avg_sq.lerp_(grad.square(), weight=1.0 - beta2)

                    de_nom = exp_avg_sq.sqrt().div_(math.sqrt(bias_correction2)).add_(group['eps'])

                    p.addcdiv_(exp_avg / bias_correction1, de_nom, value=-group['lr'])

        return loss

AdamWSN

Bases: BaseOptimizer

AdamW with subset norm and subspace momentum scaling.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
subset_size int

Number of weights per second moment subset. -1 uses half the first tensor dimension for matrices and the full size for vectors.

-1
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False

Examples:

from torch import nn
from pytorch_optimizer import AdamWSN

sn_params = [module.weight for module in model.modules() if isinstance(module, nn.Linear)]
sn_param_ids = {id(p) for p in sn_params}
regular_params = [p for p in model.parameters() if id(p) not in sn_param_ids]
param_groups = [{'params': regular_params, 'sn': False}, {'params': sn_params, 'sn': True}]
optimizer = AdamWSN(param_groups, lr=1e-3, weight_decay=1e-2, subset_size=-1)
Source code in pytorch_optimizer/optimizer/snsm.py
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
class AdamWSN(BaseOptimizer):
    """AdamW with subset norm and subspace momentum scaling.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        subset_size: Number of weights per second moment subset. `-1` uses half the first tensor dimension for
            matrices and the full size for vectors.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    Examples:
        ```python
        from torch import nn
        from pytorch_optimizer import AdamWSN

        sn_params = [module.weight for module in model.modules() if isinstance(module, nn.Linear)]
        sn_param_ids = {id(p) for p in sn_params}
        regular_params = [p for p in model.parameters() if id(p) not in sn_param_ids]
        param_groups = [{'params': regular_params, 'sn': False}, {'params': sn_params, 'sn': True}]
        optimizer = AdamWSN(param_groups, lr=1e-3, weight_decay=1e-2, subset_size=-1)
        ```

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        subset_size: int = -1,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'subset_size': subset_size,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdamWSN'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(grad)

                if group.get('sn'):
                    size: int = grad.numel()

                    if 'subset_size' not in state:
                        state['subset_size'] = closest_smaller_divisor_of_n_to_k(
                            size,
                            (
                                group['subset_size']
                                if group['subset_size'] > 0
                                else int(math.sqrt(size) / abs(int(group['subset_size'])))
                            ),
                        )

                    reshaped_grad = grad.view(size // state['subset_size'], state['subset_size'])
                    second_moment_update = torch.sum(reshaped_grad ** 2, dim=1, keepdim=True)  # fmt: skip
                    state['exp_avg_sq'] = torch.zeros_like(second_moment_update)
                else:
                    state['exp_avg_sq'] = torch.zeros_like(grad)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size: float = group['lr'] * bias_correction2_sq / bias_correction1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad
                size = grad.numel()

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                if group.get('sn'):
                    reshaped_grad = grad.view(size // state['subset_size'], state['subset_size'])
                    second_moment_update = torch.sum(reshaped_grad ** 2, dim=1, keepdim=True)  # fmt: skip
                else:
                    second_moment_update = grad.pow(2)

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.lerp_(second_moment_update, weight=1.0 - beta2)

                de_nom = exp_avg_sq.sqrt().add_(group['eps'])

                if group.get('sn'):
                    numerator = exp_avg.view(size // state['subset_size'], state['subset_size'])
                    norm_grad = (numerator / de_nom).reshape(p.shape)
                    p.add_(norm_grad, alpha=-step_size)
                else:
                    p.addcdiv_(exp_avg, de_nom, value=-step_size)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

        return loss

Adan

Bases: BaseOptimizer

Adaptive updates with gradient difference momentum.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for gradients, gradient differences, and squared corrected gradients.

(0.98, 0.92, 0.99)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
max_grad_norm float

Maximum gradient norm to clip.

0.0
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adan.py
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
class Adan(BaseOptimizer):
    """Adaptive updates with gradient difference momentum.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for gradients, gradient differences, and squared corrected gradients.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        max_grad_norm: Maximum gradient norm to clip.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.98, 0.92, 0.99),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        max_grad_norm: float = 0.0,
        foreach: bool | None = None,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(max_grad_norm, 'max_grad_norm')
        self.validate_non_negative(eps, 'eps')

        self.max_grad_norm = max_grad_norm
        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'max_grad_norm': max_grad_norm,
            'foreach': foreach,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Adan'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        clip_global_grad_norm: float = kwargs.get('clip_global_grad_norm', 0.0)

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['exp_avg_diff'] = torch.zeros_like(p)
                state['previous_grad'] = grad.clone().mul_(-clip_global_grad_norm)

                if group.get('adanorm'):
                    state['exp_grad_adanorm'] = torch.zeros((1,), dtype=grad.dtype, device=grad.device)

    @torch.no_grad()
    def get_global_gradient_norm(self) -> torch.Tensor | float:
        if self.defaults['max_grad_norm'] == 0.0:
            return 1.0

        global_grad_norm = get_global_gradient_norm(self.param_groups)
        global_grad_norm.sqrt_().add_(self.defaults['eps'])

        return torch.clamp(self.defaults['max_grad_norm'] / global_grad_norm, max=1.0)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        if group.get('foreach') is False:
            return False

        if group.get('use_gc') or group.get('adanorm'):
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        exp_avg_diffs: list[torch.Tensor],
        prev_grads: list[torch.Tensor],
        clip_global_grad_norm: torch.Tensor | float,
    ) -> None:
        beta1, beta2, beta3 = group['betas']
        lr = group['lr']

        bias_correction1: float = self.debias(beta1, group['step'])
        bias_correction2: float = self.debias(beta2, group['step'])
        bias_correction3_sq: float = math.sqrt(self.debias(beta3, group['step']))

        if self.maximize:
            torch._foreach_neg_(grads)

        if isinstance(clip_global_grad_norm, torch.Tensor):
            clip_global_grad_norm = clip_global_grad_norm.item()

        torch._foreach_mul_(grads, clip_global_grad_norm)

        grad_diffs = torch._foreach_add(prev_grads, grads)

        torch._foreach_mul_(exp_avgs, beta1)
        torch._foreach_add_(exp_avgs, grads, alpha=1.0 - beta1)

        torch._foreach_mul_(exp_avg_diffs, beta2)
        torch._foreach_add_(exp_avg_diffs, grad_diffs, alpha=1.0 - beta2)

        torch._foreach_mul_(grad_diffs, beta2)
        torch._foreach_add_(grad_diffs, grads)

        torch._foreach_mul_(exp_avg_sqs, beta3)
        torch._foreach_addcmul_(exp_avg_sqs, grad_diffs, grad_diffs, value=1.0 - beta3)

        de_noms = torch._foreach_sqrt(exp_avg_sqs)
        torch._foreach_div_(de_noms, bias_correction3_sq)
        torch._foreach_add_(de_noms, group['eps'])

        if group['weight_decouple']:
            torch._foreach_mul_(params, 1.0 - lr * group['weight_decay'])

        torch._foreach_addcdiv_(params, exp_avgs, de_noms, value=-lr / bias_correction1)
        torch._foreach_addcdiv_(params, exp_avg_diffs, de_noms, value=-lr * beta2 / bias_correction2)

        if not group['weight_decouple']:
            torch._foreach_div_(params, 1.0 + lr * group['weight_decay'])

        torch._foreach_copy_(prev_grads, torch._foreach_neg(grads))

    def _step_per_param(self, group: ParamGroup, clip_global_grad_norm: torch.Tensor | float) -> None:
        beta1, beta2, beta3 = group['betas']

        bias_correction1: float = self.debias(beta1, group['step'])
        bias_correction2: float = self.debias(beta2, group['step'])
        bias_correction3_sq: float = math.sqrt(self.debias(beta3, group['step']))

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            exp_avg, exp_avg_sq, exp_avg_diff = state['exp_avg'], state['exp_avg_sq'], state['exp_avg_diff']
            grad_diff = state['previous_grad']

            p, grad, exp_avg, exp_avg_sq, exp_avg_diff, grad_diff = self.view_as_real(
                p, grad, exp_avg, exp_avg_sq, exp_avg_diff, grad_diff
            )

            grad.mul_(clip_global_grad_norm)

            if group.get('use_gc'):
                centralize_gradient(grad, gc_conv_only=False)

            grad_diff.add_(grad)

            s_grad = self.get_adanorm_gradient(
                grad=grad,
                adanorm=group.get('adanorm', False),
                exp_grad_norm=state.get('exp_grad_adanorm', None),
                r=group.get('adanorm_r', None),
            )

            exp_avg.lerp_(s_grad, weight=1.0 - beta1)
            exp_avg_diff.lerp_(grad_diff, weight=1.0 - beta2)

            grad_diff.mul_(beta2).add_(grad)
            exp_avg_sq.mul_(beta3).addcmul_(grad_diff, grad_diff, value=1.0 - beta3)

            de_nom = exp_avg_sq.sqrt().div_(bias_correction3_sq).add_(group['eps'])

            if group['weight_decouple']:
                p.mul_(1.0 - group['lr'] * group['weight_decay'])

            p.addcdiv_(exp_avg, de_nom, value=-group['lr'] / bias_correction1)
            p.addcdiv_(exp_avg_diff, de_nom, value=-group['lr'] * beta2 / bias_correction2)

            if not group['weight_decouple']:
                p.div_(1.0 + group['lr'] * group['weight_decay'])

            grad.neg_()
            state['previous_grad'].copy_(
                torch.view_as_complex(grad) if torch.is_complex(state['previous_grad']) else grad
            )

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        clip_global_grad_norm = self.get_global_gradient_norm()

        for group in self.param_groups:
            self.init_group(group, clip_global_grad_norm=clip_global_grad_norm)
            group['step'] += 1

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_avg_sq', 'exp_avg_diff', 'previous_grad']
                )
                if params:
                    self._step_foreach(
                        group,
                        params,
                        grads,
                        state_dict['exp_avg'],
                        state_dict['exp_avg_sq'],
                        state_dict['exp_avg_diff'],
                        state_dict['previous_grad'],
                        clip_global_grad_norm,
                    )
            else:
                self._step_per_param(group, clip_global_grad_norm)

        return loss

AdaNorm

Bases: BaseOptimizer

Adam with adaptive gradient norm correction.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.99)
r float

EMA factor. Preferred values are between 0.9 and 0.99.

0.95
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adanorm.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
class AdaNorm(BaseOptimizer):
    """Adam with adaptive gradient norm correction.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        r: EMA factor. Preferred values are between 0.9 and 0.99.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.99),
        r: float = 0.95,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        ams_bound: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'r': r,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'ams_bound': ams_bound,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdaNorm'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_var'] = torch.zeros_like(p)
                state['exp_grad_norm'] = torch.zeros((1,), dtype=p.dtype, device=p.device)

                if group['ams_bound']:
                    state['max_exp_avg_var'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size: float = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=group['lr'],
                bias_correction1=bias_correction1,
            )

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_var = state['exp_avg'], state['exp_avg_var']

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                s_grad = self.get_adanorm_gradient(
                    grad=grad,
                    adanorm=True,
                    exp_grad_norm=state['exp_grad_norm'],
                    r=group['r'],
                )

                exp_avg.lerp_(s_grad, weight=1.0 - beta1)
                exp_avg_var.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                de_nom = self.apply_ams_bound(
                    ams_bound=group['ams_bound'],
                    exp_avg_sq=exp_avg_var,
                    max_exp_avg_sq=state.get('max_exp_avg_var', None),
                    eps=group['eps'],
                )
                de_nom.div_(bias_correction2_sq)

                p.addcdiv_(exp_avg, de_nom, value=-step_size)

        return loss

AdaPNM

Bases: BaseOptimizer

Adam with positive negative momentum.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Momentum decay, second moment decay, and positive negative momentum mixing coefficient.

(0.9, 0.999, 1.0)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

True
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adapnm.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
class AdaPNM(BaseOptimizer):
    """Adam with positive negative momentum.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Momentum decay, second moment decay, and positive negative momentum mixing coefficient.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999, 1.0),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        ams_bound: bool = True,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'ams_bound': ams_bound,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdaPNM'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['neg_exp_avg'] = torch.zeros_like(p)

                if group['ams_bound']:
                    state['max_exp_avg_sq'] = torch.zeros_like(p)

                if group.get('adanorm'):
                    state['exp_grad_adanorm'] = torch.zeros((1,), dtype=grad.dtype, device=grad.device)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2, beta3 = group['betas']

            beta1_p2: float = beta1 ** 2  # fmt: skip
            noise_norm: float = math.sqrt((1 + beta3) ** 2 + beta3 ** 2)  # fmt: skip

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size: float = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False), step_size=group['lr'], bias_correction1=bias_correction1
            )

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                exp_avg_sq = state['exp_avg_sq']

                if group['step'] % 2 == 1:
                    exp_avg, neg_exp_avg = state['exp_avg'], state['neg_exp_avg']
                else:
                    exp_avg, neg_exp_avg = state['neg_exp_avg'], state['exp_avg']

                p, grad, exp_avg, neg_exp_avg, exp_avg_sq = self.view_as_real(
                    p, grad, exp_avg, neg_exp_avg, exp_avg_sq
                )

                s_grad = self.get_adanorm_gradient(
                    grad=grad,
                    adanorm=group.get('adanorm', False),
                    exp_grad_norm=state.get('exp_grad_adanorm', None),
                    r=group.get('adanorm_r', None),
                )

                exp_avg.lerp_(s_grad, weight=1.0 - beta1_p2)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                de_nom = self.apply_ams_bound(
                    ams_bound=group['ams_bound'],
                    exp_avg_sq=exp_avg_sq,
                    max_exp_avg_sq=state.get('max_exp_avg_sq', None),
                    eps=group['eps'],
                )
                de_nom.div_(bias_correction2_sq)

                pn_momentum = exp_avg.mul(1.0 + beta3).add_(neg_exp_avg, alpha=-beta3).mul_(1.0 / noise_norm)

                p.addcdiv_(pn_momentum, de_nom, value=-step_size)

        return loss

AdaShift

Bases: BaseOptimizer

Adaptive updates with temporally decorrelated gradient moments.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
keep_num int

Number of gradients used to compute first moment estimation.

10
reduce_func Callable | None

Function applied to squared gradients to reduce correlation. If None, no function is applied.

max
eps float

Term added to the denominator to improve numerical stability.

1e-10
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adashift.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
class AdaShift(BaseOptimizer):
    """Adaptive updates with temporally decorrelated gradient moments.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        keep_num: Number of gradients used to compute first moment estimation.
        reduce_func: Function applied to squared gradients to reduce correlation. If None, no function is applied.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        keep_num: int = 10,
        reduce_func: Callable | None = torch.max,
        eps: float = 1e-10,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_positive(keep_num, 'keep_num')
        self.validate_non_negative(eps, 'eps')

        self.reduce_func: Callable = reduce_func if reduce_func is not None else lambda x: x
        self.maximize = maximize

        defaults: Defaults = {'lr': lr, 'betas': betas, 'keep_num': keep_num, 'eps': eps, **kwargs}

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdaShift'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['grad_queue'] = deque(maxlen=group['keep_num'])
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            exp_weight_sum: int = sum(beta1 ** i for i in range(group['keep_num']))  # fmt: skip
            first_grad_weight: float = beta1 ** (group['keep_num'] - 1) / exp_weight_sum
            last_grad_weight: float = 1.0 / exp_weight_sum

            bias_correction: float = self.debias(beta2, max(1, group['step'] - group['keep_num']))

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                grad_queue = state['grad_queue']
                offset_grad = grad_queue[0] if len(grad_queue) == group['keep_num'] else None
                grad_queue.append(grad.clone())

                exp_avg = state['exp_avg']
                if offset_grad is None:
                    exp_avg.mul_(beta1).add_(grad, alpha=last_grad_weight)
                    continue

                exp_avg.sub_(offset_grad, alpha=first_grad_weight).mul_(beta1).add_(grad, alpha=last_grad_weight)

                reduced_grad_sq = self.reduce_func(offset_grad.square())

                exp_avg_sq = state['exp_avg_sq']
                exp_avg_sq.lerp_(reduced_grad_sq, weight=1.0 - beta2)

                update = exp_avg.clone()
                if group.get('cautious'):
                    self.apply_cautious(update, grad)

                update.div_(exp_avg_sq.div(bias_correction).sqrt_().add_(group['eps']))

                p.add_(update, alpha=-group['lr'])

        return loss

AdaSmooth

Bases: BaseOptimizer

Adaptive updates with effective ratio smoothing.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Lower and upper smoothing bounds for the effective ratio adaptation.

(0.5, 0.99)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-06
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adasmooth.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
class AdaSmooth(BaseOptimizer):
    """Adaptive updates with effective ratio smoothing.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Lower and upper smoothing bounds for the effective ratio adaptation.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.5, 0.99),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        eps: float = 1e-6,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdaSmooth'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['prev_param'] = p.clone()
                state['s'] = torch.zeros_like(p)
                state['n'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                s, n, prev_param, exp_avg_sq = state['s'], state['n'], state['prev_param'], state['exp_avg_sq']

                p, grad, s, n, prev_param, exp_avg_sq = self.view_as_real(p, grad, s, n, prev_param, exp_avg_sq)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                p_diff = p - prev_param

                s.add_(p_diff)
                n.add_(p_diff.abs())

                c = s.abs().div_(n.add(group['eps']))
                c.mul_(beta2 - beta1).add_(1.0 - beta2)

                c_p2 = c.pow(2)

                exp_avg_sq.lerp_(grad.square(), weight=c_p2)

                step_size = torch.full_like(exp_avg_sq, fill_value=group['lr'])
                step_size.div_((exp_avg_sq + group['eps']).sqrt()).mul_(grad)

                state['prev_param'].copy_(torch.view_as_complex(p) if torch.is_complex(state['prev_param']) else p)

                p.add_(-step_size)

        return loss

AdaTAM

Bases: BaseOptimizer

Adam with gradient momentum alignment scaling.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for gradient momentum and squared gradients.

(0.9, 0.999)
decay_rate float

Decay rate for the gradient momentum alignment average.

0.9
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/tam.py
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
class AdaTAM(BaseOptimizer):
    """Adam with gradient momentum alignment scaling.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for gradient momentum and squared gradients.
        decay_rate: Decay rate for the gradient momentum alignment average.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        decay_rate: float = 0.9,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_range(decay_rate, 'decay_rate', 0.0, 1.0)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'decay_rate': decay_rate,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdaTAM'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['s'] = torch.zeros_like(grad)
                state['exp_avg'] = torch.zeros_like(grad)
                state['exp_avg_sq'] = torch.zeros_like(grad)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']
            decay_rate: float = group['decay_rate']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p,
                    grad,
                    group['lr'],
                    group['weight_decay'],
                    group['weight_decouple'],
                    group['fixed_decay'],
                )

                s, exp_avg, exp_avg_sq = state['s'], state['exp_avg'], state['exp_avg_sq']

                corr = normalize(exp_avg, p=2.0, dim=0).mul_(normalize(grad, p=2.0, dim=0))
                s.lerp_(corr, weight=1.0 - decay_rate)

                d = ((1.0 + s) / 2.0).add_(group['eps']).mul_(grad)

                exp_avg.mul_(beta1).add_(d)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                p.addcdiv_(exp_avg, exp_avg_sq.sqrt().add_(group['eps']), value=-group['lr'])

        return loss

AdEMAMix

Bases: BaseOptimizer

Adam with a mixture of fast and slow gradient momentum.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for fast gradient momentum, squared gradients, and slow gradient momentum.

(0.9, 0.999, 0.9999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
alpha float

Weight of slow momentum relative to fast momentum.

5.0
t_alpha_beta3 float | None

Number of steps to warm up alpha and the slow momentum decay. None disables warmup.

None
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
Source code in pytorch_optimizer/optimizer/ademamix.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
class AdEMAMix(BaseOptimizer):
    """Adam with a mixture of fast and slow gradient momentum.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for fast gradient momentum, squared gradients, and slow gradient momentum.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        alpha: Weight of slow momentum relative to fast momentum.
        t_alpha_beta3: Number of steps to warm up `alpha` and the slow momentum decay. `None` disables warmup.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999, 0.9999),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        alpha: float = 5.0,
        t_alpha_beta3: float | None = None,
        eps: float = 1e-8,
        maximize: bool = False,
        foreach: bool | None = None,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(alpha, 'alpha')
        self.validate_non_negative(t_alpha_beta3, 't_alpha_beta3')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'alpha': alpha,
            't_alpha_beta3': t_alpha_beta3,
            'eps': eps,
            'foreach': foreach,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AdEMAMix'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['exp_avg_slow'] = torch.zeros_like(p)

    @staticmethod
    def schedule_alpha(t_alpha_beta3: float | None, step: int, alpha: float) -> float:
        return alpha if t_alpha_beta3 is None else min(step * alpha / t_alpha_beta3, alpha)

    @staticmethod
    def schedule_beta3(t_alpha_beta3: float | None, step: int, beta1: float, beta3: float) -> float:
        if t_alpha_beta3 is None:
            return beta3

        log_beta1, log_beta3 = math.log(beta1), math.log(beta3)

        return min(
            math.exp(
                log_beta1 * log_beta3 / ((1.0 - step / t_alpha_beta3) * log_beta3 + (step / t_alpha_beta3) * log_beta1)
            ),
            beta3,
        )

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        exp_avg_slows: list[torch.Tensor],
        bias_correction1: float,
        bias_correction2_sq: float,
        alpha_t: float,
        beta3_t: float,
    ) -> None:
        beta1, beta2, _ = group['betas']

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params, grads, group['lr'], group['weight_decay'], group['weight_decouple'], group['fixed_decay']
        )

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - beta1)
        torch._foreach_mul_(exp_avg_sqs, beta2)
        torch._foreach_addcmul_(exp_avg_sqs, grads, grads, value=1.0 - beta2)
        torch._foreach_lerp_(exp_avg_slows, grads, weight=1.0 - beta3_t)

        de_noms = torch._foreach_sqrt(exp_avg_sqs)
        torch._foreach_div_(de_noms, bias_correction2_sq)
        torch._foreach_add_(de_noms, group['eps'])

        if group.get('cautious'):
            updates = [exp_avg.clone() for exp_avg in exp_avgs]
            for update, grad in zip(updates, grads):
                self.apply_cautious(update, grad)
            torch._foreach_div_(updates, bias_correction1)
        else:
            updates = torch._foreach_div(exp_avgs, bias_correction1)

        torch._foreach_add_(updates, exp_avg_slows, alpha=alpha_t)
        torch._foreach_div_(updates, de_noms)

        if group.get('stable_adamw'):
            rms = [self.get_stable_adamw_rms(grad, exp_avg_sq) for grad, exp_avg_sq in zip(grads, exp_avg_sqs)]
            torch._foreach_div_(updates, rms)

        torch._foreach_add_(params, updates, alpha=-group['lr'])

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2, beta3 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            alpha_t: float = self.schedule_alpha(group['t_alpha_beta3'], group['step'], group['alpha'])
            beta3_t: float = self.schedule_beta3(group['t_alpha_beta3'], group['step'], beta1, beta3)

            if self.can_use_foreach(group, group.get('foreach')):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_avg_sq', 'exp_avg_slow']
                )
                for batch in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(
                        group,
                        batch['params'],
                        batch['grads'],
                        batch['exp_avg'],
                        batch['exp_avg_sq'],
                        batch['exp_avg_slow'],
                        bias_correction1,
                        bias_correction2_sq,
                        alpha_t,
                        beta3_t,
                    )
                continue

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                exp_avg, exp_avg_sq, exp_avg_slow = state['exp_avg'], state['exp_avg_sq'], state['exp_avg_slow']

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)
                exp_avg_slow.lerp_(grad, weight=1.0 - beta3_t)

                de_nom = exp_avg_sq.sqrt().div_(bias_correction2_sq).add_(group['eps'])

                update = exp_avg.clone()
                if group.get('cautious'):
                    self.apply_cautious(update, grad)

                update.div_(bias_correction1).add_(exp_avg_slow, alpha=alpha_t).div_(de_nom)

                if group.get('stable_adamw'):
                    update.div_(self.get_stable_adamw_rms(grad, exp_avg_sq))

                p.add_(update, alpha=-group['lr'])

        return loss

ADOPT

Bases: BaseOptimizer

Adaptive updates using the previous second moment and gradient clipping.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.9999)
clip_lambda Callable[[float], float]

Function to clip gradient. Default is step ** 0.25.

lambda step: pow(step, 0.25)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
eps float

Term added to the denominator to improve numerical stability.

1e-06
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adopt.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
class ADOPT(BaseOptimizer):
    """Adaptive updates using the previous second moment and gradient clipping.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        clip_lambda: Function to clip gradient. Default is `step ** 0.25`.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.9999),
        clip_lambda: Callable[[float], float] = lambda step: math.pow(step, 0.25),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        foreach: bool | None = None,
        eps: float = 1e-6,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.clip_lambda = clip_lambda
        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'foreach': foreach,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'ADOPT'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        if group.get('foreach') is False:
            return False

        if group.get('cautious') or group.get('stable_adamw'):
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
    ) -> None:
        beta1, beta2 = group['betas']
        lr = group['lr']
        eps = group['eps']

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=lr,
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        if group['step'] == 1:
            torch._foreach_addcmul_(exp_avg_sqs, grads, grads)
            return

        de_noms = torch._foreach_sqrt(exp_avg_sqs)
        torch._foreach_clamp_min_(de_noms, eps)

        normed_grads = torch._foreach_div(grads, de_noms)
        if self.clip_lambda is not None:
            clip: float = self.clip_lambda(group['step'])
            torch._foreach_clamp_min_(normed_grads, -clip)
            torch._foreach_clamp_max_(normed_grads, clip)

        torch._foreach_lerp_(exp_avgs, normed_grads, weight=1.0 - beta1)

        torch._foreach_add_(params, exp_avgs, alpha=-lr)

        torch._foreach_mul_(exp_avg_sqs, beta2)
        torch._foreach_addcmul_(exp_avg_sqs, grads, grads, value=1.0 - beta2)

    def _step_per_param(self, group: ParamGroup) -> None:
        beta1, beta2 = group['betas']
        lr: float = group['lr']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

            p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

            self.apply_weight_decay(
                p=p,
                grad=grad,
                lr=lr,
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=group['fixed_decay'],
            )

            if group['step'] == 1:
                exp_avg_sq.addcmul_(grad, grad.conj())
                continue

            de_nom = exp_avg_sq.sqrt().clamp_(min=group['eps'])

            normed_grad = grad.div(de_nom)
            if self.clip_lambda is not None:
                clip = self.clip_lambda(group['step'])
                normed_grad.clamp_(-clip, clip)

            exp_avg.lerp_(normed_grad, weight=1.0 - beta1)

            if group.get('cautious'):
                update = exp_avg.clone()
                self.apply_cautious(update, normed_grad)
            else:
                update = exp_avg

            if group.get('stable_adamw'):
                update = update / self.get_stable_adamw_rms(grad, exp_avg_sq)

            p.add_(update, alpha=-lr)

            exp_avg_sq.mul_(beta2).addcmul_(grad, grad.conj(), value=1.0 - beta2)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_avg_sq']
                )
                if params:
                    self._step_foreach(group, params, grads, state_dict['exp_avg'], state_dict['exp_avg_sq'])
            else:
                self._step_per_param(group)

        return loss

agc

agc(p, grad, agc_eps=0.001, agc_clip_val=0.01, eps=1e-06)

Clip gradients relative to their parameter unit norms.

Parameters:

Name Type Description Default
p Tensor

Parameter tensor.

required
grad Tensor

Gradient tensor with the same shape as p.

required
agc_eps float

Lower bound for parameter unit norms.

0.001
agc_clip_val float

Maximum gradient-to-parameter norm ratio.

0.01
eps float

Lower bound for gradient unit norms.

1e-06

Returns:

Type Description
Tensor

torch.Tensor: Clipped gradient tensor. The input gradient remains unchanged.

Source code in pytorch_optimizer/optimizer/agc.py
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
def agc(
    p: torch.Tensor, grad: torch.Tensor, agc_eps: float = 1e-3, agc_clip_val: float = 1e-2, eps: float = 1e-6
) -> torch.Tensor:
    """Clip gradients relative to their parameter unit norms.

    Args:
        p: Parameter tensor.
        grad: Gradient tensor with the same shape as `p`.
        agc_eps: Lower bound for parameter unit norms.
        agc_clip_val: Maximum gradient-to-parameter norm ratio.
        eps: Lower bound for gradient unit norms.

    Returns:
        torch.Tensor: Clipped gradient tensor. The input gradient remains unchanged.

    """
    max_norm = unit_norm(p).clamp_min_(agc_eps).mul_(agc_clip_val)
    g_norm = unit_norm(grad).clamp_min_(eps)

    clipped_grad = grad * (max_norm / g_norm)

    return torch.where(g_norm > max_norm, clipped_grad, grad)

AggMo

Bases: BaseOptimizer

SGD with aggregated momentum buffers at multiple decay rates.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the momentum buffers to aggregate.

(0.0, 0.9, 0.99)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/aggmo.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
class AggMo(BaseOptimizer):
    """SGD with aggregated momentum buffers at multiple decay rates.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the momentum buffers to aggregate.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.0, 0.9, 0.99),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AggMo'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['momentum_buffer'] = {beta: torch.zeros_like(p) for beta in group['betas']}

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            betas = group['betas']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                for beta in betas:
                    buf = state['momentum_buffer'][beta]
                    buf.mul_(beta).add_(grad)

                    p.add_(buf, alpha=-group['lr'] / len(betas))

        return loss

Aida

Bases: BaseOptimizer

Adaptive updates using alternating gradient and momentum projections.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for gradient momentum and the squared projected gradient residual.

(0.9, 0.999)
k int

Number of alternating gradient and momentum projections per update.

2
xi float

Term used in vector projections to avoid division by zero.

1e-20
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
rectify bool

Perform the rectified update similar to RAdam.

False
n_sma_threshold int

Minimum effective simple moving average length for rectification.

5
degenerated_to_sgd bool

Use an SGD update before the moving average reaches the rectification threshold.

True
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/aida.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
class Aida(BaseOptimizer):
    """Adaptive updates using alternating gradient and momentum projections.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for gradient momentum and the squared projected gradient residual.
        k: Number of alternating gradient and momentum projections per update.
        xi: Term used in vector projections to avoid division by zero.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        rectify: Perform the rectified update similar to RAdam.
        n_sma_threshold: Minimum effective simple moving average length for rectification.
        degenerated_to_sgd: Use an SGD update before the moving average reaches the rectification threshold.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        k: int = 2,
        xi: float = 1e-20,
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        rectify: bool = False,
        n_sma_threshold: int = 5,
        degenerated_to_sgd: bool = True,
        ams_bound: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(k, 'k')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(xi, 'xi')
        self.validate_non_negative(eps, 'eps')

        self.k = k
        self.xi = xi
        self.n_sma_threshold = n_sma_threshold
        self.degenerated_to_sgd = degenerated_to_sgd
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'rectify': rectify,
            'ams_bound': ams_bound,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Aida'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_var'] = torch.zeros_like(p)

                if group['ams_bound']:
                    state['max_exp_avg_var'] = torch.zeros_like(p)

                if group.get('adanorm'):
                    state['exp_grad_adanorm'] = torch.zeros((1,), dtype=grad.dtype, device=grad.device)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size, n_sma = self.get_rectify_step_size(
                is_rectify=group['rectify'],
                step=group['step'],
                lr=group['lr'],
                beta2=beta2,
                n_sma_threshold=self.n_sma_threshold,
                degenerated_to_sgd=self.degenerated_to_sgd,
            )

            step_size = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=step_size,
                bias_correction1=bias_correction1,
            )

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                s_grad = self.get_adanorm_gradient(
                    grad=grad,
                    adanorm=group.get('adanorm', False),
                    exp_grad_norm=state.get('exp_grad_adanorm', None),
                    r=group.get('adanorm_r', None),
                )

                exp_avg, exp_avg_var = state['exp_avg'], state['exp_avg_var']
                exp_avg.lerp_(s_grad, weight=1.0 - beta1)

                proj_g = grad.detach().clone()
                proj_m = exp_avg.detach().clone()

                for _ in range(self.k):
                    proj_sum_gm = torch.sum(torch.mul(proj_g, proj_m))

                    scalar_g = proj_sum_gm / (torch.sum(torch.pow(proj_g, 2)).add_(self.xi))
                    scalar_m = proj_sum_gm / (torch.sum(torch.pow(proj_m, 2)).add_(self.xi))

                    proj_g.mul_(scalar_g)
                    proj_m.mul_(scalar_m)

                grad_residual = proj_m - proj_g
                exp_avg_var.mul_(beta2).addcmul_(grad_residual, grad_residual, value=1.0 - beta2).add_(group['eps'])

                de_nom = self.apply_ams_bound(
                    ams_bound=group['ams_bound'],
                    exp_avg_sq=exp_avg_var,
                    max_exp_avg_sq=state.get('max_exp_avg_var', None),
                    eps=group['eps'],
                )

                if not group['rectify']:
                    de_nom.div_(bias_correction2_sq)
                    p.addcdiv_(exp_avg, de_nom, value=-step_size)
                    continue

                if n_sma >= self.n_sma_threshold:
                    p.addcdiv_(exp_avg, de_nom, value=-step_size)
                elif step_size > 0:
                    p.add_(exp_avg, alpha=-step_size)

        return loss

Alice

Bases: BaseOptimizer

Adaptive subspace updates with full rank gradient compensation.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.02
betas Betas

Decay rates for gradient momentum, squared gradients, and subspace statistics. Set the third value to 0 for Alice-0.

(0.9, 0.9, 0.999)
alpha float

Update scaling factor.

0.3
alpha_c float

Scaling factor for the compensation update.

0.4
update_interval int

Number of steps between subspace updates.

200
rank int

Dimension of the low rank subspace.

256
gamma float

Maximum multiplicative growth of the scaled update norm.

1.01
leading_basis int

Number of leading subspace basis vectors to update.

40
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/racs.py
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
class Alice(BaseOptimizer):
    """Adaptive subspace updates with full rank gradient compensation.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for gradient momentum, squared gradients, and subspace statistics. Set the third value
            to 0 for Alice-0.
        alpha: Update scaling factor.
        alpha_c: Scaling factor for the compensation update.
        update_interval: Number of steps between subspace updates.
        rank: Dimension of the low rank subspace.
        gamma: Maximum multiplicative growth of the scaled update norm.
        leading_basis: Number of leading subspace basis vectors to update.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 0.02,
        betas: Betas = (0.9, 0.9, 0.999),
        alpha: float = 0.3,
        alpha_c: float = 0.4,
        update_interval: int = 200,
        rank: int = 256,
        gamma: float = 1.01,
        leading_basis: int = 40,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_range(alpha, 'alpha', 0.0, 1.0)
        self.validate_range(alpha_c, 'alpha_c', 0.0, 1.0)
        self.validate_positive(update_interval, 'update_interval')
        self.validate_positive(rank, 'rank')
        self.validate_positive(gamma, 'gamma')
        self.validate_positive(leading_basis, 'leading_basis')
        self.validate_non_negative(rank - leading_basis, 'rank - leading_basis')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'alpha': alpha,
            'alpha_c': alpha_c,
            'update_interval': update_interval,
            'rank': rank,
            'gamma': gamma,
            'leading_basis': leading_basis,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Alice'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

    @staticmethod
    def subspace_iteration(
        a: torch.Tensor, mat: torch.Tensor, num_steps: int = 1
    ) -> tuple[torch.Tensor, torch.Tensor]:
        r"""Perform subspace iteration."""
        u = mat
        for _ in range(num_steps):
            u, _ = torch.linalg.qr(a @ u)

        vals, vecs = torch.linalg.eigh(u.T @ a @ u)
        return vals, u @ vecs

    def switch(self, q: torch.Tensor, u_prev: torch.Tensor, rank: int, leading_basis: int) -> torch.Tensor:
        vals, vecs = self.subspace_iteration(q.to(torch.float32), u_prev.to(torch.float32), num_steps=1)

        leading_indices = torch.argsort(vals, descending=True)[:leading_basis]
        u_t1 = vecs[:, leading_indices]

        u_c, _ = torch.linalg.qr(torch.eye(q.shape[0], device=q.device) - u_t1 @ u_t1.T)
        u_t2 = u_c[:, :rank - leading_basis]  # fmt: skip

        return torch.cat([u_t1, u_t2], dim=1).to(q.dtype)

    @staticmethod
    def compensation(
        grad: torch.Tensor,
        u: torch.Tensor,
        p: torch.Tensor,
        phi: torch.Tensor,
        gamma: float,
        decay_rate: float,
        rank: int,
    ) -> tuple[torch.Tensor, torch.Tensor]:
        m = grad.size(0)

        sigma = u.T @ grad

        p.lerp_(grad.pow(2).sum(dim=0) - sigma.pow(2).sum(dim=0), weight=1.0 - decay_rate).clamp_min_(
            1e-8
        )

        c_t = math.sqrt(m - rank) * (grad - u @ sigma) / p.sqrt() if m >= rank else torch.zeros_like(grad)

        n = gamma / max(torch.norm(c_t) / phi, gamma) if phi.item() > 0 else torch.ones_like(phi)

        c_t.mul_(n)
        phi = torch.norm(c_t)

        return c_t, phi

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2, beta3 = group['betas']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad
                if grad.is_sparse:
                    raise NoSparseGradientError(str(self))

                if torch.is_complex(p):
                    raise NoComplexParameterError(str(self))

                state = self.state[p]

                self.maximize_gradient(grad, maximize=self.maximize)

                if grad.ndim < 2:
                    grad = grad.reshape(len(grad), 1)
                elif grad.ndim > 2:
                    grad = grad.reshape(len(grad), -1)

                has_state = 'U' in state
                if not has_state:
                    m, n = grad.shape
                    rank = min(group['rank'], m)

                    state['U'] = torch.zeros((m, rank), dtype=p.dtype, device=p.device)
                    state['Q'] = torch.zeros((rank, rank), dtype=p.dtype, device=p.device)

                    state['m'] = torch.zeros((rank, n), dtype=p.dtype, device=p.device)
                    state['v'] = torch.zeros((rank, n), dtype=p.dtype, device=p.device)

                    state['p'] = torch.zeros((n,), dtype=p.dtype, device=p.device)
                    state['phi'] = torch.zeros((1,), dtype=p.dtype, device=p.device)

                rank = state['U'].size(1)
                leading_basis = min(group['leading_basis'], rank)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                q, u, m, v = state['Q'], state['U'], state['m'], state['v']

                if not has_state or group['step'] % group['update_interval'] == 0:
                    q_t = beta3 * (u @ q @ u.T) + (1.0 - beta3) * (grad @ grad.T)
                    u = self.switch(q_t, u, rank, leading_basis)
                    state['U'] = u

                sigma = u.T @ grad

                q.lerp_(sigma @ sigma.T, weight=1.0 - beta3)
                m.lerp_(sigma, weight=1.0 - beta1)
                v.lerp_(sigma.pow(2), weight=1.0 - beta2)

                c_t, phi = self.compensation(grad, u, state['p'], state['phi'], group['gamma'], beta1, rank)

                update = u @ (m / v.sqrt().add_(group['eps']))
                update.add_(c_t, alpha=group['alpha_c'])

                p.add_(update.view_as(p), alpha=-group['lr'] * group['alpha'])

                state['phi'] = phi

        return loss

subspace_iteration(a, mat, num_steps=1) staticmethod

Perform subspace iteration.

Source code in pytorch_optimizer/optimizer/racs.py
220
221
222
223
224
225
226
227
228
229
230
@staticmethod
def subspace_iteration(
    a: torch.Tensor, mat: torch.Tensor, num_steps: int = 1
) -> tuple[torch.Tensor, torch.Tensor]:
    r"""Perform subspace iteration."""
    u = mat
    for _ in range(num_steps):
        u, _ = torch.linalg.qr(a @ u)

    vals, vecs = torch.linalg.eigh(u.T @ a @ u)
    return vals, u @ vecs

AliG

Bases: BaseOptimizer

Gradient steps scaled by loss with an optional learning rate cap.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
max_lr float | None

Maximum learning rate.

None
projection_fn Callable | None

Projection function to enforce constraints.

None
momentum float

Momentum factor.

0.0
adjusted_momentum bool

If True, use PyTorch like momentum instead of standard Nesterov momentum.

False
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/alig.py
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
class AliG(BaseOptimizer):
    """Gradient steps scaled by loss with an optional learning rate cap.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        max_lr: Maximum learning rate.
        projection_fn: Projection function to enforce constraints.
        momentum: Momentum factor.
        adjusted_momentum: If True, use PyTorch like momentum instead of standard Nesterov momentum.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        max_lr: float | None = None,
        projection_fn: Callable | None = None,
        momentum: float = 0.0,
        adjusted_momentum: bool = False,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(max_lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0)

        self.projection_fn = projection_fn
        self.maximize = maximize

        defaults: Defaults = {'max_lr': max_lr, 'adjusted_momentum': adjusted_momentum, 'momentum': momentum}

        super().__init__(params, defaults)

        if self.projection_fn is not None:
            self.projection_fn()

    def __str__(self) -> str:
        return 'AliG'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        momentum: float = kwargs.get('momentum', 0.9)

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0 and momentum > 0.0:
                state['momentum_buffer'] = torch.zeros_like(p)

    @torch.no_grad()
    def compute_step_size(self, loss: float) -> float:
        """Divide the loss by the sum of squared gradient norms plus a stability term."""
        global_grad_norm = get_global_gradient_norm(self.param_groups)
        global_grad_norm.add_(1e-6)

        return loss / global_grad_norm.item()

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        if closure is None:
            raise NoClosureError('AliG', '(e.g. `optimizer.step(lambda: float(loss))`).')

        loss = closure()

        un_clipped_step_size: float = self.compute_step_size(loss)

        for group in self.param_groups:
            momentum = group['momentum']

            self.init_group(group, momentum=momentum)
            group['step'] += 1

            step_size = group['step_size'] = (
                min(un_clipped_step_size, group['max_lr']) if group['max_lr'] is not None else un_clipped_step_size
            )

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                p, grad, buffer = self.view_as_real(p, grad, state.get('momentum_buffer', None))

                p.add_(grad, alpha=-step_size)

                if buffer is not None:
                    if group['adjusted_momentum']:
                        buffer.mul_(momentum).sub_(grad)
                        p.add_(buffer, alpha=step_size * momentum)
                    else:
                        buffer.mul_(momentum).add_(grad, alpha=-step_size)
                        p.add_(buffer, alpha=momentum)

            if self.projection_fn is not None:
                self.projection_fn()

        return loss

compute_step_size(loss)

Divide the loss by the sum of squared gradient norms plus a stability term.

Source code in pytorch_optimizer/optimizer/alig.py
80
81
82
83
84
85
86
@torch.no_grad()
def compute_step_size(self, loss: float) -> float:
    """Divide the loss by the sum of squared gradient norms plus a stability term."""
    global_grad_norm = get_global_gradient_norm(self.param_groups)
    global_grad_norm.add_(1e-6)

    return loss / global_grad_norm.item()

Amos

Bases: BaseOptimizer

Adaptive updates with weight decay toward a target parameter scale.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
beta float

A float slightly less than 1. Recommended to set 1 - beta approximately the same magnitude as the learning rate, similar to beta2 in Adam.

0.999
momentum float

Momentum factor.

0.0
extra_l2 float

Additional L2 regularization.

0.0
c_coef float

Coefficient for decay_factor_c.

0.25
d_coef float

Coefficient for decay_factor_d.

0.25
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
eps float

Term added to the denominator to improve numerical stability.

1e-18
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/amos.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
class Amos(BaseOptimizer):
    """Adaptive updates with weight decay toward a target parameter scale.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        beta: A float slightly less than 1. Recommended to set `1 - beta` approximately the same magnitude as the
            learning rate, similar to beta2 in Adam.
        momentum: Momentum factor.
        extra_l2: Additional L2 regularization.
        c_coef: Coefficient for decay_factor_c.
        d_coef: Coefficient for decay_factor_d.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        beta: float = 0.999,
        momentum: float = 0.0,
        extra_l2: float = 0.0,
        c_coef: float = 0.25,
        d_coef: float = 0.25,
        foreach: bool | None = None,
        eps: float = 1e-18,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0, range_type='[)')
        self.validate_range(beta, 'beta', 0.0, 1.0, range_type='[)')
        self.validate_non_negative(extra_l2, 'extra_l2')
        self.validate_non_negative(eps, 'eps')

        self.c_coef = c_coef
        self.d_coef = d_coef
        self.foreach = foreach
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'beta': beta,
            'momentum': momentum,
            'extra_l2': extra_l2,
            'foreach': foreach,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Amos'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg_sq'] = torch.zeros((1,), dtype=p.dtype, device=p.device)
                state['decay'] = torch.zeros((1,), dtype=p.dtype, device=p.device)
                if group['momentum'] > 0.0:
                    state['exp_avg'] = torch.zeros_like(p)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        if group.get('foreach') is False:
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    @staticmethod
    def get_scale(p: torch.Tensor) -> float:
        """Return the target weight scale from the parameter shape."""
        if len(p.shape) == 1:
            return 0.5
        if len(p.shape) == 2:
            return math.sqrt(2 / p.size(1))
        return math.sqrt(1 / p.size(1))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        decays: list[torch.Tensor],
    ) -> None:
        lr_sq: float = math.sqrt(group['lr'])
        lr_p2: float = math.pow(group['lr'], 2)

        beta: float = group['beta']
        bias_correction: float = self.debias(beta, group['step'])

        if self.maximize:
            torch._foreach_neg_(grads)

        g2 = [grad.pow(2).mean() for grad in grads]
        init_lrs: list[float] = [group['lr'] * self.get_scale(p) for p in params]

        torch._foreach_mul_(exp_avg_sqs, beta)
        torch._foreach_add_(exp_avg_sqs, g2, alpha=1.0 - beta)

        r_v_hat = torch._foreach_add(exp_avg_sqs, group['eps'])
        torch._foreach_reciprocal_(r_v_hat)
        torch._foreach_mul_(r_v_hat, bias_correction)

        df_c = torch._foreach_mul(decays, self.c_coef * lr_sq)
        torch._foreach_add_(df_c, 1.0)
        foreach_rsqrt_(df_c)

        d_step_sizes = [self.d_coef * math.sqrt(step_size) for step_size in init_lrs]
        df_d = torch._foreach_mul(decays, d_step_sizes)
        torch._foreach_add_(df_d, 1.0)

        torch._foreach_mul_(df_c, r_v_hat)
        torch._foreach_mul_(df_c, lr_p2)
        torch._foreach_mul_(df_c, g2)

        updates = torch._foreach_div(params, 2.0)
        torch._foreach_mul_(updates, torch._foreach_sub(df_c, group['extra_l2']))

        torch._foreach_sqrt_(r_v_hat)
        torch._foreach_mul_(r_v_hat, init_lrs)
        torch._foreach_add_(updates, torch._foreach_mul(grads, r_v_hat))

        torch._foreach_div_(updates, df_d)

        torch._foreach_mul_(decays, torch._foreach_add(df_c, 1.0))
        torch._foreach_add_(decays, df_c)

        if group['momentum'] > 0.0:
            torch._foreach_lerp_(exp_avgs, updates, weight=1.0 - group['momentum'])
            torch._foreach_copy_(updates, exp_avgs)

        torch._foreach_sub_(params, updates)

    def _step_per_param(self, group: ParamGroup) -> None:
        momentum, beta = group['momentum'], group['beta']

        lr_sq: float = math.sqrt(group['lr'])
        lr_p2: float = math.pow(group['lr'], 2)
        bias_correction: float = self.debias(beta, group['step'])

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            g2 = grad.pow(2).mean()
            init_lr: float = group['lr'] * self.get_scale(p)

            exp_avg_sq = state['exp_avg_sq']
            exp_avg_sq.lerp_(g2, weight=1.0 - beta)

            r_v_hat = bias_correction / (exp_avg_sq + group['eps'])

            decay = state['decay']
            decay_factor_c = torch.rsqrt(1.0 + self.c_coef * lr_sq * decay)
            decay_factor_d = torch.reciprocal(1.0 + self.d_coef * math.sqrt(init_lr) * decay)

            gamma = decay_factor_c * lr_p2 * r_v_hat * g2

            update = p.clone()
            update.mul_((gamma - group['extra_l2']) / 2.0)
            update.add_(r_v_hat.sqrt() * grad, alpha=init_lr)
            update.mul_(decay_factor_d)

            decay.mul_(1.0 + gamma).add_(gamma)

            if momentum > 0.0:
                exp_avg = state['exp_avg']
                exp_avg.lerp_(update, weight=1.0 - momentum)

                update.copy_(exp_avg)

            p.add_(-update)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_avg_sq', 'decay']
                )
                if params:
                    self._step_foreach(
                        group,
                        params,
                        grads,
                        state_dict['exp_avg'],
                        state_dict['exp_avg_sq'],
                        state_dict['decay'],
                    )
            else:
                self._step_per_param(group)

        return loss

get_scale(p) staticmethod

Return the target weight scale from the parameter shape.

Source code in pytorch_optimizer/optimizer/amos.py
 94
 95
 96
 97
 98
 99
100
101
@staticmethod
def get_scale(p: torch.Tensor) -> float:
    """Return the target weight scale from the parameter shape."""
    if len(p.shape) == 1:
        return 0.5
    if len(p.shape) == 2:
        return math.sqrt(2 / p.size(1))
    return math.sqrt(1 / p.size(1))

Ano

Bases: BaseOptimizer

Ano optimizer with adaptive momentum and sign based updates.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.0001
betas Betas

Coefficients used for computing running averages of gradient and the squared gradient.

(0.92, 0.99)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
logarithmic_schedule bool

Enable adaptive beta1 scheduling based on step count.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/ano.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
class Ano(BaseOptimizer):
    """Ano optimizer with adaptive momentum and sign based updates.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Coefficients used for computing running averages of gradient and the squared gradient.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        logarithmic_schedule: Enable adaptive beta1 scheduling based on step count.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-4,
        betas: Betas = (0.92, 0.99),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        logarithmic_schedule: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.logarithmic_schedule = logarithmic_schedule
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Ano'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            if self.logarithmic_schedule:
                max_t = max(2, group['step'])
                beta1 = 1.0 - 1.0 / math.log(max_t)

            bias_correction2: float = self.debias(beta2, group['step'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                exp_avg.lerp_(grad, weight=1.0 - beta1)

                square_grad = grad.square()
                sign_term = torch.sign(square_grad - exp_avg_sq)
                exp_avg_sq.mul_(beta2).addcmul_(sign_term, square_grad, value=1.0 - beta2)

                de_nom = square_grad.copy_(exp_avg_sq).div_(bias_correction2).sqrt_().add_(group['eps'])

                p.addcdiv_(grad.abs().mul_(exp_avg.sign()), de_nom, value=-group['lr'])

        return loss

APOLLO

Bases: BaseOptimizer

AdamW with low rank gradient projection and norm based update scaling.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.01
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
scale_type SCALE_TYPE

Compute update scaling per 'tensor' or per 'channel'.

'tensor'
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
correct_bias bool

Whether to correct bias in Adam.

True
eps float

Term added to the denominator to improve numerical stability.

1e-06
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/apollo.py
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
class APOLLO(BaseOptimizer):
    """AdamW with low rank gradient projection and norm based update scaling.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        scale_type: Compute update scaling per `'tensor'` or per `'channel'`.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        correct_bias: Whether to correct bias in Adam.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-2,
        betas: Betas = (0.9, 0.999),
        scale_type: SCALE_TYPE = 'tensor',
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        correct_bias: bool = True,
        eps: float = 1e-6,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'scale_type': scale_type,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'correct_bias': correct_bias,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'APOLLO'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            step_size: float = group['lr']
            if group['correct_bias']:
                bias_correction1: float = self.debias(beta1, group['step'])
                bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))
                step_size *= bias_correction2_sq / bias_correction1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

                if 'rank' in group and p.dim() > 1:
                    if 'projector' not in state:
                        state['projector'] = GaLoreProjector(
                            rank=group['rank'],
                            update_proj_gap=group['update_proj_gap'],
                            scale=group['scale'],
                            projection_type=group['projection_type'],
                        )

                    grad = state['projector'].project(grad, group['step'], from_random_matrix=True)

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                de_nom = exp_avg_sq.sqrt().add_(group['eps'])

                norm_grad = exp_avg / de_nom
                if 'rank' in group and p.dim() > 1:
                    if group['scale_type'] == 'channel':
                        norm_dim: int = 0 if norm_grad.shape[0] < norm_grad.shape[1] else 1
                        scaling_factor = torch.norm(norm_grad, dim=norm_dim) / (torch.norm(grad, dim=norm_dim) + 1e-8)
                        if norm_dim == 1:
                            scaling_factor = scaling_factor.unsqueeze(1)
                    else:
                        scaling_factor = torch.norm(norm_grad) / (torch.norm(grad) + 1e-8)

                    scaling_grad = grad * scaling_factor

                    scaling_grad_norm = torch.norm(scaling_grad)
                    if 'scaling_grad' in state:
                        limiter = (
                            max(
                                scaling_grad_norm / (state['scaling_grad'] + 1e-8),
                                1.01,
                            )
                            / 1.01
                        )

                        scaling_grad.div_(limiter)
                        scaling_grad_norm.div_(limiter)

                    state['scaling_grad'] = scaling_grad_norm

                    norm_grad = scaling_grad * np.sqrt(group['scale'])
                    norm_grad = state['projector'].project_back(norm_grad)

                p.add_(norm_grad, alpha=-step_size)

                self.apply_weight_decay(
                    p,
                    grad,
                    lr=step_size,
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

        return loss

ApolloDQN

Bases: BaseOptimizer

Adaptive updates with a diagonal quasi-Newton preconditioner.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.01
init_lr float | None

Initial learning rate (default lr / 1000).

1e-05
beta float

Coefficient used for computing running averages of gradient.

0.9
rebound str

Rectified bound for diagonal Hessian. Options: 'constant', 'belief'.

'constant'
weight_decay float

Weight decay coefficient.

0.0
weight_decay_type str

Type of weight decay. Options: 'l2', 'decoupled', 'stable'.

'l2'
warmup_steps int

Number of warmup steps.

500
eps float

Term added to the denominator to improve numerical stability.

0.0001
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/apollo.py
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
class ApolloDQN(BaseOptimizer):
    """Adaptive updates with a diagonal quasi-Newton preconditioner.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        init_lr: Initial learning rate (default lr / 1000).
        beta: Coefficient used for computing running averages of gradient.
        rebound: Rectified bound for diagonal Hessian. Options: 'constant', 'belief'.
        weight_decay: Weight decay coefficient.
        weight_decay_type: Type of weight decay. Options: 'l2', 'decoupled', 'stable'.
        warmup_steps: Number of warmup steps.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-2,
        init_lr: float | None = 1e-5,
        beta: float = 0.9,
        rebound: str = 'constant',
        weight_decay: float = 0.0,
        weight_decay_type: str = 'l2',
        warmup_steps: int = 500,
        eps: float = 1e-4,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(beta, 'beta', 0.0, 1.0, range_type='[]')
        self.validate_options(rebound, 'rebound', ['constant', 'belief'])
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_options(weight_decay_type, 'weight_decay_type', ['l2', 'decoupled', 'stable'])
        self.validate_non_negative(eps, 'eps')

        self.lr = lr
        self.warmup_steps = warmup_steps
        self.init_lr: float = init_lr if init_lr is not None else lr / 1000.0
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'init_lr': self.init_lr,
            'beta': beta,
            'rebound': rebound,
            'weight_decay': weight_decay,
            'weight_decay_type': weight_decay_type,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'ApolloDQN'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg_grad'] = torch.zeros_like(p)
                state['approx_hessian'] = torch.zeros_like(p)
                state['update'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            current_lr: float = (
                group['lr']
                if group['step'] >= self.warmup_steps
                else (self.lr - group['init_lr']) * group['step'] / self.warmup_steps + group['init_lr']
            )

            weight_decay, eps = group['weight_decay'], group['eps']

            bias_correction: float = self.debias(group['beta'], group['step'])
            alpha: float = (1.0 - group['beta']) / bias_correction

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg_grad, b, d_p = state['exp_avg_grad'], state['approx_hessian'], state['update']

                p, grad, exp_avg_grad, b, d_p = self.view_as_real(p, grad, exp_avg_grad, b, d_p)

                if weight_decay > 0.0 and group['weight_decay_type'] == 'l2':
                    grad.add_(p, alpha=weight_decay)

                delta_grad = grad - exp_avg_grad
                if group['rebound'] == 'belief':
                    rebound = delta_grad.norm(p=np.inf)
                else:
                    rebound = 1e-2
                    eps /= rebound

                exp_avg_grad.add_(delta_grad, alpha=alpha)

                de_nom = d_p.norm(p=4).add_(eps)
                d_p.div_(de_nom)

                v_sq = d_p.mul(d_p)
                delta = delta_grad.div_(de_nom).mul_(d_p).sum().mul(-alpha) - b.mul(v_sq).sum()

                b.addcmul_(v_sq, delta)

                de_nom = b.abs().clamp_(min=rebound)
                if group['rebound'] == 'belief':
                    de_nom.add_(eps / alpha)

                d_p.copy_(exp_avg_grad.div(de_nom))

                if weight_decay > 0.0 and group['weight_decay_type'] != 'l2':
                    decay = weight_decay
                    if group['weight_decay_type'] == 'stable':
                        decay = weight_decay / de_nom.mean().item()

                    d_p.add_(p, alpha=decay)

                p.add_(d_p, alpha=-current_lr)

        return loss

ASGD

Bases: BaseOptimizer

Adaptive SGD with estimation of the local smoothness (curvature).

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.01
amplifier float

Coefficient controlling the maximum learning rate growth per step.

0.02
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
theta float

Initial ratio of consecutive learning rates, updated after each step.

1.0
dampening float

Scale of the local smoothness bound on the learning rate.

1.0
eps float

Term added to denominator to improve numerical stability.

1e-05
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/sgd.py
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
class ASGD(BaseOptimizer):
    """Adaptive SGD with estimation of the local smoothness (curvature).

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        amplifier: Coefficient controlling the maximum learning rate growth per step.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        theta: Initial ratio of consecutive learning rates, updated after each step.
        dampening: Scale of the local smoothness bound on the learning rate.
        eps: Term added to denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-2,
        amplifier: float = 0.02,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        theta: float = 1.0,
        dampening: float = 1.0,
        eps: float = 1e-5,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_non_negative(amplifier, 'amplifier')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'amplifier': amplifier,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'theta': theta,
            'dampening': dampening,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'ASGD'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        pass

    @staticmethod
    def get_norms_by_group(group: ParamGroup, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]:
        """Compute global parameter and gradient L2 norms for a parameter group."""
        p_norm = torch.zeros(1, dtype=torch.float32, device=device)

        for p in group['params']:
            if p.grad is None:
                continue

            p_norm.add_(p.norm().pow(2))

        p_norm.sqrt_()
        g_norm = get_global_gradient_norm([group], device).sqrt_()

        return p_norm, g_norm

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            device = group['params'][0].device

            if 'prev_param_norm' not in group and 'prev_grad_norm' not in group:
                group['prev_param_norm'], group['prev_grad_norm'] = self.get_norms_by_group(group, device)

            group['curr_param_norm'], group['curr_grad_norm'] = self.get_norms_by_group(group, device)

            param_diff_norm: float = (group['curr_param_norm'] - group['prev_param_norm']).item()
            grad_diff_norm: float = (group['curr_grad_norm'] - group['prev_grad_norm']).item()

            new_lr: float = group['lr'] * math.sqrt(1 + group['amplifier'] * group['theta'])
            if param_diff_norm > 0 and grad_diff_norm > 0:
                new_lr = min(new_lr, param_diff_norm / (group['dampening'] * grad_diff_norm)) + group['eps']

            group['theta'] = new_lr / group['lr']
            group['lr'] = new_lr

            group['prev_param_norm'].copy_(group['curr_param_norm'])
            group['prev_grad_norm'].copy_(group['curr_grad_norm'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad
                if grad.is_sparse:
                    raise NoSparseGradientError(str(self))

                self.maximize_gradient(grad, maximize=self.maximize)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                p.add_(grad, alpha=-new_lr)

        return loss

get_norms_by_group(group, device) staticmethod

Compute global parameter and gradient L2 norms for a parameter group.

Source code in pytorch_optimizer/optimizer/sgd.py
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
@staticmethod
def get_norms_by_group(group: ParamGroup, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]:
    """Compute global parameter and gradient L2 norms for a parameter group."""
    p_norm = torch.zeros(1, dtype=torch.float32, device=device)

    for p in group['params']:
        if p.grad is None:
            continue

        p_norm.add_(p.norm().pow(2))

    p_norm.sqrt_()
    g_norm = get_global_gradient_norm([group], device).sqrt_()

    return p_norm, g_norm

AvaGrad

Bases: BaseOptimizer

Adaptive updates with a lagged second moment preconditioner.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.1
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

0.1
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/avagrad.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
class AvaGrad(BaseOptimizer):
    """Adaptive updates with a lagged second moment preconditioner.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-1,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-1,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'gamma': None,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'AvaGrad'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))
            prev_bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step'] - 1))

            step_size: float = group['lr']
            if group['step'] > 1:
                step_size: float = self.apply_adam_debias(
                    adam_debias=group.get('adam_debias', False),
                    step_size=group['gamma'] * group['lr'],
                    bias_correction1=bias_correction1,
                )

            squared_norm: float = 0.0
            num_params: float = 0.0

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                sqrt_exp_avg_sq = exp_avg_sq.sqrt()

                if group['step'] > 1:
                    de_nom = sqrt_exp_avg_sq.div(prev_bias_correction2_sq).add_(group['eps'])

                    p.addcdiv_(exp_avg, de_nom, value=-step_size)

                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                param_wise_lr = sqrt_exp_avg_sq.div_(bias_correction2_sq).add_(group['eps'])
                squared_norm += param_wise_lr.norm(-2) ** -2
                num_params += param_wise_lr.numel()

            group['gamma'] = 0.0 if num_params == 0.0 else 1.0 / math.sqrt(squared_norm / num_params)

        return loss

BCOS

Bases: BaseOptimizer

Stochastic approximation with block coordinate optimal step sizes.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
beta float

Decay rate for momentum and its second moment estimator.

0.9
beta2 float | None

Separate second moment decay rate. None uses beta.

None
mode Mode

Search direction and estimator: 'g' for gradients, 'm' for momentum, or 'c' for conditional momentum.

'c'
simple_cond bool

Use the simplified conditional estimator in mode c.

False
weight_decay float

Weight decay coefficient.

0.1
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
eps float

Term added to the denominator to improve numerical stability.

1e-06
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/bcos.py
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
class BCOS(BaseOptimizer):
    """Stochastic approximation with block coordinate optimal step sizes.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        beta: Decay rate for momentum and its second moment estimator.
        beta2: Separate second moment decay rate. `None` uses `beta`.
        mode: Search direction and estimator: `'g'` for gradients, `'m'` for momentum, or `'c'` for conditional
            momentum.
        simple_cond: Use the simplified conditional estimator in mode `c`.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        beta: float = 0.9,
        beta2: float | None = None,
        mode: Mode = 'c',
        simple_cond: bool = False,
        weight_decay: float = 0.1,
        weight_decouple: bool = True,
        eps: float = 1e-6,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(beta, 'beta', 0.0, 1.0)
        self.validate_options(mode, 'mode', ['g', 'm', 'c'])
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.mode = mode
        self.simple_cond = simple_cond
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'beta': beta,
            'beta2': beta2,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'eps': eps,
            **kwargs,
        }
        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'BCOS'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if self.mode in ('m', 'c') and 'm' not in state:
                state['m'] = grad.clone()
                self.maximize_gradient(state['m'], maximize=self.maximize)

            if self.mode in ('g', 'm') and 'v' not in state:
                state['v'] = grad.square()

    def compute_v(self, grad: torch.Tensor, m: torch.Tensor, beta: float, beta2: float | None) -> torch.Tensor:
        g2 = grad.square()

        if self.simple_cond:
            beta_v: float = 1.0 - (1.0 - beta) ** 2 if beta2 is None else beta2
            return beta_v * m.square() + (1.0 - beta_v) * g2

        return (
            (3.0 * beta ** 2 - 2.0 * beta ** 3) * m.square()
            + (1.0 - beta) ** 2 * g2
            + 2.0 * beta * (1.0 - beta) ** 2 * m * grad
        )  # fmt: skip

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta, beta2 = group['beta'], group['beta2']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=False,
                )

                old_m: torch.Tensor | None = state.get('m', None)

                if self.mode in ('m', 'c'):
                    m = state['m']
                    m.lerp_(grad, weight=1.0 - beta)
                    d = m
                else:
                    d = grad

                if self.mode in ('g', 'm'):
                    beta_v: float = beta if beta2 is None else beta2

                    v = state['v']
                    v.lerp_(d.square(), weight=1.0 - beta_v)
                else:
                    v: torch.Tensor = self.compute_v(grad, old_m, beta, beta2)

                p.addcdiv_(d, v.sqrt().add_(group['eps']), value=-group['lr'])

        return loss

BSAM

Bases: BaseOptimizer

Bayesian sharpness-aware minimization with noisy parameter perturbations.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
num_data int

Number of training data.

required
lr float

Learning rate.

0.5
betas Betas

Decay rates for gradient momentum and the squared curvature estimate.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0001
rho float

Size of the neighborhood for computing the max loss.

0.05
adaptive bool

Elementwise Adaptive SAM.

False
damping float

Damping to stabilize the method.

0.1
**kwargs dict

Parameters for optimizer.

{}
Source code in pytorch_optimizer/optimizer/sam.py
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
class BSAM(BaseOptimizer):
    """Bayesian sharpness-aware minimization with noisy parameter perturbations.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        num_data: Number of training data.
        lr: Learning rate.
        betas: Decay rates for gradient momentum and the squared curvature estimate.
        weight_decay: Weight decay coefficient.
        rho: Size of the neighborhood for computing the max loss.
        adaptive: Elementwise Adaptive SAM.
        damping: Damping to stabilize the method.
        **kwargs (dict): Parameters for optimizer.

    """

    def __init__(
        self,
        params: ParamsT,
        num_data: int,
        lr: float = 5e-1,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 1e-4,
        rho: float = 0.05,
        adaptive: bool = False,
        damping: float = 0.1,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(rho, 'rho')
        self.validate_non_negative(num_data, 'num_data')
        self.validate_non_negative(damping, 'damping')

        self.num_data = num_data
        self.damping = damping

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'rho': rho,
            'adaptive': adaptive,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'bSAM'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            state = self.state[p]

            if 's' not in state:
                state['s'] = torch.ones_like(p)
                state['noisy_gradient'] = torch.zeros_like(p.grad)
                state['momentum'] = torch.zeros_like(p)

    @torch.no_grad()
    def first_step(self):
        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            for p in group['params']:
                if p.grad is None:
                    continue

                state = self.state[p]

                noise = torch.normal(0.0, 1 / (self.num_data * state['s']))

                p.add_(noise)

    @torch.no_grad()
    def second_step(self):
        for group in self.param_groups:
            for p in group['params']:
                if p.grad is None:
                    continue

                state = self.state[p]

                state['noisy_gradient'] = p.grad.clone()

                e_w = (torch.pow(p, 2) if group['adaptive'] else 1.0) * group['rho'] * p.grad / state['s']

                p.add_(e_w)

    @torch.no_grad()
    def third_step(self):
        for group in self.param_groups:
            beta1, beta2 = group['betas']
            weight_decay = group['weight_decay']
            for p in group['params']:
                if p.grad is None:
                    continue

                state = self.state[p]

                momentum, s = state['momentum'], state['s']
                momentum.lerp_(p.grad * weight_decay, weight=1.0 - beta1)

                var = (torch.sqrt(s).mul_(p.grad.abs()).add_(weight_decay + self.damping)).pow_(2)
                s.lerp_(var, weight=1.0 - beta2)

                p.add_(momentum / s, alpha=-group['lr'])

    @torch.no_grad()
    def step(self, closure: Closure = None):
        if closure is None:
            raise NoClosureError(str(self))

        self.first_step()

        with torch.enable_grad():
            closure()

        self.second_step()

        with torch.enable_grad():
            loss = closure()

        self.third_step()

        return loss

CAME

Bases: BaseOptimizer

Factored adaptive updates with confidence weighted momentum.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.0002
betas Betas

Decay rates for gradient momentum, squared gradients, and squared update residuals.

(0.9, 0.999, 0.9999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
clip_threshold float

Maximum root mean square of the preconditioned update.

1.0
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
eps1 float

Stability constant added to squared gradients.

1e-30
eps2 float

Stability constant added to squared update residuals.

1e-16
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/came.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
class CAME(BaseOptimizer):
    """Factored adaptive updates with confidence weighted momentum.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for gradient momentum, squared gradients, and squared update residuals.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        clip_threshold: Maximum root mean square of the preconditioned update.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        eps1: Stability constant added to squared gradients.
        eps2: Stability constant added to squared update residuals.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 2e-4,
        betas: Betas = (0.9, 0.999, 0.9999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        clip_threshold: float = 1.0,
        ams_bound: bool = False,
        eps1: float = 1e-30,
        eps2: float = 1e-16,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps1, 'eps1')
        self.validate_non_negative(eps2, 'eps2')

        self.clip_threshold = clip_threshold
        self.eps1 = eps1
        self.eps2 = eps2
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'ams_bound': ams_bound,
            'eps1': eps1,
            'eps2': eps2,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'CAME'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            grad_shape: tuple[int, ...] = grad.shape
            factored: bool = self.get_options(grad_shape)

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)

                if factored:
                    state['exp_avg_sq_row'] = torch.zeros(grad_shape[:-1], dtype=grad.dtype, device=grad.device)
                    state['exp_avg_sq_col'] = torch.zeros(
                        grad_shape[:-2] + grad_shape[-1:], dtype=grad.dtype, device=grad.device
                    )
                    state['exp_avg_res_row'] = torch.zeros(grad_shape[:-1], dtype=grad.dtype, device=grad.device)
                    state['exp_avg_res_col'] = torch.zeros(
                        grad_shape[:-2] + grad_shape[-1:], dtype=grad.dtype, device=grad.device
                    )
                else:
                    state['exp_avg_sq'] = torch.zeros_like(grad)

                if group['ams_bound']:
                    state['exp_avg_sq_hat'] = torch.zeros_like(grad)

                state['RMS'] = 0.0

    @staticmethod
    def get_options(shape: tuple[int, ...]) -> bool:
        """Return whether the gradient supports factored second moments."""
        return len(shape) >= 2

    @staticmethod
    def get_rms(x: torch.Tensor) -> torch.Tensor:
        """Compute the root mean square of a tensor."""
        return x.norm(2) / math.sqrt(x.numel())

    @staticmethod
    def approximate_sq_grad(
        exp_avg_sq_row: torch.Tensor,
        exp_avg_sq_col: torch.Tensor,
        output: torch.Tensor,
    ):
        """Write a factored inverse root second moment approximation to `output`."""
        r_factor: torch.Tensor = (exp_avg_sq_row / exp_avg_sq_row.mean(dim=-1, keepdim=True)).rsqrt_().unsqueeze(-1)
        c_factor: torch.Tensor = exp_avg_sq_col.unsqueeze(-2).rsqrt()
        torch.mul(r_factor, c_factor, out=output)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2, beta3 = group['betas']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                grad_shape: tuple[int, ...] = grad.shape
                factored: bool = self.get_options(grad_shape)

                state['RMS'] = self.get_rms(p)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                update = torch.mul(grad, grad).add_(self.eps1)

                if factored:
                    exp_avg_sq_row, exp_avg_sq_col = state['exp_avg_sq_row'], state['exp_avg_sq_col']

                    exp_avg_sq_row.mul_(beta2).add_(update.mean(dim=-1), alpha=1.0 - beta2)
                    exp_avg_sq_col.mul_(beta2).add_(update.mean(dim=-2), alpha=1.0 - beta2)

                    self.approximate_sq_grad(exp_avg_sq_row, exp_avg_sq_col, update)
                else:
                    exp_avg_sq = state['exp_avg_sq']
                    exp_avg_sq.mul_(beta2).add_(update, alpha=1.0 - beta2)
                    torch.rsqrt(exp_avg_sq, out=update)

                if group['ams_bound']:
                    exp_avg_sq_hat = state['exp_avg_sq_hat']
                    torch.max(exp_avg_sq_hat, 1 / update, out=exp_avg_sq_hat)
                    torch.rsqrt(exp_avg_sq_hat / beta2, out=update)

                update.mul_(grad)

                update.div_((self.get_rms(update) / self.clip_threshold).clamp_(min=1.0))

                exp_avg = state['exp_avg']
                exp_avg.lerp_(update, weight=1.0 - beta1)

                res = update - exp_avg
                res.pow_(2).add_(self.eps2)

                if factored:
                    exp_avg_res_row, exp_avg_res_col = state['exp_avg_res_row'], state['exp_avg_res_col']

                    exp_avg_res_row.mul_(beta3).add_(res.mean(dim=-1), alpha=1.0 - beta3)
                    exp_avg_res_col.mul_(beta3).add_(res.mean(dim=-2), alpha=1.0 - beta3)

                    self.approximate_sq_grad(exp_avg_res_row, exp_avg_res_col, update)
                    update.mul_(exp_avg)
                else:
                    update = exp_avg

                p.add_(update, alpha=-group['lr'])

        return loss

approximate_sq_grad(exp_avg_sq_row, exp_avg_sq_col, output) staticmethod

Write a factored inverse root second moment approximation to output.

Source code in pytorch_optimizer/optimizer/came.py
120
121
122
123
124
125
126
127
128
129
@staticmethod
def approximate_sq_grad(
    exp_avg_sq_row: torch.Tensor,
    exp_avg_sq_col: torch.Tensor,
    output: torch.Tensor,
):
    """Write a factored inverse root second moment approximation to `output`."""
    r_factor: torch.Tensor = (exp_avg_sq_row / exp_avg_sq_row.mean(dim=-1, keepdim=True)).rsqrt_().unsqueeze(-1)
    c_factor: torch.Tensor = exp_avg_sq_col.unsqueeze(-2).rsqrt()
    torch.mul(r_factor, c_factor, out=output)

get_options(shape) staticmethod

Return whether the gradient supports factored second moments.

Source code in pytorch_optimizer/optimizer/came.py
110
111
112
113
@staticmethod
def get_options(shape: tuple[int, ...]) -> bool:
    """Return whether the gradient supports factored second moments."""
    return len(shape) >= 2

get_rms(x) staticmethod

Compute the root mean square of a tensor.

Source code in pytorch_optimizer/optimizer/came.py
115
116
117
118
@staticmethod
def get_rms(x: torch.Tensor) -> torch.Tensor:
    """Compute the root mean square of a tensor."""
    return x.norm(2) / math.sqrt(x.numel())

centralize_gradient(grad, gc_conv_only=False)

Subtract the mean of each gradient channel in place.

Parameters:

Name Type Description Default
grad Tensor

Gradient tensor.

required
gc_conv_only bool

If False, apply GC to both convolutional and fully connected layers. If True, apply only to convolutional layers.

False
Source code in pytorch_optimizer/optimizer/gradient_centralization.py
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
def centralize_gradient(grad: torch.Tensor, gc_conv_only: bool = False) -> None:
    """Subtract the mean of each gradient channel in place.

    Args:
        grad: Gradient tensor.
        gc_conv_only: If False, apply GC to both convolutional and fully connected layers. If True, apply only to
            convolutional layers.

    """
    size: int = grad.dim()
    if (gc_conv_only and size > 3) or (not gc_conv_only and size > 1):
        grad.add_(-grad.mean(dim=tuple(range(1, size)), keepdim=True))

Conda

Bases: BaseOptimizer

Adam with gradient projection in a basis derived from momentum.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
update_proj_gap int

Number of steps between low rank projection updates.

2000
scale float

Scaling factor for the projected update.

1.0
projection_type PROJECTION_TYPE

The type of the projection.

'std'
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/conda.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
class Conda(BaseOptimizer):
    """Adam with gradient projection in a basis derived from momentum.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        update_proj_gap: Number of steps between low rank projection updates.
        scale: Scaling factor for the projected update.
        projection_type: The type of the projection.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        update_proj_gap: int = 2000,
        scale: float = 1.0,
        projection_type: PROJECTION_TYPE = 'std',
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_positive(update_proj_gap, 'update_proj_gap')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'update_proj_gap': update_proj_gap,
            'scale': scale,
            'projection_type': projection_type,
            'eps': eps,
            **kwargs,
        }
        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Conda'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)

                shape = list(p.shape)
                if p.dim() == 2:
                    rank = min(shape)
                    if group['projection_type'] in ('left', 'full', 'reverse_std'):
                        shape[0] = rank
                    if group['projection_type'] in ('right', 'full', 'reverse_std'):
                        shape[1] = rank

                state['exp_avg_sq'] = p.new_zeros(shape)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size: float = group['lr'] * bias_correction2_sq / bias_correction1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']
                exp_avg.lerp_(grad, weight=1.0 - beta1)

                if p.dim() == 2:
                    if 'projector' not in state:
                        state['projector'] = GaLoreProjector(
                            rank=None,
                            update_proj_gap=group['update_proj_gap'],
                            scale=group['scale'],
                            projection_type=group['projection_type'],
                        )

                    grad = state['projector'].project(grad, group['step'], exp_avg)
                    exp_avg = state['projector'].project(exp_avg, group['step'])

                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                de_nom = exp_avg_sq.sqrt().add_(group['eps'])

                norm_grad = exp_avg / de_nom

                if p.dim() == 2:
                    norm_grad = state['projector'].project_back(norm_grad)

                p.add_(norm_grad, alpha=-step_size)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=True,
                    fixed_decay=False,
                )

        return loss

DAdaptAdaGrad

Bases: BaseOptimizer

AdaGrad with D-Adaptation.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Multiplier for the adapted step size. Use 1.0 unless training is unstable.

1.0
momentum float

Momentum factor.

0.0
d0 float

Initial estimate of the distance to the optimum.

1e-06
growth_rate float

Maximum multiplicative growth of the distance estimate per step.

float('inf')
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

0.0
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/dadapt.py
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
class DAdaptAdaGrad(BaseOptimizer):
    """AdaGrad with D-Adaptation.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Multiplier for the adapted step size. Use `1.0` unless training is unstable.
        momentum: Momentum factor.
        d0: Initial estimate of the distance to the optimum.
        growth_rate: Maximum multiplicative growth of the distance estimate per step.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1.0,
        momentum: float = 0.0,
        d0: float = 1e-6,
        growth_rate: float = float('inf'),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        eps: float = 0.0,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0, range_type='[)')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'momentum': momentum,
            'd': d0,
            'growth_rate': growth_rate,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'k': 0,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'DAdaptAdaGrad'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if 'alpha_k' not in state:
                state['alpha_k'] = torch.full_like(p, fill_value=1e-6)
                state['sk'] = torch.zeros_like(p)
                state['x0'] = torch.clone(p)
                if p.grad.is_sparse:
                    state['weighted_sk'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:  # noqa: PLR0912, PLR0915
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        group = self.param_groups[0]
        device = group['params'][0].device

        d, lr = group['d'], group['lr']
        d_lr: float = d * lr

        g_sq = torch.tensor([0.0], device=device)
        sk_sq_weighted_change = torch.tensor([0.0], device=device)
        sk_l1_change = torch.tensor([0.0], device=device)
        if 'gsq_weighted' not in group:
            group['gsq_weighted'] = torch.tensor([0.0], device=device)
        if 'sk_sq_weighted' not in group:
            group['sk_sq_weighted'] = torch.tensor([0.0], device=device)
        if 'sk_l1' not in group:
            group['sk_l1'] = torch.tensor([0.0], device=device)

        gsq_weighted = group['gsq_weighted']
        sk_sq_weighted = group['sk_sq_weighted']
        sk_l1 = group['sk_l1']

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            eps = group['eps']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                sk, alpha_k = state['sk'], state['alpha_k']

                if grad.is_sparse:
                    weighted_sk = state['weighted_sk']

                    grad = grad.coalesce()

                    vk = grad._values().pow(2)
                    sk_masked = sk.sparse_mask(grad).coalesce()
                    old_sk_l1_masked = sk_masked._values().abs().sum()

                    sk.add_(grad, alpha=d_lr)

                    sk_masked = sk.sparse_mask(grad).coalesce()
                    alpha_k_masked = alpha_k.sparse_mask(grad).coalesce()
                    weighted_sk_masked = weighted_sk.sparse_mask(grad).coalesce()

                    # update alpha before step
                    alpha_k_p1_masked = alpha_k_masked._values() + vk

                    alpha_k_delta_masked = alpha_k_p1_masked - alpha_k_masked._values()
                    alpha_k_delta = torch.sparse_coo_tensor(
                        grad.indices(),
                        alpha_k_delta_masked,
                        grad.shape,
                        check_invariants=False,
                    )
                    alpha_k.add_(alpha_k_delta)

                    de_nom = torch.sqrt(alpha_k_p1_masked + eps)

                    grad_sq = vk.div(de_nom).sum()
                    g_sq.add_(grad_sq)

                    # update weighted sk sq tracking
                    weighted_sk_p1_masked = sk_masked._values().pow(2).div(de_nom)

                    sk_sq_weighted_change.add_(weighted_sk_p1_masked.sum() - weighted_sk_masked._values().sum())

                    weighted_sk_p1_delta_masked = weighted_sk_p1_masked - weighted_sk_masked._values()
                    weighted_sk_p1_delta = torch.sparse_coo_tensor(
                        grad.indices(),
                        weighted_sk_p1_delta_masked,
                        grad.shape,
                        check_invariants=False,
                    )
                    weighted_sk.add_(weighted_sk_p1_delta)

                    sk_l1_masked = sk_masked._values().abs().sum()
                    sk_l1_change.add_(sk_l1_masked - old_sk_l1_masked)
                else:
                    self.apply_weight_decay(
                        p=p,
                        grad=grad,
                        lr=group['lr'],
                        weight_decay=group['weight_decay'],
                        weight_decouple=group['weight_decouple'],
                        fixed_decay=group['fixed_decay'],
                    )

                    old_sk_sq_weighted_param = sk.pow(2).div(torch.sqrt(alpha_k) + eps).sum()
                    old_sk_l1_param = sk.abs().sum()

                    alpha_k.add_(grad.pow(2))
                    grad_sq = grad.pow(2).div(torch.sqrt(alpha_k) + eps).sum()
                    g_sq.add_(grad_sq)

                    sk.add_(grad, alpha=d_lr)

                    sk_sq_weighted_param = sk.pow(2).div(torch.sqrt(alpha_k) + eps).sum()
                    sk_l1_param = sk.abs().sum()

                    sk_sq_weighted_change.add_(sk_sq_weighted_param - old_sk_sq_weighted_param)
                    sk_l1_change.add_(sk_l1_param - old_sk_l1_param)

        sk_sq_weighted.add_(sk_sq_weighted_change)
        gsq_weighted.add_(g_sq, alpha=d_lr ** 2)  # fmt: skip
        sk_l1.add_(sk_l1_change)

        if sk_l1 == 0:
            return loss

        if lr > 0.0:
            d_hat = (sk_sq_weighted - gsq_weighted) / sk_l1
            d = group['d'] = max(d, min(d_hat.item(), d * group['growth_rate']))

        for group in self.param_groups:
            group['gsq_weighted'] = gsq_weighted
            group['sk_sq_weighted'] = sk_sq_weighted
            group['sk_l1'] = sk_l1
            group['d'] = d

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                state = self.state[p]

                alpha_k, sk, x0 = state['alpha_k'], state['sk'], state['x0']

                if grad.is_sparse:
                    grad = grad.coalesce()

                    sk_masked = sk.sparse_mask(grad).coalesce()._values()
                    alpha_k_masked = alpha_k.sparse_mask(grad).coalesce()._values()
                    x0_masked = x0.sparse_mask(grad).coalesce()._values()
                    p_masked = p.sparse_mask(grad).coalesce()._values()

                    loc_masked = x0_masked - sk_masked.div(torch.sqrt(alpha_k_masked + group['eps']))

                    loc_delta_masked = loc_masked - p_masked
                    loc_delta = torch.sparse_coo_tensor(
                        grad.indices(),
                        loc_delta_masked,
                        grad.shape,
                        check_invariants=False,
                    )
                    p.add_(loc_delta)
                else:
                    z = x0 - sk.div(alpha_k.sqrt().add_(group['eps']))

                    if group['momentum'] > 0.0:
                        p.lerp_(z, weight=1.0 - group['momentum'])
                    else:
                        p.copy_(z)

            group['k'] += 1

        return loss

DAdaptAdam

Bases: BaseOptimizer

Adam with D-Adaptation V3.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Multiplier for the adapted step size. Use 1.0 unless training is unstable.

1.0
betas Betas

Decay rates for the gradient mean and squared gradients.

(0.9, 0.999)
d0 float

Initial estimate of the distance to the optimum.

1e-06
growth_rate float

Maximum multiplicative growth of the distance estimate per step.

float('inf')
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
bias_correction bool

Apply bias correction to the moment estimates.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/dadapt.py
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
class DAdaptAdam(BaseOptimizer):
    """Adam with D-Adaptation V3.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Multiplier for the adapted step size. Use `1.0` unless training is unstable.
        betas: Decay rates for the gradient mean and squared gradients.
        d0: Initial estimate of the distance to the optimum.
        growth_rate: Maximum multiplicative growth of the distance estimate per step.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        bias_correction: Apply bias correction to the moment estimates.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1.0,
        betas: Betas = (0.9, 0.999),
        d0: float = 1e-6,
        growth_rate: float = float('inf'),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        bias_correction: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'd': d0,
            'growth_rate': growth_rate,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'bias_correction': bias_correction,
            'step': 0,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'DAdaptAdam'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['s'] = torch.zeros_like(p)
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        group = self.param_groups[0]
        device = group['params'][0].device

        beta1, beta2 = group['betas']

        beta2_sq: float = math.sqrt(beta2)

        d: float = group['d']
        lr: float = group['lr']

        bias_correction1: float = 1.0 - beta1 ** (group['step'] + 1)
        bias_correction2_sq: float = math.sqrt(1.0 - beta2 ** (group['step'] + 1))
        bias_correction: float = bias_correction1 / bias_correction2_sq

        d_lr: float = self.apply_adam_debias(
            not group['bias_correction'], step_size=d * lr, bias_correction1=bias_correction
        )

        sk_l1 = torch.tensor([0.0], device=device)
        numerator_acc = torch.tensor([0.0], device=device)

        if 'numerator_weighted' not in group:
            group['numerator_weighted'] = torch.tensor([0.0], device=device)
        numerator_weighted = group['numerator_weighted']

        for group in self.param_groups:
            self.init_group(group)

            group['step'] += 1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq, s = state['exp_avg'], state['exp_avg_sq'], state['s']

                de_nom = exp_avg_sq.sqrt().add_(group['eps'])
                numerator_acc.add_(torch.dot(grad.flatten(), s.div(de_nom).flatten()), alpha=d_lr)

                exp_avg.mul_(beta1).add_(grad, alpha=d_lr * (1.0 - beta1))
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                s.mul_(beta2_sq).add_(grad, alpha=d_lr * (1.0 - beta2_sq))

                sk_l1.add_(s.abs().sum())

        if sk_l1 == 0:
            return loss

        numerator_weighted.lerp_(numerator_acc, weight=1.0 - beta2_sq)  # fmt: skip

        if lr > 0.0:
            d_hat = numerator_weighted / ((1.0 - beta2_sq) * sk_l1)
            d = max(d, min(d_hat.item(), d * group['growth_rate']))

        for group in self.param_groups:
            group['numerator_weighted'] = numerator_weighted
            group['d'] = d

            for p in group['params']:
                if p.grad is None:
                    continue

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                de_nom = exp_avg_sq.sqrt().add_(group['eps'])

                self.apply_weight_decay(
                    p=p,
                    grad=None,
                    lr=d_lr,
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                p.addcdiv_(exp_avg, de_nom, value=-1.0)

        return loss

DAdaptAdan

Bases: BaseOptimizer

Adan with D-Adaptation.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Multiplier for the adapted step size. Use 1.0 unless training is unstable.

1.0
betas Betas

Decay rates for gradients, gradient differences, and squared corrected gradients.

(0.98, 0.92, 0.99)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
d0 float

Initial estimate of the distance to the optimum.

1e-06
growth_rate float

Maximum multiplicative growth of the distance estimate per step.

float('inf')
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/dadapt.py
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
class DAdaptAdan(BaseOptimizer):
    """Adan with D-Adaptation.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Multiplier for the adapted step size. Use `1.0` unless training is unstable.
        betas: Decay rates for gradients, gradient differences, and squared corrected gradients.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        d0: Initial estimate of the distance to the optimum.
        growth_rate: Maximum multiplicative growth of the distance estimate per step.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1.0,
        betas: Betas = (0.98, 0.92, 0.99),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        d0: float = 1e-6,
        growth_rate: float = float('inf'),
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'd': d0,
            'growth_rate': growth_rate,
            'k': 0,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'DAdaptAdan'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if 'exp_avg' not in state:
                state['s'] = torch.zeros_like(p)
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['exp_avg_diff'] = torch.zeros_like(p)
                state['previous_grad'] = -grad.clone()

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        group = self.param_groups[0]

        beta1, beta2, beta3 = group['betas']
        growth_rate = group['growth_rate']

        d, lr = group['d'], group['lr']
        d_lr = float(d * lr)

        g_sq = torch.tensor([0.0], device=group['params'][0].device)
        sk_sq_weighted = torch.tensor([0.0], device=group['params'][0].device)
        sk_l1 = torch.tensor([0.0], device=group['params'][0].device)
        if 'gsq_weighted' not in group:
            group['gsq_weighted'] = torch.tensor([0.0], device=group['params'][0].device)
        gsq_weighted = group['gsq_weighted']

        for group in self.param_groups:
            self.init_group(group)

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                grad_diff = state['previous_grad']
                grad_diff.add_(grad)

                exp_avg, exp_avg_sq, exp_avg_diff = state['exp_avg'], state['exp_avg_sq'], state['exp_avg_diff']

                exp_avg.mul_(beta1).add_(grad, alpha=d_lr * (1.0 - beta1))
                exp_avg_diff.mul_(beta2).add_(grad_diff, alpha=d_lr * (1.0 - beta2))

                grad_diff.mul_(beta2).add_(grad)
                grad_diff = to_real(grad_diff * grad_diff.conj())
                exp_avg_sq.mul_(beta3).addcmul_(grad_diff, grad_diff, value=1.0 - beta3)

                grad_power = to_real(grad * grad.conj())
                de_nom = exp_avg_sq.sqrt().add_(group['eps'])

                g_sq.add_(grad_power.div_(de_nom).sum())

                s = state['s']
                s.mul_(beta3).add_(grad, alpha=d_lr * (1.0 - beta3))

                sk_sq_weighted.add_(to_real(s * s.conj()).div_(de_nom).sum())
                sk_l1.add_(s.abs().sum())

                state['previous_grad'].copy_(-grad)

        if sk_l1 == 0:
            return loss

        gsq_weighted.mul_(beta3).add_(g_sq, alpha=(d_lr ** 2) * (1.0 - beta3))  # fmt: skip

        if lr > 0.0:
            d_hat = (sk_sq_weighted / (1.0 - beta3) - gsq_weighted) / sk_l1
            d = max(d, min(d_hat, d * growth_rate))

        for group in self.param_groups:
            group['step'] += 1

            group['gsq_weighted'] = gsq_weighted
            group['d'] = d
            for p in group['params']:
                if p.grad is None:
                    continue

                state = self.state[p]

                exp_avg, exp_avg_sq, exp_avg_diff = state['exp_avg'], state['exp_avg_sq'], state['exp_avg_diff']

                de_nom = exp_avg_sq.sqrt().add_(group['eps'])

                if group['weight_decouple']:
                    p.mul_(1.0 - d_lr * group['weight_decay'])

                p.addcdiv_(exp_avg, de_nom, value=-1.0)
                p.addcdiv_(exp_avg_diff, de_nom, value=-beta2)

                if not group['weight_decouple']:
                    p.div_(1.0 + d_lr * group['weight_decay'])

            group['k'] += 1

        return loss

DAdaptLion

Bases: BaseOptimizer

Lion with D-Adaptation V3.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Multiplier for the adapted step size. Use 1.0 unless training is unstable.

1.0
betas Betas

Decay rates for update interpolation and gradient momentum.

(0.9, 0.999)
d0 float

Initial estimate of the distance to the optimum.

1e-06
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/dadapt.py
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
class DAdaptLion(BaseOptimizer):
    """Lion with D-Adaptation V3.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Multiplier for the adapted step size. Use `1.0` unless training is unstable.
        betas: Decay rates for update interpolation and gradient momentum.
        d0: Initial estimate of the distance to the optimum.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1.0,
        betas: Betas = (0.9, 0.999),
        d0: float = 1e-6,
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'd': d0,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'step': 0,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'DAdaptLion'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['s'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        group = self.param_groups[0]
        device = group['params'][0].device

        if 'numerator_weighted' not in group:
            group['numerator_weighted'] = torch.tensor([0.0], device=device)
        numerator_weighted = group['numerator_weighted']

        sk_l1 = torch.tensor([0.0], device=device)
        numerator_accumulator = torch.tensor([0.0], device=device)

        beta1, beta2 = group['betas']
        beta2_sq = math.sqrt(beta2)

        d, lr = group['d'], group['lr']
        d_lr: float = d * lr

        for group in self.param_groups:
            self.init_group(group)

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=d_lr,
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                exp_avg, s = state['exp_avg'], state['s']

                update = exp_avg.clone().lerp_(grad, weight=1.0 - beta1).sign_()
                p.add_(update, alpha=-d_lr)

                exp_avg.mul_(beta2).add_(grad, alpha=(1.0 - beta2) * d_lr)

                numerator_accumulator.add_(torch.dot(update.flatten(), s.flatten()), alpha=d_lr)
                s.mul_(beta2_sq).add_(update, alpha=(1.0 - beta2_sq) * d_lr)

                sk_l1.add_(s.abs().sum())

        numerator_weighted.lerp_(numerator_accumulator, weight=1.0 - beta2_sq)

        if sk_l1 == 0:
            return loss

        if lr > 0.0:
            d_hat: float = (numerator_weighted / ((1.0 - beta2_sq) * sk_l1)).item()
            d = max(d, d_hat)

        for group in self.param_groups:
            group['step'] += 1

            group['numerator_weighted'] = numerator_weighted
            group['d'] = d

        return loss

DAdaptSGD

Bases: BaseOptimizer

SGD with D-Adaptation V3.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Multiplier for the adapted step size. Use 1.0 unless training is unstable.

1.0
momentum float

Momentum factor.

0.9
d0 float

Initial estimate of the distance to the optimum.

1e-06
growth_rate float

Maximum multiplicative growth of the distance estimate per step.

float('inf')
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/dadapt.py
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
class DAdaptSGD(BaseOptimizer):
    """SGD with D-Adaptation V3.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Multiplier for the adapted step size. Use `1.0` unless training is unstable.
        momentum: Momentum factor.
        d0: Initial estimate of the distance to the optimum.
        growth_rate: Maximum multiplicative growth of the distance estimate per step.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1.0,
        momentum: float = 0.9,
        d0: float = 1e-6,
        growth_rate: float = float('inf'),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0, range_type='[)')
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'momentum': momentum,
            'd': d0,
            'growth_rate': growth_rate,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'step': 0,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'DAdaptSGD'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['z'] = p.clone()
                state['s'] = torch.zeros_like(p)
                state['x0'] = p.clone()

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        group = self.param_groups[0]
        device = group['params'][0].device

        sk_sq = torch.tensor([0.0], device=device)
        if 'numerator_weighted' not in group:
            group['numerator_weighted'] = torch.tensor([0.0], device=device)
        numerator_weighted = group['numerator_weighted']

        if group['step'] == 0:
            group['g0_norm'] = get_global_gradient_norm(self.param_groups).sqrt_().item()
        g0_norm = group['g0_norm']

        if g0_norm == 0:
            return loss

        d, lr = group['d'], group['lr']
        d_lr: float = d * lr / g0_norm

        for group in self.param_groups:
            self.init_group(group)

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=None,
                    lr=d_lr,
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                s = state['s']
                numerator_weighted.add_(torch.dot(grad.flatten(), s.flatten()), alpha=d_lr)

                s.add_(grad, alpha=d_lr)
                sk_sq.add_(s.pow(2).sum())

        if lr > 0.0:
            d_hat = 2.0 * numerator_weighted / sk_sq.sqrt()
            d = max(d, min(d_hat.item(), d * group['growth_rate']))

        for group in self.param_groups:
            group['step'] += 1

            group['numerator_weighted'] = numerator_weighted
            group['d'] = d

            for p in group['params']:
                if p.grad is None:
                    continue

                state = self.state[p]

                z = state['z']
                z.copy_(state['x0'] - state['s'])

                p.lerp_(z, weight=1.0 - group['momentum'])

        return loss

DASH

Bases: BaseOptimizer

Accelerated Shampoo with batched blocks and Newton-Denman-Beavers inverse roots.

Implements local, layerwise DASH from https://arxiv.org/abs/2602.02016. Equal-sized blocks, including left and right factors, share batched storage. Scalars and squeezed vectors use one-sided inverse-square-root preconditioning; higher-order tensors are flattened after the first non-singleton dimension. Adam grafting rescales each block independently. Optimizer states use at least float32 precision.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the gradient EMA and Shampoo statistics.

(0.9, 0.95)
grafting_beta float | None

Decay rate for Adam grafting statistics. None uses beta2.

None
weight_decay float

Decoupled weight decay coefficient.

0.0
block_size int

Maximum block dimension. Edge blocks are processed without padding.

1024
precondition_frequency int

Number of steps between inverse-root updates.

10
start_preconditioning_step int

First step using Shampoo. Earlier steps use Adam grafting.

1
inverse_root_method Literal['newton_db', 'eigh']

Inverse-root solver: 'newton_db' or 'eigh'.

'newton_db'
matrix_scaling Literal['power', 'frobenius']

Newton-DB scaling: 'power' or 'frobenius'.

'power'
newton_steps int

Number of iterations per Newton-DB square-root computation.

10
power_iteration_steps int

Number of iterations for spectral-radius estimation.

10
power_iteration_vectors int

Number of parallel starting vectors for power iteration.

16
momentum float

Momentum coefficient for the grafted update.

0.0
nesterov bool

Whether to use Nesterov update momentum.

True
correct_bias bool

Whether to correct bias in the Adam grafting direction.

True
eps float

Term added to the Adam grafting denominator.

1e-08
matrix_eps float

Diagonal regularization for inverse roots. Must be positive.

1e-10
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/dash.py
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
class DASH(BaseOptimizer):
    """Accelerated Shampoo with batched blocks and Newton-Denman-Beavers inverse roots.

    Implements local, layerwise DASH from https://arxiv.org/abs/2602.02016. Equal-sized blocks, including left and
    right factors, share batched storage. Scalars and squeezed vectors use one-sided inverse-square-root
    preconditioning; higher-order tensors are flattened after the first non-singleton dimension.
    Adam grafting rescales each block independently. Optimizer states use at least float32 precision.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the gradient EMA and Shampoo statistics.
        grafting_beta: Decay rate for Adam grafting statistics. `None` uses `beta2`.
        weight_decay: Decoupled weight decay coefficient.
        block_size: Maximum block dimension. Edge blocks are processed without padding.
        precondition_frequency: Number of steps between inverse-root updates.
        start_preconditioning_step: First step using Shampoo. Earlier steps use Adam grafting.
        inverse_root_method: Inverse-root solver: `'newton_db'` or `'eigh'`.
        matrix_scaling: Newton-DB scaling: `'power'` or `'frobenius'`.
        newton_steps: Number of iterations per Newton-DB square-root computation.
        power_iteration_steps: Number of iterations for spectral-radius estimation.
        power_iteration_vectors: Number of parallel starting vectors for power iteration.
        momentum: Momentum coefficient for the grafted update.
        nesterov: Whether to use Nesterov update momentum.
        correct_bias: Whether to correct bias in the Adam grafting direction.
        eps: Term added to the Adam grafting denominator.
        matrix_eps: Diagonal regularization for inverse roots. Must be positive.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.95),
        grafting_beta: float | None = None,
        weight_decay: float = 0.0,
        block_size: int = 1024,
        precondition_frequency: int = 10,
        start_preconditioning_step: int = 1,
        inverse_root_method: Literal['newton_db', 'eigh'] = 'newton_db',
        matrix_scaling: Literal['power', 'frobenius'] = 'power',
        newton_steps: int = 10,
        power_iteration_steps: int = 10,
        power_iteration_vectors: int = 16,
        momentum: float = 0.0,
        nesterov: bool = True,
        correct_bias: bool = True,
        eps: float = 1e-8,
        matrix_eps: float = 1e-10,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)

        if grafting_beta is not None:
            self.validate_range(grafting_beta, 'grafting_beta', 0.0, 1.0)

        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(momentum, 'momentum', 0.0, 1.0)
        self.validate_non_negative(eps, 'eps')
        self.validate_positive(matrix_eps, 'matrix_eps')

        self.validate_options(inverse_root_method, 'inverse_root_method', ['newton_db', 'eigh'])
        self.validate_options(matrix_scaling, 'matrix_scaling', ['power', 'frobenius'])

        self.validate_positive(block_size, 'block_size')
        self.validate_positive(precondition_frequency, 'precondition_frequency')
        self.validate_positive(start_preconditioning_step, 'start_preconditioning_step')
        self.validate_positive(newton_steps, 'newton_steps')
        self.validate_positive(power_iteration_steps, 'power_iteration_steps')
        self.validate_positive(power_iteration_vectors, 'power_iteration_vectors')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'grafting_beta': grafting_beta,
            'weight_decay': weight_decay,
            'block_size': block_size,
            'precondition_frequency': precondition_frequency,
            'start_preconditioning_step': start_preconditioning_step,
            'inverse_root_method': inverse_root_method,
            'matrix_scaling': matrix_scaling,
            'newton_steps': newton_steps,
            'power_iteration_steps': power_iteration_steps,
            'power_iteration_vectors': power_iteration_vectors,
            'momentum': momentum,
            'nesterov': nesterov,
            'correct_bias': correct_bias,
            'eps': eps,
            'matrix_eps': matrix_eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'DASH'

    def _restore_state_types(self, value, saved_value):
        if isinstance(saved_value, torch.Tensor) and saved_value.is_floating_point():
            return saved_value.to(device=value.device)

        return super()._restore_state_types(value, saved_value)

    @staticmethod
    def partition(grad: torch.Tensor, block_size: int) -> Iterator[tuple[tuple[int, int, int, int], torch.Tensor]]:
        """Yield up to four batches of equal-sized blocks and their matrix bounds."""
        rows, cols = grad.shape

        row_start = 0
        for row_size in (rows // block_size * block_size, rows % block_size):
            col_start = 0
            for col_size in (cols // block_size * block_size, cols % block_size):
                if row_size and col_size:
                    height, width = min(row_size, block_size), min(col_size, block_size)
                    region = grad[row_start : row_start + row_size, col_start : col_start + col_size]
                    blocks = region.reshape(row_size // height, height, col_size // width, width)

                    yield (row_start, col_start, row_size, col_size), blocks.transpose(1, 2).reshape(-1, height, width)

                col_start += col_size

            row_start += row_size

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            if p.grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['blocks'] = []

                shape = p.squeeze().shape
                state['shape'] = (shape[0], math.prod(shape[1:])) if len(shape) > 1 else (p.numel(), 1)
                state['one_sided'] = len(shape) < 2

                dtype = torch.float64 if p.dtype == torch.float64 else torch.float32

                for _, block in self.partition(p.grad.reshape(state['shape']), group['block_size']):
                    batch, rows, cols = block.shape

                    sizes = (
                        [(batch, rows, rows)]
                        if state['one_sided']
                        else (
                            [(2 * batch, rows, rows)] if rows == cols else [(batch, rows, rows), (batch, cols, cols)]
                        )
                    )

                    block_state = {
                        'exp_avg_sq': torch.zeros_like(block, dtype=dtype),
                        'statistics': [torch.zeros(size, device=p.device, dtype=dtype) for size in sizes],
                        'inverse_roots': [],
                    }

                    if group['betas'][0] > 0.0:
                        block_state['exp_avg'] = torch.zeros_like(block, dtype=dtype)

                    if group['momentum'] > 0.0:
                        block_state['momentum'] = torch.zeros_like(block, dtype=dtype)

                    state['blocks'].append(block_state)

    @staticmethod
    def scale_matrix(matrix: torch.Tensor, group: ParamGroup) -> torch.Tensor:
        """Estimate batched matrix scales using the reference's bfloat16 power iteration."""
        dtype = matrix.dtype
        matrix = matrix.to(torch.bfloat16)

        if group['matrix_scaling'] == 'frobenius':
            return torch.linalg.vector_norm(matrix, dim=(-2, -1), keepdim=True).to(dtype)

        matrix.diagonal(dim1=-2, dim2=-1).add_(1e-6)

        scale = batched_power_iteration(
            matrix, num_iters=group['power_iteration_steps'], num_vectors=group['power_iteration_vectors']
        )

        return scale.to(dtype).mul_(2.0)

    def compute_inverse_root(self, matrix: torch.Tensor, root: int, group: ParamGroup) -> torch.Tensor:
        """Compute regularized batched inverse roots without modifying the statistics."""
        regularized = matrix.clone()
        regularized.diagonal(dim1=-2, dim2=-1).add_(group['matrix_eps'])

        if group['inverse_root_method'] == 'eigh':
            values, vectors = torch.linalg.eigh(regularized)

            # Match Distributed Shampoo's spectral shift after regularized eigendecomposition.
            values.add_(group['matrix_eps'] - values.amin(dim=-1, keepdim=True).clamp_max_(0.0)).pow_(-1.0 / root)

            return (vectors * values.unsqueeze(-2)) @ vectors.transpose(-2, -1)

        scale = self.scale_matrix(regularized, group).clamp_min_(torch.finfo(matrix.dtype).tiny)
        if root == 4:
            regularized = compute_power_newton_db(regularized, scale, group['newton_steps'], inverse=False)
            scale.sqrt_()

        return compute_power_newton_db(regularized, scale, group['newton_steps'], inverse=True)

    def update_block(
        self,
        grad: torch.Tensor,
        state: dict,
        one_sided: bool,
        group: ParamGroup,
    ) -> torch.Tensor:
        """Update statistics and return a grafted, optionally momentum-filtered block batch."""
        beta1, beta2 = group['betas']
        grafting_beta = beta2 if group['grafting_beta'] is None else group['grafting_beta']
        step = group['step']

        bias1 = self.debias(beta1, step) if group['correct_bias'] else 1.0
        bias2 = self.debias(grafting_beta, step) if group['correct_bias'] else 1.0

        batch = grad.shape[0]
        grad = grad.to(state['exp_avg_sq'].dtype)

        statistics, inverse_roots = state['statistics'], state['inverse_roots']
        statistics[0][:batch].baddbmm_(grad, grad.transpose(1, 2), beta=beta2, alpha=1.0 - beta2)
        if not one_sided:
            statistics[-1][-batch:].baddbmm_(grad.transpose(1, 2), grad, beta=beta2, alpha=1.0 - beta2)

        state['exp_avg_sq'].mul_(grafting_beta).addcmul_(grad, grad, value=1.0 - grafting_beta)
        if beta1 > 0.0:
            grad = state['exp_avg'].lerp_(grad, weight=1.0 - beta1)

        graft = grad.div(bias1).div_(
            state['exp_avg_sq'].div(bias2).sqrt_().add_(group['eps']).clamp_min_(torch.finfo(grad.dtype).tiny)
        )

        start = group['start_preconditioning_step']
        if step >= start:
            graft_norm = torch.linalg.vector_norm(graft, dim=(1, 2), keepdim=True)
            del graft

            if not inverse_roots or step % group['precondition_frequency'] == 0:
                inverse_roots[:] = [
                    self.compute_inverse_root(statistic, 2 if one_sided else 4, group) for statistic in statistics
                ]

            update = inverse_roots[0][:batch] @ grad
            if not one_sided:
                update = update @ inverse_roots[-1][-batch:]

            update.mul_(graft_norm / torch.linalg.vector_norm(update, dim=(1, 2), keepdim=True).add_(1e-16))
        else:
            update = graft

        if group['momentum'] > 0.0:
            momentum = state['momentum'].mul_(group['momentum']).add_(update)
            update = update.add_(momentum, alpha=group['momentum']) if group['nesterov'] else momentum

        return update

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            for p in group['params']:
                if p.grad is None:
                    continue

                state = self.state[p]

                grad = p.grad.reshape(state['shape'])
                self.maximize_gradient(grad, self.maximize)

                matrix = p.reshape(state['shape'])

                for (bounds, block), block_state in zip(self.partition(grad, group['block_size']), state['blocks']):
                    direction = self.update_block(block, block_state, state['one_sided'], group)

                    row, col, rows, cols = bounds

                    height, width = block.shape[1:]

                    direction = direction.to(p.dtype).reshape(rows // height, cols // width, height, width)
                    direction = direction.transpose(1, 2)

                    region = matrix[row : row + rows, col : col + cols]
                    region = region.view(rows // height, height, cols // width, width)

                    self.apply_weight_decay(
                        region, None, group['lr'], group['weight_decay'], weight_decouple=True, fixed_decay=False
                    )

                    region.add_(direction, alpha=-group['lr'])

                if matrix.data_ptr() != p.data_ptr():
                    p.copy_(matrix.reshape_as(p))

        return loss

compute_inverse_root(matrix, root, group)

Compute regularized batched inverse roots without modifying the statistics.

Source code in pytorch_optimizer/optimizer/dash.py
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
def compute_inverse_root(self, matrix: torch.Tensor, root: int, group: ParamGroup) -> torch.Tensor:
    """Compute regularized batched inverse roots without modifying the statistics."""
    regularized = matrix.clone()
    regularized.diagonal(dim1=-2, dim2=-1).add_(group['matrix_eps'])

    if group['inverse_root_method'] == 'eigh':
        values, vectors = torch.linalg.eigh(regularized)

        # Match Distributed Shampoo's spectral shift after regularized eigendecomposition.
        values.add_(group['matrix_eps'] - values.amin(dim=-1, keepdim=True).clamp_max_(0.0)).pow_(-1.0 / root)

        return (vectors * values.unsqueeze(-2)) @ vectors.transpose(-2, -1)

    scale = self.scale_matrix(regularized, group).clamp_min_(torch.finfo(matrix.dtype).tiny)
    if root == 4:
        regularized = compute_power_newton_db(regularized, scale, group['newton_steps'], inverse=False)
        scale.sqrt_()

    return compute_power_newton_db(regularized, scale, group['newton_steps'], inverse=True)

partition(grad, block_size) staticmethod

Yield up to four batches of equal-sized blocks and their matrix bounds.

Source code in pytorch_optimizer/optimizer/dash.py
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
@staticmethod
def partition(grad: torch.Tensor, block_size: int) -> Iterator[tuple[tuple[int, int, int, int], torch.Tensor]]:
    """Yield up to four batches of equal-sized blocks and their matrix bounds."""
    rows, cols = grad.shape

    row_start = 0
    for row_size in (rows // block_size * block_size, rows % block_size):
        col_start = 0
        for col_size in (cols // block_size * block_size, cols % block_size):
            if row_size and col_size:
                height, width = min(row_size, block_size), min(col_size, block_size)
                region = grad[row_start : row_start + row_size, col_start : col_start + col_size]
                blocks = region.reshape(row_size // height, height, col_size // width, width)

                yield (row_start, col_start, row_size, col_size), blocks.transpose(1, 2).reshape(-1, height, width)

            col_start += col_size

        row_start += row_size

scale_matrix(matrix, group) staticmethod

Estimate batched matrix scales using the reference's bfloat16 power iteration.

Source code in pytorch_optimizer/optimizer/dash.py
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
@staticmethod
def scale_matrix(matrix: torch.Tensor, group: ParamGroup) -> torch.Tensor:
    """Estimate batched matrix scales using the reference's bfloat16 power iteration."""
    dtype = matrix.dtype
    matrix = matrix.to(torch.bfloat16)

    if group['matrix_scaling'] == 'frobenius':
        return torch.linalg.vector_norm(matrix, dim=(-2, -1), keepdim=True).to(dtype)

    matrix.diagonal(dim1=-2, dim2=-1).add_(1e-6)

    scale = batched_power_iteration(
        matrix, num_iters=group['power_iteration_steps'], num_vectors=group['power_iteration_vectors']
    )

    return scale.to(dtype).mul_(2.0)

update_block(grad, state, one_sided, group)

Update statistics and return a grafted, optionally momentum-filtered block batch.

Source code in pytorch_optimizer/optimizer/dash.py
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
def update_block(
    self,
    grad: torch.Tensor,
    state: dict,
    one_sided: bool,
    group: ParamGroup,
) -> torch.Tensor:
    """Update statistics and return a grafted, optionally momentum-filtered block batch."""
    beta1, beta2 = group['betas']
    grafting_beta = beta2 if group['grafting_beta'] is None else group['grafting_beta']
    step = group['step']

    bias1 = self.debias(beta1, step) if group['correct_bias'] else 1.0
    bias2 = self.debias(grafting_beta, step) if group['correct_bias'] else 1.0

    batch = grad.shape[0]
    grad = grad.to(state['exp_avg_sq'].dtype)

    statistics, inverse_roots = state['statistics'], state['inverse_roots']
    statistics[0][:batch].baddbmm_(grad, grad.transpose(1, 2), beta=beta2, alpha=1.0 - beta2)
    if not one_sided:
        statistics[-1][-batch:].baddbmm_(grad.transpose(1, 2), grad, beta=beta2, alpha=1.0 - beta2)

    state['exp_avg_sq'].mul_(grafting_beta).addcmul_(grad, grad, value=1.0 - grafting_beta)
    if beta1 > 0.0:
        grad = state['exp_avg'].lerp_(grad, weight=1.0 - beta1)

    graft = grad.div(bias1).div_(
        state['exp_avg_sq'].div(bias2).sqrt_().add_(group['eps']).clamp_min_(torch.finfo(grad.dtype).tiny)
    )

    start = group['start_preconditioning_step']
    if step >= start:
        graft_norm = torch.linalg.vector_norm(graft, dim=(1, 2), keepdim=True)
        del graft

        if not inverse_roots or step % group['precondition_frequency'] == 0:
            inverse_roots[:] = [
                self.compute_inverse_root(statistic, 2 if one_sided else 4, group) for statistic in statistics
            ]

        update = inverse_roots[0][:batch] @ grad
        if not one_sided:
            update = update @ inverse_roots[-1][-batch:]

        update.mul_(graft_norm / torch.linalg.vector_norm(update, dim=(1, 2), keepdim=True).add_(1e-16))
    else:
        update = graft

    if group['momentum'] > 0.0:
        momentum = state['momentum'].mul_(group['momentum']).add_(update)
        update = update.add_(momentum, alpha=group['momentum']) if group['nesterov'] else momentum

    return update

DeMo

Bases: SGD, BaseOptimizer

SGD with compressed distributed momentum exchange.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
compression_decay float

Decay rate for the residual momentum buffer.

0.999
compression_top_k int

Maximum DCT coefficients to retain per block.

32
compression_chunk int

Maximum size of each DCT block dimension.

64
process_group ProcessGroup | None

Distributed process group for compressed gradient exchange. None uses the default group.

None
weight_decay float

Weight decay coefficient.

0.0
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/demo.py
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
class DeMo(torch.optim.SGD, BaseOptimizer):  # pragma: no cover
    """SGD with compressed distributed momentum exchange.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        compression_decay: Decay rate for the residual momentum buffer.
        compression_top_k: Maximum DCT coefficients to retain per block.
        compression_chunk: Maximum size of each DCT block dimension.
        process_group: Distributed process group for compressed gradient exchange. `None` uses the default group.
        weight_decay: Weight decay coefficient.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        compression_decay: float = 0.999,
        compression_top_k: int = 32,
        compression_chunk: int = 64,
        weight_decay: float = 0.0,
        process_group: ProcessGroup | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(compression_decay, 'compression_decay', 0.0, 1.0, range_type='[)')
        self.validate_positive(compression_top_k, 'compression_top_k')
        self.validate_positive(compression_chunk, 'compression_chunk')

        self.weight_decay = weight_decay

        self.compression_decay = compression_decay
        self.compression_top_k = compression_top_k
        self.compression_chunk = compression_chunk
        self.process_group = process_group

        self.data_transmit: int = 0
        self.data_receive: int = 0

        self.maximize = maximize

        super().__init__(
            params,
            lr=lr,
            foreach=False,
            momentum=0.0,
            dampening=0.0,
            nesterov=False,
            maximize=False,
            weight_decay=0.0,
            **kwargs,
        )

        self.demo_state = {}
        self.init_demo_states()
        self.init_parameters()

        self.default_dtype: torch.dtype = self.find_dtype()
        self.transform = TransformDCT(self.param_groups, self.compression_chunk, norm='ortho')
        self.compress = CompressDCT()

    def __str__(self) -> str:
        return 'DeMo'

    def find_dtype(self) -> torch.dtype:
        """Return the data type of the first optimizer parameter."""
        for group in self.param_groups:
            for p in group['params']:
                if p.requires_grad:
                    return p.dtype
        return torch.float32

    def init_demo_states(self) -> None:
        for group in self.param_groups:
            for p in group['params']:
                if p.requires_grad:
                    self.demo_state[p] = {}

    def init_parameters(self) -> None:
        for group in self.param_groups:
            group['step'] = 0
            for p in group['params']:
                if p.requires_grad:
                    state = self.demo_state.get(p, {})

                    state['delta'] = torch.zeros_like(p)

    def demo_all_gather(self, sparse_idx, sparse_val):
        world_size: int = get_world_size() if self.process_group is None else self.process_group.size()

        sparse_idx_list = [torch.zeros_like(sparse_idx) for _ in range(world_size)]
        sparse_val_list = [torch.zeros_like(sparse_val) for _ in range(world_size)]

        sparse_idx_handle = all_gather(sparse_idx_list, sparse_idx, group=self.process_group, async_op=True)
        sparse_val_handle = all_gather(sparse_val_list, sparse_val, group=self.process_group, async_op=True)

        sparse_idx_handle.wait()
        sparse_val_handle.wait()

        return sparse_idx_list, sparse_val_list

    @torch.no_grad()
    def init_group(self):
        pass

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        self.data_transmit = 0
        self.data_receive = 0

        for group in self.param_groups:
            if 'step' in group:
                group['step'] += 1
            else:
                group['step'] = 1

            lr = group['lr']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad
                if grad.is_sparse:
                    raise NoSparseGradientError(str(self))

                if torch.is_complex(p):
                    raise NoComplexParameterError(str(self))

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.demo_state.get(p, {})

                self.apply_weight_decay(
                    p,
                    grad,
                    lr=lr,
                    weight_decay=self.weight_decay,
                    weight_decouple=True,
                    fixed_decay=False,
                )

                if self.compression_decay != 1:
                    state['delta'].mul_(self.compression_decay)

                state['delta'].add_(grad, alpha=lr)

                sparse_idx, sparse_val, x_shape = self.compress.compress(
                    self.transform.encode(state['delta']), self.compression_top_k
                )

                transmit_grad = self.transform.decode(self.compress.decompress(p, sparse_idx, sparse_val, x_shape))

                state['delta'].sub_(transmit_grad)

                sparse_idx_gather, sparse_val_gather = self.demo_all_gather(sparse_idx, sparse_val)

                self.data_transmit += sparse_idx.nbytes + sparse_val.nbytes
                for si, v in zip(sparse_idx_gather, sparse_val_gather):
                    self.data_receive += si.nbytes + v.nbytes

                new_grad = self.transform.decode(
                    self.compress.batch_decompress(p, sparse_idx_gather, sparse_val_gather, x_shape)
                )

                if p.grad is None:
                    p.grad = new_grad
                else:
                    p.grad.copy_(new_grad)

                p.grad.sign_()

        return super().step(closure)

find_dtype()

Return the data type of the first optimizer parameter.

Source code in pytorch_optimizer/optimizer/demo.py
358
359
360
361
362
363
364
def find_dtype(self) -> torch.dtype:
    """Return the data type of the first optimizer parameter."""
    for group in self.param_groups:
        for p in group['params']:
            if p.requires_grad:
                return p.dtype
    return torch.float32

DiffGrad

Bases: BaseOptimizer

Adam updates scaled by changes between consecutive gradients.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
rectify bool

Perform the rectified update similar to RAdam.

False
n_sma_threshold int

Minimum effective simple moving average length for rectification.

5
degenerated_to_sgd bool

Use an SGD update before the moving average reaches the rectification threshold.

True
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
Source code in pytorch_optimizer/optimizer/diffgrad.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
class DiffGrad(BaseOptimizer):
    """Adam updates scaled by changes between consecutive gradients.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        rectify: Perform the rectified update similar to RAdam.
        n_sma_threshold: Minimum effective simple moving average length for rectification.
        degenerated_to_sgd: Use an SGD update before the moving average reaches the rectification threshold.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.

    """

    _supports_compiled_foreach = True

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        rectify: bool = False,
        n_sma_threshold: int = 5,
        degenerated_to_sgd: bool = True,
        ams_bound: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        foreach: bool | None = None,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.n_sma_threshold = n_sma_threshold
        self.degenerated_to_sgd = degenerated_to_sgd
        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'rectify': rectify,
            'ams_bound': ams_bound,
            'eps': eps,
            'foreach': foreach,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'diffGrad'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['previous_grad'] = torch.zeros_like(p)

                if group['ams_bound']:
                    state['max_exp_avg_sq'] = torch.zeros_like(p)

                if group.get('adanorm'):
                    state['exp_grad_adanorm'] = torch.zeros((1,), dtype=grad.dtype, device=grad.device)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        return not group.get('adanorm') and self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        state_dict: dict[str, list[torch.Tensor]],
        step_size: float | torch.Tensor,
        is_rectified: bool,
        apply_update: bool,
    ) -> None:
        beta1, beta2 = group['betas']
        exp_avgs, exp_avg_sqs = state_dict['exp_avg'], state_dict['exp_avg_sq']
        previous_grads = state_dict['previous_grad']

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=group['lr'],
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - beta1)

        torch._foreach_mul_(exp_avg_sqs, beta2)
        torch._foreach_addcmul_(exp_avg_sqs, grads, grads, value=1.0 - beta2)

        if not group['rectify'] or is_rectified:
            de_noms = self.apply_ams_bound_foreach(
                group['ams_bound'], exp_avg_sqs, state_dict.get('max_exp_avg_sq', []), group['eps']
            )

            torch._foreach_sub_(previous_grads, grads)
            torch._foreach_abs_(previous_grads)
            if isinstance(step_size, torch.Tensor):
                # Inductor cannot fuse the native foreach sigmoid.
                for previous_grad in previous_grads:
                    previous_grad.sigmoid_()
            else:
                torch._foreach_sigmoid_(previous_grads)
            torch._foreach_mul_(previous_grads, exp_avgs)

            foreach_addcdiv_(params, previous_grads, de_noms, value=-step_size)
        else:
            if group['ams_bound']:
                torch._foreach_maximum_(state_dict['max_exp_avg_sq'], exp_avg_sqs)

            if apply_update:
                foreach_add_(params, exp_avgs, alpha=-step_size)

        torch._foreach_copy_(previous_grads, grads)

    def _step_per_param(self, group: ParamGroup, step_size: float, n_sma: float) -> None:
        beta1, beta2 = group['betas']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            self.apply_weight_decay(
                p=p,
                grad=grad,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=group['fixed_decay'],
            )

            state = self.state[p]

            exp_avg, exp_avg_sq, previous_grad = state['exp_avg'], state['exp_avg_sq'], state['previous_grad']

            p, grad, exp_avg, exp_avg_sq, previous_grad = self.view_as_real(
                p, grad, exp_avg, exp_avg_sq, previous_grad
            )

            s_grad = self.get_adanorm_gradient(
                grad=grad,
                adanorm=group.get('adanorm', False),
                exp_grad_norm=state.get('exp_grad_adanorm', None),
                r=group.get('adanorm_r', None),
            )

            exp_avg.lerp_(s_grad, weight=1.0 - beta1)

            exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

            de_nom = self.apply_ams_bound(
                ams_bound=group['ams_bound'],
                exp_avg_sq=exp_avg_sq,
                max_exp_avg_sq=state.get('max_exp_avg_sq', None),
                eps=group['eps'],
            )

            dfc = previous_grad
            dfc.sub_(grad).abs_().sigmoid_().mul_(exp_avg)

            if not group['rectify'] or n_sma >= self.n_sma_threshold:
                p.addcdiv_(dfc, de_nom, value=-step_size)
            elif step_size > 0:
                p.add_(exp_avg, alpha=-step_size)

            state['previous_grad'].copy_(
                torch.view_as_complex(grad) if torch.is_complex(state['previous_grad']) else grad
            )

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])

            step_size, n_sma = self.get_rectify_step_size(
                is_rectify=group['rectify'],
                step=group['step'],
                lr=group['lr'],
                beta2=beta2,
                n_sma_threshold=self.n_sma_threshold,
                degenerated_to_sgd=self.degenerated_to_sgd,
            )

            if not group['rectify']:
                step_size = step_size * math.sqrt(self.debias(beta2, group['step']))

            step_size = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=step_size,
                bias_correction1=bias_correction1,
            )

            if self._can_use_foreach(group):
                state_keys = ['exp_avg', 'exp_avg_sq', 'previous_grad']

                if group['ams_bound']:
                    state_keys.append('max_exp_avg_sq')

                params, grads, state_dict = self.collect_trainable_params(group, self.state, state_keys=state_keys)

                for tensors in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(
                        group,
                        tensors['params'],
                        tensors['grads'],
                        tensors,
                        step_size,
                        n_sma >= self.n_sma_threshold,
                        bool(step_size > 0),
                    )
            else:
                self._step_per_param(group, step_size, n_sma)

        return loss

DistributedMuon

Bases: BaseOptimizer

Distributed momentum updates with Newton-Schulz matrix orthogonalization.

Set use_muon=True for hidden weight matrices and use_muon=False for AdamW groups, such as embeddings, classifier heads, biases, and gains. Pass higher dimensional weights directly. The orthogonal update uses a flattened matrix view. Requires an initialized distributed process group.

Parameters:

Name Type Description Default
params ParamsT

Parameter group dictionaries with a use_muon flag for each group.

required
lr float

Learning rate.

0.02
momentum float

Momentum factor.

0.95
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
nesterov bool

Use Nesterov momentum.

True
ns_steps int

Number of Newton-Schulz iterations.

5
ns_coeffs NewtonSchulzWeights

Newton-Schulz coefficients or preset name.

'original'
use_adjusted_lr bool

Scale orthogonal updates using the Moonlight shape adjustment.

False
adamw_lr float

Learning rate for parameters in the AdamW groups.

0.0003
adamw_betas Betas

Decay rates for the first and second moments in the AdamW groups.

(0.9, 0.95)
adamw_wd float

Weight decay for parameters in the AdamW groups.

0.0
adamw_eps float

Numerical stability constant for the AdamW groups.

1e-10
maximize bool

Maximize the objective instead of minimizing it.

False

Examples:

from pytorch_optimizer import DistributedMuon

hidden_weights = [p for p in model.body.parameters() if p.ndim >= 2]
hidden_gains_biases = [p for p in model.body.parameters() if p.ndim < 2]
non_hidden_params = [*model.head.parameters(), *model.embed.parameters()]

param_groups = [
    dict(params=hidden_weights, lr=0.02, weight_decay=0.01, use_muon=True),
    dict(
        params=hidden_gains_biases + non_hidden_params,
        lr=3e-4,
        betas=(0.9, 0.95),
        weight_decay=0.01,
        use_muon=False,
    ),
]

optimizer = DistributedMuon(param_groups)
Source code in pytorch_optimizer/optimizer/muon.py
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
class DistributedMuon(BaseOptimizer):  # pragma: no cover
    """Distributed momentum updates with Newton-Schulz matrix orthogonalization.

    Set `use_muon=True` for hidden weight matrices and `use_muon=False` for AdamW groups,
    such as embeddings, classifier heads, biases, and gains. Pass higher dimensional
    weights directly. The orthogonal update uses a flattened matrix view.
    Requires an initialized distributed process group.

    Args:
        params: Parameter group dictionaries with a `use_muon` flag for each group.
        lr: Learning rate.
        momentum: Momentum factor.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        nesterov: Use Nesterov momentum.
        ns_steps: Number of Newton-Schulz iterations.
        ns_coeffs: Newton-Schulz coefficients or preset name.
        use_adjusted_lr: Scale orthogonal updates using the Moonlight shape adjustment.
        adamw_lr: Learning rate for parameters in the AdamW groups.
        adamw_betas: Decay rates for the first and second moments in the AdamW groups.
        adamw_wd: Weight decay for parameters in the AdamW groups.
        adamw_eps: Numerical stability constant for the AdamW groups.
        maximize: Maximize the objective instead of minimizing it.

    Examples:
        ```python
        from pytorch_optimizer import DistributedMuon

        hidden_weights = [p for p in model.body.parameters() if p.ndim >= 2]
        hidden_gains_biases = [p for p in model.body.parameters() if p.ndim < 2]
        non_hidden_params = [*model.head.parameters(), *model.embed.parameters()]

        param_groups = [
            dict(params=hidden_weights, lr=0.02, weight_decay=0.01, use_muon=True),
            dict(
                params=hidden_gains_biases + non_hidden_params,
                lr=3e-4,
                betas=(0.9, 0.95),
                weight_decay=0.01,
                use_muon=False,
            ),
        ]

        optimizer = DistributedMuon(param_groups)
        ```

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 2e-2,
        momentum: float = 0.95,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        nesterov: bool = True,
        ns_steps: int = 5,
        ns_coeffs: NewtonSchulzWeights = 'original',
        use_adjusted_lr: bool = False,
        adamw_lr: float = 3e-4,
        adamw_betas: Betas = (0.9, 0.95),
        adamw_wd: float = 0.0,
        adamw_eps: float = 1e-10,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_learning_rate(adamw_lr)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(momentum, 'momentum', 0.0, 1.0, range_type='[)')
        self.validate_positive(ns_steps, 'ns_steps')
        self.validate_betas(adamw_betas)
        self.validate_non_negative(adamw_wd, 'adamw_wd')
        self.validate_non_negative(adamw_eps, 'adamw_eps')
        ns_coeffs = get_newton_schulz_weights(ns_coeffs)

        self.maximize = maximize

        self.world_size: int = get_world_size()
        self.rank: int = get_rank()

        for group in params:
            group = cast(ParamGroup, group)
            if 'use_muon' not in group:
                raise ValueError('`use_muon` must be set.')

            if group['use_muon']:
                group['lr'] = group.get('lr', lr)
                group['momentum'] = group.get('momentum', momentum)
                group['nesterov'] = group.get('nesterov', nesterov)
                group['weight_decay'] = group.get('weight_decay', weight_decay)
                group['ns_steps'] = group.get('ns_steps', ns_steps)
                group['ns_coeffs'] = get_newton_schulz_weights(group.get('ns_coeffs', ns_coeffs))
                group['use_adjusted_lr'] = group.get('use_adjusted_lr', use_adjusted_lr)
            else:
                group['lr'] = group.get('lr', adamw_lr)
                group['betas'] = group.get('betas', adamw_betas)
                group['eps'] = group.get('eps', adamw_eps)
                group['weight_decay'] = group.get('weight_decay', adamw_wd)

            group['weight_decouple'] = group.get('weight_decouple', weight_decouple)

        super().__init__(params, kwargs)

    def __str__(self) -> str:
        return 'DistributedMuon'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                p.grad = torch.zeros_like(p)

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0 and not group['use_muon']:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if group['use_muon']:
                params = group['params']
                padded_params = params + [torch.empty_like(params[-1])] * (
                    self.world_size - len(params) % self.world_size
                )

                for i in range(len(params))[:: self.world_size]:
                    if i + self.rank < len(params):
                        p = params[i + self.rank]

                        grad = p.grad

                        self.maximize_gradient(grad, maximize=self.maximize)

                        state = self.state[p]
                        if len(state) == 0:
                            state['momentum_buffer'] = torch.zeros_like(p)

                        self.apply_weight_decay(
                            p,
                            grad=grad,
                            lr=group['lr'],
                            weight_decay=group['weight_decay'],
                            weight_decouple=group['weight_decouple'],
                            fixed_decay=False,
                        )

                        buf = state['momentum_buffer']
                        buf.lerp_(grad, weight=1.0 - group['momentum'])

                        update = grad.lerp_(buf, weight=group['momentum']) if group['nesterov'] else buf
                        if update.ndim > 2:
                            update = update.view(len(update), -1)

                        update = zero_power_via_newton_schulz_5(
                            update, num_steps=group['ns_steps'], weights=group['ns_coeffs']
                        )

                        if group.get('cautious'):
                            self.apply_cautious(update, grad)

                        lr = get_adjusted_lr(group['lr'], p.size(), use_adjusted_lr=group['use_adjusted_lr'])

                        p.add_(update.reshape(p.shape), alpha=-lr)

                    all_gather(padded_params[i:i + self.world_size], padded_params[i + self.rank])  # fmt: skip
            else:
                for p in group['params']:
                    grad = p.grad

                    self.maximize_gradient(grad, maximize=self.maximize)

                    self.apply_weight_decay(
                        p,
                        grad=grad,
                        lr=group['lr'],
                        weight_decay=group['weight_decay'],
                        weight_decouple=group['weight_decouple'],
                        fixed_decay=False,
                    )

                    state = self.state[p]
                    exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                    beta1, beta2 = group['betas']

                    bias_correction1: float = self.debias(beta1, group['step'])
                    bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

                    exp_avg.lerp_(grad, weight=1.0 - beta1)
                    exp_avg_sq.lerp_(grad.square(), weight=1.0 - beta2)

                    de_nom = exp_avg_sq.sqrt().div_(bias_correction2_sq).add_(group['eps'])

                    p.addcdiv_(exp_avg / bias_correction1, de_nom, value=-group['lr'])

        return loss

DualAdam

Bases: BaseOptimizer

Adam with a decaying inverse Adam update contribution.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Coefficients used for computing running averages of gradient and squared gradient.

(0.9, 0.999)
switch_rate float

Linear decay rate for the inverse Adam update contribution.

0.01
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/dual_adam.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
class DualAdam(BaseOptimizer):
    """Adam with a decaying inverse Adam update contribution.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Coefficients used for computing running averages of gradient and squared gradient.
        switch_rate: Linear decay rate for the inverse Adam update contribution.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        switch_rate: float = 1e-2,
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_range(switch_rate, 'switch_rate', 0.0, 1.0, range_type='[]')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'switch_rate': switch_rate,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'DualAdam'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2: float = self.debias(beta2, group['step'])

            inverse_adam_rate: float = max(0.0, 1.0 - group['step'] * group['switch_rate'])
            use_inverse_adam: bool = inverse_adam_rate >= group['switch_rate']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                exp_avg_hat = exp_avg.div(bias_correction1)
                de_nom = exp_avg_sq.div(bias_correction2).sqrt_().add_(group['eps'])

                if use_inverse_adam:
                    update = de_nom.reciprocal().lerp_(de_nom, weight=inverse_adam_rate)

                    p.addcmul_(exp_avg_hat, update, value=-group['lr'])
                else:
                    p.addcdiv_(exp_avg_hat, de_nom, value=-group['lr'])

        return loss

DynamicLossScaler

Adjust the loss scale in response to low precision gradient overflow.

Increase the scale after an overflow free window and decrease it when the overflow fraction reaches the tolerance.

References

Parameters:

Name Type Description Default
init_scale float

Initial loss scale.

2.0 ** 15
scale_factor float

Multiplier for increasing or decreasing the scale.

2.0
scale_window int

Number of overflow free iterations between scale increases.

2000
tolerance float

Fraction of overflowing iterations that triggers a scale decrease.

0.0
threshold float | None

Optional lower bound for the scale.

None
Source code in pytorch_optimizer/optimizer/fp16.py
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
class DynamicLossScaler:
    """Adjust the loss scale in response to low precision gradient overflow.

    Increase the scale after an overflow free window and decrease it when the overflow
    fraction reaches the tolerance.

    References:
        - https://docs.nvidia.com/deeplearning/performance/mixed-precision-training/index.html#lossscaling
        - https://github.com/pytorch/fairseq/blob/main/fairseq/optim/fp16_optimizer.py
        - https://github.com/facebookresearch/ParlAI/blob/main/parlai/utils/fp16.py

    Args:
        init_scale: Initial loss scale.
        scale_factor: Multiplier for increasing or decreasing the scale.
        scale_window: Number of overflow free iterations between scale increases.
        tolerance: Fraction of overflowing iterations that triggers a scale decrease.
        threshold: Optional lower bound for the scale.

    """

    def __init__(
        self,
        init_scale: float = 2.0 ** 15,
        scale_factor: float = 2.0,
        scale_window: int = 2000,
        tolerance: float = 0.00,
        threshold: float | None = None,
    ):  # fmt: skip
        self.loss_scale = init_scale
        self.scale_factor = scale_factor
        self.scale_window = scale_window
        self.tolerance = tolerance
        self.threshold = threshold

        self.iter: int = 0
        self.last_overflow_iter: int = -1
        self.last_rescale_iter: int = -1
        self.overflows_since_rescale: int = 0
        self.has_overflow_serial: bool = False

    def update_scale(self, overflow: bool):
        """Update the loss scale after checking the current gradients.

        Args:
            overflow: Whether the current gradients contain NaN or infinite values.

        """
        iter_since_rescale: int = self.iter - self.last_rescale_iter

        if overflow:
            # calculate how often we overflowed already
            self.last_overflow_iter = self.iter
            self.overflows_since_rescale += 1

            pct_overflow: float = self.overflows_since_rescale / float(iter_since_rescale)
            if pct_overflow >= self.tolerance:
                # decrease loss scale by the scale factor
                self.decrease_loss_scale()

                # reset iterations
                self.last_rescale_iter = self.iter
                self.overflows_since_rescale = 0
        elif (self.iter - self.last_overflow_iter) % self.scale_window == 0:
            # increase the loss scale by scale factor
            self.loss_scale *= self.scale_factor
            self.last_rescale_iter = self.iter

        self.iter += 1

    def decrease_loss_scale(self):
        """Divide the loss scale by `scale_factor`, respecting the optional lower bound."""
        self.loss_scale /= self.scale_factor
        if self.threshold is not None:
            self.loss_scale = max(self.loss_scale, self.threshold)

decrease_loss_scale()

Divide the loss scale by scale_factor, respecting the optional lower bound.

Source code in pytorch_optimizer/optimizer/fp16.py
81
82
83
84
85
def decrease_loss_scale(self):
    """Divide the loss scale by `scale_factor`, respecting the optional lower bound."""
    self.loss_scale /= self.scale_factor
    if self.threshold is not None:
        self.loss_scale = max(self.loss_scale, self.threshold)

update_scale(overflow)

Update the loss scale after checking the current gradients.

Parameters:

Name Type Description Default
overflow bool

Whether the current gradients contain NaN or infinite values.

required
Source code in pytorch_optimizer/optimizer/fp16.py
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
def update_scale(self, overflow: bool):
    """Update the loss scale after checking the current gradients.

    Args:
        overflow: Whether the current gradients contain NaN or infinite values.

    """
    iter_since_rescale: int = self.iter - self.last_rescale_iter

    if overflow:
        # calculate how often we overflowed already
        self.last_overflow_iter = self.iter
        self.overflows_since_rescale += 1

        pct_overflow: float = self.overflows_since_rescale / float(iter_since_rescale)
        if pct_overflow >= self.tolerance:
            # decrease loss scale by the scale factor
            self.decrease_loss_scale()

            # reset iterations
            self.last_rescale_iter = self.iter
            self.overflows_since_rescale = 0
    elif (self.iter - self.last_overflow_iter) % self.scale_window == 0:
        # increase the loss scale by scale factor
        self.loss_scale *= self.scale_factor
        self.last_rescale_iter = self.iter

    self.iter += 1

EmoFact

Bases: BaseOptimizer

Factored adaptive updates with EmoNavi loss driven scaling.

Supply a loss closure to step() to enable loss driven scaling.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for row/column gradient RMS averages and the vector second moment.

(0.9, 0.999)
use_shadow bool

Blend parameters with a running shadow copy based on loss trends.

False
shadow_weight float

Interpolation weight for shadow copy updates during a shadow correction.

0.05
weight_decay float

Weight decay coefficient.

0.01
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/emonavi.py
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
class EmoFact(BaseOptimizer):
    """Factored adaptive updates with EmoNavi loss driven scaling.

    Supply a loss closure to `step()` to enable loss driven scaling.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for row/column gradient RMS averages and the vector second moment.
        use_shadow: Blend parameters with a running shadow copy based on loss trends.
        shadow_weight: Interpolation weight for shadow copy updates during a shadow correction.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        use_shadow: bool = False,
        shadow_weight: float = 0.05,
        weight_decay: float = 1e-2,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_range(shadow_weight, 'shadow_weight', 0.0, 1.0)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        self.lr = lr

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'use_shadow': use_shadow,
            'shadow_weight': shadow_weight,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'EmoFact'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                if group['use_shadow']:
                    state['shadow'] = p.clone()

                shape = p.size()

                if len(shape) >= 2:
                    r_shape = [shape[0]] + [1] * (len(shape) - 1)
                    state['exp_avg_r'] = torch.zeros(r_shape, dtype=p.dtype, device=p.device)

                    c_shape = [1, *list(shape[1:])]
                    state['exp_avg_c'] = torch.zeros(c_shape, dtype=p.dtype, device=p.device)
                else:
                    state['exp_avg'] = torch.zeros_like(p)
                    state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss = 0.0
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            emo_drive, ratio, trust = get_emo_drive(self.state, loss, group['use_shadow'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                if group['use_shadow']:
                    shadow = state['shadow']
                    if ratio > 0.0:
                        p.mul_(1.0 - ratio).add_(shadow, alpha=abs(trust))
                    else:
                        leap_ratio = 0.1 * abs(trust)
                        shadow.lerp_(p, weight=leap_ratio)

                if grad.dim() >= 2:
                    exp_avg_r, exp_avg_c = state['exp_avg_r'], state['exp_avg_c']

                    grad_p2 = grad.pow(2)
                    r_sq = (
                        torch.mean(grad_p2, dim=tuple(range(1, grad.dim())), keepdim=True).add_(group['eps']).sqrt_()
                    )
                    c_sq = torch.mean(grad_p2, dim=0, keepdim=True).add_(group['eps']).sqrt_()

                    exp_avg_r.lerp_(r_sq, weight=1.0 - beta1)
                    exp_avg_c.lerp_(c_sq, weight=1.0 - beta1)

                    de_nom = (exp_avg_r * exp_avg_c).sqrt_().add_(group['eps'])

                    update = grad / de_nom
                else:
                    exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                    exp_avg.lerp_(grad, weight=1.0 - beta1)
                    exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                    de_nom = exp_avg_sq.sqrt().add_(group['eps'])

                    update = exp_avg / de_nom

                update.sign_()

                p.add_(update, alpha=-group['lr'] * emo_drive)

        self.prev_loss = loss

        return loss

EmoLynx

Bases: BaseOptimizer

Sign based momentum updates with EmoNavi loss driven scaling.

Supply a loss closure to step() to enable loss driven scaling.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for update interpolation and gradient momentum.

(0.9, 0.99)
use_shadow bool

Blend parameters with a running shadow copy based on loss trends.

False
shadow_weight float

Interpolation weight for shadow copy updates during a shadow correction.

0.05
weight_decay float

Weight decay coefficient.

0.01
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/emonavi.py
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
class EmoLynx(BaseOptimizer):
    """Sign based momentum updates with EmoNavi loss driven scaling.

    Supply a loss closure to `step()` to enable loss driven scaling.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for update interpolation and gradient momentum.
        use_shadow: Blend parameters with a running shadow copy based on loss trends.
        shadow_weight: Interpolation weight for shadow copy updates during a shadow correction.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.99),
        use_shadow: bool = False,
        shadow_weight: float = 0.05,
        weight_decay: float = 1e-2,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_range(shadow_weight, 'shadow_weight', 0.0, 1.0)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'use_shadow': use_shadow,
            'shadow_weight': shadow_weight,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'EmoLynx'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                if group['use_shadow']:
                    state['shadow'] = p.clone()
                state['exp_avg'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss = 0.0
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            emo_drive, ratio, trust = get_emo_drive(self.state, loss, group['use_shadow'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                if group['use_shadow']:
                    shadow = state['shadow']
                    if ratio > 0.0:
                        p.mul_(1.0 - ratio).add_(shadow, alpha=abs(trust))
                    else:
                        leap_ratio = 0.1 * abs(trust)
                        shadow.lerp_(p, weight=leap_ratio)

                exp_avg = state['exp_avg']

                blended_grad = grad.mul(1.0 - beta1).add_(exp_avg, alpha=beta1).sign_()
                exp_avg.lerp_(grad, weight=1.0 - beta2)

                p.add_(blended_grad, alpha=-group['lr'] * emo_drive)

        return loss

EmoNavi

Bases: BaseOptimizer

Adam style updates with loss driven momentum scaling and optional shadow weights.

Supply a loss closure to step() to enable loss driven scaling.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
use_shadow bool

Blend parameters with a running shadow copy based on loss trends.

False
shadow_weight float

Interpolation weight for shadow copy updates during a shadow correction.

0.05
weight_decay float

Weight decay coefficient.

0.01
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/emonavi.py
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
class EmoNavi(BaseOptimizer):
    """Adam style updates with loss driven momentum scaling and optional shadow weights.

    Supply a loss closure to `step()` to enable loss driven scaling.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        use_shadow: Blend parameters with a running shadow copy based on loss trends.
        shadow_weight: Interpolation weight for shadow copy updates during a shadow correction.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        use_shadow: bool = False,
        shadow_weight: float = 0.05,
        weight_decay: float = 1e-2,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_range(shadow_weight, 'shadow_weight', 0.0, 1.0)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.use_shadow = use_shadow
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'use_shadow': use_shadow,
            'shadow_weight': shadow_weight,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'EmoNavi'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

                if group['use_shadow']:
                    state['shadow'] = p.clone()

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss = 0.0
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            emo_drive, ratio, trust = get_emo_drive(self.state, loss, group['use_shadow'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                if group['use_shadow']:
                    shadow = state['shadow']

                    if ratio > 0.0:
                        p.mul_(1.0 - ratio).add_(state['shadow'], alpha=abs(trust))
                        shadow.lerp_(p, weight=group['shadow_weight'])
                    else:
                        leap_ratio: float = 0.1 * abs(trust)
                        shadow.lerp_(p, weight=leap_ratio)

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                de_nom = exp_avg_sq.sqrt().add_(group['eps'])

                p.addcdiv_(exp_avg, de_nom, value=-group['lr'] * emo_drive)

        return loss

EXAdam

Bases: BaseOptimizer

Adam with adaptive cross moment corrections.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/exadam.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
class EXAdam(BaseOptimizer):
    """Adam with adaptive cross moment corrections.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'EXAdam'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2: float = self.debias(beta2, group['step'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                d1 = 1.0 + exp_avg_sq.div(exp_avg_sq.add(group['eps'])) * (1.0 - bias_correction2)

                exp_avg_p2 = exp_avg.pow(2)
                d2 = 1.0 + exp_avg_p2.div(exp_avg_p2.add(group['eps'])) * (1.0 - bias_correction1)

                m_tilde = exp_avg.div(bias_correction1) * d1
                v_tilde = exp_avg_sq.div(bias_correction2) * d2

                g_tilde = grad.div(bias_correction1) * d1

                update = (m_tilde + g_tilde) / v_tilde.sqrt().add_(group['eps'])

                p.add_(update, alpha=-group['lr'])

        return loss

FAdam

Bases: BaseOptimizer

Natural gradient Adam using diagonal empirical Fisher information.

The adaptive stabilizer is min(eps, eps_2 * RMS(grad)) ** (2 * p). Checkpoint loading preserves the saved momentum and Fisher state dtypes.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for natural gradient momentum and diagonal empirical Fisher estimates.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.1
clip float

RMS cap for natural gradients and preconditioned weight decay.

1.0
p float

Exponent applied to the Fisher information diagonal.

0.5
eps float

Upper bound on the adaptive epsilon before applying the exponent.

1e-08
momentum_dtype dtype

Dtype of momentum.

float32
fim_dtype dtype

Data type of the Fisher information diagonal.

float32
maximize bool

Maximize the objective instead of minimizing it.

False
eps_2 float

Gradient RMS multiplier for the adaptive epsilon.

0.01
Source code in pytorch_optimizer/optimizer/fadam.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
class FAdam(BaseOptimizer):
    """Natural gradient Adam using diagonal empirical Fisher information.

    The adaptive stabilizer is `min(eps, eps_2 * RMS(grad)) ** (2 * p)`.
    Checkpoint loading preserves the saved momentum and Fisher state dtypes.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for natural gradient momentum and diagonal empirical Fisher estimates.
        weight_decay: Weight decay coefficient.
        clip: RMS cap for natural gradients and preconditioned weight decay.
        p: Exponent applied to the Fisher information diagonal.
        eps: Upper bound on the adaptive epsilon before applying the exponent.
        momentum_dtype: Dtype of momentum.
        fim_dtype: Data type of the Fisher information diagonal.
        maximize: Maximize the objective instead of minimizing it.
        eps_2: Gradient RMS multiplier for the adaptive epsilon.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.1,
        clip: float = 1.0,
        p: float = 0.5,
        eps: float = 1e-8,
        momentum_dtype: torch.dtype = torch.float32,
        fim_dtype: torch.dtype = torch.float32,
        maximize: bool = False,
        eps_2: float = 0.01,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_positive(clip, 'clip')
        self.validate_positive(p, 'p')
        self.validate_non_negative(eps, 'eps')
        self.validate_non_negative(eps_2, 'eps_2')

        self.momentum_dtype = momentum_dtype
        self.fim_dtype = fim_dtype
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'clip': clip,
            'p': p,
            'eps': eps,
            'eps_2': eps_2,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'FAdam'

    def _restore_state_types(self, value, saved_value):
        if isinstance(saved_value, torch.Tensor) and saved_value.is_floating_point():
            return saved_value.to(device=value.device)

        return super()._restore_state_types(value, saved_value)

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['momentum'] = torch.zeros_like(p, dtype=self.momentum_dtype)
                state['fim'] = torch.zeros_like(p, dtype=self.fim_dtype)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            curr_beta2: float = self.debias_beta(beta2, group['step'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                momentum, fim = state['momentum'], state['fim']

                fim.mul_(curr_beta2).addcmul_(grad, grad, value=1.0 - curr_beta2)

                rms_grad = grad.pow(2).mean().sqrt_()
                curr_eps = min(group['eps'], group['eps_2'] * rms_grad) if rms_grad > 0 else group['eps']

                fim_base = fim.pow(group['p']).add_(curr_eps ** (2.0 * group['p']))
                grad_nat = grad / fim_base

                rms = grad_nat.pow(2).mean().sqrt_()
                divisor = max(1, rms / group['clip'])
                grad_nat.div_(divisor)

                momentum.lerp_(grad_nat, weight=1.0 - beta1)

                grad_weights = p / fim_base

                rms = torch.pow(grad_weights, 2).mean().sqrt_()
                divisor = max(1, rms / group['clip'])
                grad_weights.div_(divisor)

                grad_weights.mul_(group['weight_decay']).add_(momentum)

                p.add_(grad_weights, alpha=-group['lr'])

        return loss

Fira

Bases: BaseOptimizer

AdamW with low rank updates and full rank gradient compensation.

Add rank, update_proj_gap, scale, and projection_type to parameter groups containing matrix weights to enable projection.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
eps float

Term added to the denominator to improve numerical stability.

1e-06
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/fira.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
class Fira(BaseOptimizer):
    """AdamW with low rank updates and full rank gradient compensation.

    Add `rank`, `update_proj_gap`, `scale`, and `projection_type` to parameter groups
    containing matrix weights to enable projection.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        eps: float = 1e-6,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Fira'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size: float = group['lr'] * bias_correction2_sq / bias_correction1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                full_grad = grad

                if 'rank' in group and p.dim() == 2:
                    if 'projector' not in state:
                        state['projector'] = GaLoreProjector(
                            rank=group['rank'],
                            update_proj_gap=group['update_proj_gap'],
                            scale=group['scale'],
                            projection_type=group['projection_type'],
                        )

                    grad = state['projector'].project(grad, group['step'])

                if 'exp_avg' not in state:
                    state['exp_avg'] = torch.zeros_like(grad)
                    state['exp_avg_sq'] = torch.zeros_like(grad)

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']
                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                de_nom = exp_avg_sq.sqrt().add_(group['eps'])

                norm_grad = exp_avg / de_nom

                if 'rank' in group and p.dim() == 2:
                    sub_grad = state['projector'].project_back(grad)

                    norm_dim: int = 0 if norm_grad.shape[0] < norm_grad.shape[1] else 1

                    scaling_factor = torch.norm(norm_grad, dim=norm_dim) / (torch.norm(grad, dim=norm_dim) + 1e-8)
                    if norm_dim == 1:
                        scaling_factor = scaling_factor.unsqueeze(1)

                    scaling_grad = full_grad.sub(sub_grad).mul_(scaling_factor)

                    if 'scaling_grad' in state:
                        scaling_grad_norm = torch.norm(scaling_grad)

                        limiter = max(scaling_grad_norm / (state['scaling_grad'] + 1e-8), 1.01) / 1.01
                        scaling_grad.div_(limiter)

                        state['scaling_grad'] = scaling_grad_norm / limiter
                    else:
                        state['scaling_grad'] = torch.norm(scaling_grad)

                    norm_grad = state['projector'].project_back(norm_grad).add_(scaling_grad)

                p.add_(norm_grad, alpha=-step_size)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=True,
                    fixed_decay=False,
                )

        return loss

FlashAdamW

Bases: BaseOptimizer

AdamW with grouped 8-bit optimizer states and optional master weight error correction.

Supports compressed checkpoints and low precision parameters through a portable PyTorch implementation of FlashOptim style updates. Float64 parameters retain their precision, with float64 arithmetic for unquantized moments.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Coefficients used for computing running averages of gradient and squared gradient.

(0.9, 0.999)
eps float

Term added to the denominator to improve numerical stability.

1e-08
weight_decay float

Weight decay coefficient.

0.01
decouple_lr bool

Scale weight decay by lr / initial_lr instead of lr. Requires a positive initial_lr when applying nonzero weight decay at a positive learning rate.

False
quantize bool

Store Adam moments as grouped 8-bit values plus fp16 scales.

True
compress_state_dict bool

Save quantized states in checkpoints when quantize is enabled.

True
master_weight_bits int | None

Effective master weight precision for bf16/fp16 parameters. Supports None, 24, and 32.

None
check_numerics bool

Raise if low precision parameter updates are unlikely to alter the master weight.

False
fused bool

Placeholder for FlashOptim's Triton fused path. Currently unsupported in this portable backend.

False
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/flash_adamw.py
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
class FlashAdamW(BaseOptimizer):
    """AdamW with grouped 8-bit optimizer states and optional master weight error correction.

    Supports compressed checkpoints and low precision parameters through a portable
    PyTorch implementation of FlashOptim style updates.
    Float64 parameters retain their precision, with float64 arithmetic for unquantized moments.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Coefficients used for computing running averages of gradient and squared gradient.
        eps: Term added to the denominator to improve numerical stability.
        weight_decay: Weight decay coefficient.
        decouple_lr: Scale weight decay by `lr / initial_lr` instead of `lr`. Requires a positive
            `initial_lr` when applying nonzero weight decay at a positive learning rate.
        quantize: Store Adam moments as grouped 8-bit values plus fp16 scales.
        compress_state_dict: Save quantized states in checkpoints when `quantize` is enabled.
        master_weight_bits: Effective master weight precision for bf16/fp16 parameters. Supports `None`, `24`, and
            `32`.
        check_numerics: Raise if low precision parameter updates are unlikely to alter the master weight.
        fused: Placeholder for FlashOptim's Triton fused path. Currently unsupported in this portable backend.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 1e-2,
        decouple_lr: bool = False,
        quantize: bool = True,
        compress_state_dict: bool = True,
        master_weight_bits: int | None = None,
        check_numerics: bool = False,
        fused: bool = False,
        maximize: bool = False,
        eps: float = 1e-8,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(eps, 'eps')
        self.validate_non_negative(weight_decay, 'weight_decay')

        if master_weight_bits not in VALID_MASTER_WEIGHT_BITS:
            raise ValueError(f'master_weight_bits must be one of {VALID_MASTER_WEIGHT_BITS}')

        if fused:
            raise NotImplementedError('FlashAdamW fused Triton kernels are not available in this portable backend')

        self.maximize = maximize
        self.compress_state_dict = compress_state_dict
        self.check_numerics = check_numerics
        self.master_byte_width = BITS_TO_BYTES[master_weight_bits]
        self.param_absmax: dict[int, float] = {}

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'eps': eps,
            'weight_decay': weight_decay,
            'decouple_lr': decouple_lr,
            'quantize': quantize,
            'master_byte_width': self.master_byte_width,
            **kwargs,
        }

        super().__init__(params, defaults)

        for group in self.param_groups:
            group.setdefault('initial_lr', group['lr'])

        if master_weight_bits is not None and all(
            p.dtype == torch.float32 for group in self.param_groups for p in group['params']
        ):
            raise ValueError('master_weight_bits has no effect when all parameters are fp32')

    def __str__(self) -> str:
        return 'FlashAdamW'

    def maybe_check_numerics(self, p: torch.Tensor, lr: float, master_byte_width: int) -> None:
        if not self.check_numerics or p.dtype == torch.float32 or lr == 0.0:
            return

        max_abs = self.param_absmax.get(id(p))
        if max_abs is None:
            self.param_absmax[id(p)] = max_abs = float(p.detach().abs().max().item()) if p.numel() > 0 else 0.0

        if max_abs <= 0.0 or not math.isfinite(max_abs):
            return

        bits: int = max(DTYPE_WIDTHS[p.dtype], master_byte_width) * 8
        resolution: float = max_abs * 2.0 ** (-(bits - 1))

        if lr * 0.1 < resolution:
            raise ArithmeticError('learning rate is too small to update low-precision FlashAdamW parameters')

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        group.setdefault('initial_lr', group.get('lr'))

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]
            if 'exp_avg' not in state and _quantized_key('exp_avg') not in state:
                store_state(state, 'exp_avg', torch.zeros_like(p, dtype=torch.float32), group['quantize'], p.dtype)
                store_state(state, 'exp_avg_sq', torch.zeros_like(p, dtype=torch.float32), group['quantize'], p.dtype)

            error_bytes = group['master_byte_width'] - DTYPE_WIDTHS[p.dtype]
            if error_bytes > 0 and 'error_bits' not in state:
                error_dtype = torch.int8 if error_bytes == 1 else torch.int16
                state['error_bits'] = torch.zeros_like(p, dtype=error_dtype)

    @staticmethod
    def get_param_fp32(p: torch.Tensor, state: dict[str, Any]) -> torch.Tensor:
        return reconstruct_fp32_param(p, state['error_bits']) if 'error_bits' in state else p.to(torch.float32)

    @staticmethod
    def set_param_fp32(p: torch.Tensor, state: dict[str, Any], value: torch.Tensor, master_byte_width: int) -> None:
        p.copy_(value.to(p.dtype))
        if 'error_bits' in state:
            state['error_bits'].copy_(compute_ecc_bits(value, p, master_byte_width))

    def recompute_param_stats(self) -> None:
        for group in self.param_groups:
            for p in group['params']:
                self.param_absmax[id(p)] = float(p.detach().abs().max().item()) if p.numel() > 0 else 0.0

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2: float = self.debias(beta2, group['step'])

            weight_decay = group['weight_decay']
            if group['decouple_lr'] and weight_decay > 0.0 and group['lr'] > 0.0:
                self.validate_positive(group['initial_lr'], 'initial_lr')
                weight_decay /= group['initial_lr']

            for p in group['params']:
                if p.grad is None:
                    continue

                state = self.state[p]

                grad = p.grad.to(dtype=torch.float64 if p.dtype == torch.float64 else torch.float32)

                self.maximize_gradient(grad, maximize=self.maximize)

                self.maybe_check_numerics(p, group['lr'], group['master_byte_width'])

                exp_avg = materialize_state(state, 'exp_avg').to(dtype=grad.dtype)
                exp_avg_sq = materialize_state(state, 'exp_avg_sq').to(dtype=grad.dtype)

                param = p if p.dtype == torch.float64 else self.get_param_fp32(p, state)

                self.apply_weight_decay(
                    param,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=weight_decay,
                    weight_decouple=True,
                    fixed_decay=False,
                )

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                denominator = exp_avg_sq.div(bias_correction2).sqrt_().add_(group['eps'])
                param.addcdiv_(exp_avg.div(bias_correction1), denominator, value=-group['lr'])

                self.set_param_fp32(p, state, param, group['master_byte_width'])
                store_state(state, 'exp_avg', exp_avg, group['quantize'], p.dtype)
                store_state(state, 'exp_avg_sq', exp_avg_sq, group['quantize'], p.dtype)

        return loss

    def state_dict(self) -> dict[str, Any]:
        state_dict = super().state_dict()
        if self.compress_state_dict:
            return state_dict

        state_dict['state'] = {param_id: dict(param_state) for param_id, param_state in state_dict['state'].items()}
        for param_state in state_dict['state'].values():
            for name in ('exp_avg', 'exp_avg_sq'):
                q_key, s_key = _quantized_key(name), _scales_key(name)
                if q_key not in param_state:
                    continue
                param_state[name] = dequantize_state(
                    param_state.pop(q_key), param_state.pop(s_key), *_state_spec(name)
                )

        return state_dict

    def load_state_dict(self, state_dict: dict[str, Any]) -> None:
        super().load_state_dict(state_dict)

        for group in self.param_groups:
            group.setdefault('initial_lr', group['lr'])

        for group, saved_group in zip(self.param_groups, state_dict['param_groups']):
            for p, saved_id in zip(group['params'], saved_group['params']):
                state = self.state[p]
                if not state:
                    continue

                saved_state = state_dict['state'][saved_id]
                for key, value in saved_state.items():
                    if isinstance(value, torch.Tensor):
                        state[key] = value.to(device=p.device)

                for name in ('exp_avg', 'exp_avg_sq'):
                    if group['quantize'] and name in state:
                        store_state(state, name, state.pop(name).to(torch.float32), True, p.dtype)
                    elif group['quantize'] and _quantized_key(name) in state:
                        signed, _, _ = _state_spec(name)
                        quantized_dtype = torch.int8 if signed else torch.uint8
                        state[_quantized_key(name)] = state[_quantized_key(name)].to(quantized_dtype)
                        state[_scales_key(name)] = state[_scales_key(name)].to(torch.float16)
                    elif not group['quantize'] and _quantized_key(name) in state:
                        state[name] = dequantize_state(
                            state.pop(_quantized_key(name)), state.pop(_scales_key(name)), *_state_spec(name)
                        ).to(dtype=p.dtype)

    def get_fp32_model_state_dict(self, model: nn.Module) -> dict[str, torch.Tensor]:
        return {
            name: self.get_param_fp32(param.detach(), self.state.get(param, {})).detach().clone()
            for name, param in model.named_parameters()
        }

    @torch.inference_mode()
    def set_fp32_model_state_dict(self, model: nn.Module, state_dict: dict[str, torch.Tensor]) -> None:
        for name, param in model.named_parameters():
            if name not in state_dict:
                continue

            state = self.state[param]

            master_byte_width = next(
                group['master_byte_width']
                for group in self.param_groups
                if any(param is grouped_param for grouped_param in group['params'])
            )
            error_bytes = master_byte_width - DTYPE_WIDTHS[param.dtype]
            if error_bytes > 0 and 'error_bits' not in state:
                error_dtype = torch.int8 if error_bytes == 1 else torch.int16
                state['error_bits'] = torch.zeros_like(param, dtype=error_dtype)

            self.set_param_fp32(param, state, state_dict[name].to(torch.float32), master_byte_width)

FOCUS

Bases: BaseOptimizer

Sign based updates with attraction toward the running parameter average.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.01
betas Betas

Decay rates for gradient momentum and the running parameter average.

(0.9, 0.999)
gamma float

Controls the strength of the attraction.

0.1
weight_decay float

Weight decay coefficient.

0.0
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/focus.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
class FOCUS(BaseOptimizer):
    """Sign based updates with attraction toward the running parameter average.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for gradient momentum and the running parameter average.
        gamma: Controls the strength of the attraction.
        weight_decay: Weight decay coefficient.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-2,
        betas: Betas = (0.9, 0.999),
        gamma: float = 0.1,
        weight_decay: float = 0.0,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_range(gamma, 'gamma', 0.0, 1.0, '[)')
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.maximize = maximize

        defaults: Defaults = {'lr': lr, 'betas': betas, 'gamma': gamma, 'weight_decay': weight_decay}

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'FOCUS'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['pbar'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction2: float = self.debias(beta2, group['step'])

            weight_decay: float = group['weight_decay']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, pbar = state['exp_avg'], state['pbar']

                p, grad, exp_avg, pbar = self.view_as_real(p, grad, exp_avg, pbar)

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                pbar.lerp_(p, weight=1.0 - beta2)

                pbar_hat = pbar / bias_correction2

                if weight_decay > 0.0:
                    p.add_(pbar_hat, alpha=-group['lr'] * weight_decay)

                update = (p - pbar_hat).sign_().mul_(group['gamma']).add_(torch.sign(exp_avg))

                p.add_(update, alpha=-group['lr'])

        return loss

FriendlySAM

Bases: BaseOptimizer

Sharpness-aware minimization with momentum adjusted perturbations.

Compute gradients at the current weights before calling step(). The closure must recompute the loss and gradients at the perturbed weights.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
base_optimizer OptimizerType

Optimizer class to instantiate for the parameter update.

required
rho float

Radius of the neighborhood used to perturb parameters.

0.05
sigma float

Strength of the momentum subtraction in the perturbation gradient.

1.0
lmbda float

Decay rate for perturbation gradient momentum.

0.9
adaptive bool

Scale perturbations by the squared parameter values.

False
perturb_eps float

Stability constant for the perturbation norm.

1e-12
**kwargs dict

Options for the base optimizer.

{}

Examples:

optimizer = FriendlySAM(model.parameters(), torch.optim.AdamW, lr=1e-3)
for inputs, targets in data:
    optimizer.zero_grad()

    def closure():
        optimizer.zero_grad()
        loss = loss_fn(model(inputs), targets)
        loss.backward()
        return loss

    closure()
    optimizer.step(closure)
Source code in pytorch_optimizer/optimizer/sam.py
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
974
975
976
977
978
979
980
981
982
983
984
985
986
class FriendlySAM(BaseOptimizer):
    """Sharpness-aware minimization with momentum adjusted perturbations.

    Compute gradients at the current weights before calling `step()`. The closure
    must recompute the loss and gradients at the perturbed weights.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        base_optimizer: Optimizer class to instantiate for the parameter update.
        rho: Radius of the neighborhood used to perturb parameters.
        sigma: Strength of the momentum subtraction in the perturbation gradient.
        lmbda: Decay rate for perturbation gradient momentum.
        adaptive: Scale perturbations by the squared parameter values.
        perturb_eps: Stability constant for the perturbation norm.
        **kwargs (dict): Options for the base optimizer.

    Examples:
        ```python
        optimizer = FriendlySAM(model.parameters(), torch.optim.AdamW, lr=1e-3)
        for inputs, targets in data:
            optimizer.zero_grad()

            def closure():
                optimizer.zero_grad()
                loss = loss_fn(model(inputs), targets)
                loss.backward()
                return loss

            closure()
            optimizer.step(closure)
        ```

    """

    def __init__(
        self,
        params: ParamsT,
        base_optimizer: OptimizerType,
        rho: float = 0.05,
        sigma: float = 1.0,
        lmbda: float = 0.9,
        adaptive: bool = False,
        perturb_eps: float = 1e-12,
        **kwargs,
    ):
        self.validate_non_negative(rho, 'rho')
        self.validate_non_negative(sigma, 'sigma')
        self.validate_non_negative(lmbda, 'lmbda')
        self.validate_non_negative(perturb_eps, 'perturb_eps')

        self.perturb_eps = perturb_eps

        defaults: Defaults = {'rho': rho, 'sigma': sigma, 'lmbda': lmbda, 'adaptive': adaptive}
        defaults.update(kwargs)

        super().__init__(params, defaults)

        self.base_optimizer: Optimizer = base_optimizer(self.param_groups, **kwargs)
        self.param_groups = self.base_optimizer.param_groups

    def __str__(self) -> str:
        return 'FriendlySAM'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        pass

    @torch.no_grad()
    def first_step(self, zero_grad: bool = False) -> None:
        for group in self.param_groups:
            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad
                state = self.state[p]

                if 'momentum' not in state:
                    state['momentum'] = grad.clone()
                else:
                    momentum = state['momentum']

                    grad.sub_(momentum, alpha=group['sigma'])
                    momentum.lerp_(grad, weight=1.0 - group['lmbda'])

        grad_norm = get_global_gradient_norm(self.param_groups, weight_adaptive=True)
        grad_norm.sqrt_().squeeze_(0).add_(self.perturb_eps)

        for group in self.param_groups:
            scale = group['rho'] / grad_norm

            for i, p in enumerate(group['params']):
                if p.grad is None:
                    continue

                grad = p.grad

                self.state[p]['old_p'] = p.clone()
                self.state[f'old_grad_p_{i}']['old_grad_p'] = grad.clone()

                e_w = (torch.pow(p, 2) if group['adaptive'] else 1.0) * grad * scale.to(p)

                p.add_(e_w)

        if zero_grad:
            self.zero_grad()

    @torch.no_grad()
    def second_step(self, zero_grad: bool = False):
        for group in self.param_groups:
            for p in group['params']:
                if 'old_p' in self.state[p]:
                    p.copy_(self.state[p].pop('old_p'))

        self.base_optimizer.step()

        if zero_grad:
            self.zero_grad()

    @torch.no_grad()
    def step(self, closure: Closure = None):
        """Perturb weights, recompute gradients, and apply the base optimizer update.

        Args:
            closure: Callable that clears gradients and recomputes the loss and gradients. Compute the initial
                gradients before calling this method.

        Raises:
            NoClosureError: No closure is supplied.

        """
        if closure is None:
            raise NoClosureError(str(self))

        self.first_step(zero_grad=True)

        with torch.enable_grad():
            closure()

        self.second_step()

    def state_dict(self) -> dict:
        state = super().state_dict()
        state['base_optimizer'] = self.base_optimizer.state_dict()
        return state

    def load_state_dict(self, state_dict: dict):
        super().load_state_dict(state_dict)
        if 'base_optimizer' in state_dict:
            self.base_optimizer.load_state_dict(state_dict['base_optimizer'])
            self.param_groups = self.base_optimizer.param_groups
        else:
            self.base_optimizer.param_groups = self.param_groups

step(closure=None)

Perturb weights, recompute gradients, and apply the base optimizer update.

Parameters:

Name Type Description Default
closure Closure

Callable that clears gradients and recomputes the loss and gradients. Compute the initial gradients before calling this method.

None

Raises:

Type Description
NoClosureError

No closure is supplied.

Source code in pytorch_optimizer/optimizer/sam.py
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
@torch.no_grad()
def step(self, closure: Closure = None):
    """Perturb weights, recompute gradients, and apply the base optimizer update.

    Args:
        closure: Callable that clears gradients and recomputes the loss and gradients. Compute the initial
            gradients before calling this method.

    Raises:
        NoClosureError: No closure is supplied.

    """
    if closure is None:
        raise NoClosureError(str(self))

    self.first_step(zero_grad=True)

    with torch.enable_grad():
        closure()

    self.second_step()

Fromage

Bases: BaseOptimizer

Gradient descent scaled by the ratio of parameter and gradient norms.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.01
p_bound float | None

Restricts the optimization to a bounded set. For example, a value of 2.0 restricts parameter norms to lie within 2x their initial norms, which helps regularize the model class.

None
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/fromage.py
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
class Fromage(BaseOptimizer):
    """Gradient descent scaled by the ratio of parameter and gradient norms.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        p_bound: Restricts the optimization to a bounded set. For example, a value of 2.0 restricts parameter norms
            to lie within 2x their initial norms, which helps regularize the model class.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self, params: ParamsT, lr: float = 1e-2, p_bound: float | None = None, maximize: bool = False, **kwargs
    ):
        self.validate_learning_rate(lr)

        self.p_bound = p_bound
        self.maximize = maximize

        defaults: Defaults = {'lr': lr}

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Fromage'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0 and self.p_bound is not None:
                state['max'] = p.norm().mul_(self.p_bound)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            pre_factor: float = math.sqrt(1 + group['lr'] ** 2)

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                p, grad = self.view_as_real(p, grad)

                p_norm, g_norm = p.norm(), grad.norm()

                if p_norm > 0.0 and g_norm > 0.0:
                    p.add_(grad * (p_norm / g_norm), alpha=-group['lr'])
                else:
                    p.add_(grad, alpha=-group['lr'])

                p.div_(pre_factor)

                if self.p_bound is not None:
                    p_norm = p.norm()
                    if p_norm > state['max']:
                        p.mul_(state['max']).div_(p_norm)

        return loss

FTRL

Bases: BaseOptimizer

Follow the Regularized Leader updates with L1 and L2 penalties.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
lr_power float

Exponent controlling the accumulated gradient correction, typically -0.5.

-0.5
beta float

Offset in the adaptive update denominator.

0.0
lambda_1 float

L1 regularization parameter.

0.0
lambda_2 float

L2 regularization parameter.

0.0
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/ftrl.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
class FTRL(BaseOptimizer):
    """Follow the Regularized Leader updates with L1 and L2 penalties.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        lr_power: Exponent controlling the accumulated gradient correction, typically `-0.5`.
        beta: Offset in the adaptive update denominator.
        lambda_1: L1 regularization parameter.
        lambda_2: L2 regularization parameter.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        lr_power: float = -0.5,
        beta: float = 0.0,
        lambda_1: float = 0.0,
        lambda_2: float = 0.0,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_non_negative(beta, 'beta')
        self.validate_non_positive(lr_power, 'lr_power')
        self.validate_non_negative(lambda_1, 'lambda_1')
        self.validate_non_negative(lambda_2, 'lambda_2')

        self.maximize = maximize

        defaults: Defaults = {'lr': lr, 'lr_power': lr_power, 'beta': beta, 'lambda_1': lambda_1, 'lambda_2': lambda_2}

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'FTRL'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['z'] = torch.zeros_like(p)
                state['n'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                z, n = state['z'], state['n']

                p, grad, z, n = self.view_as_real(p, grad, z, n)

                grad_p2 = grad.pow(2)

                sigma = (n + grad_p2).pow_(-group['lr_power']).sub_(n.pow(-group['lr_power'])).div_(group['lr'])

                z.add_(grad).sub_(sigma.mul(p))
                n.add_(grad_p2)

                update = z.sign().mul_(group['lambda_1']).sub_(z)
                update.div_((group['beta'] + n.sqrt()).div_(group['lr']).add_(group['lambda_2']))

                p.copy_(update)
                p.masked_fill_(z.abs() <= group['lambda_1'], 0.0)

        return loss

GaLore

Bases: BaseOptimizer

AdamW with low rank gradient projection.

Add rank, update_proj_gap, scale, and projection_type to parameter groups containing matrix weights to enable projection.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
eps float

Term added to the denominator to improve numerical stability.

1e-06
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/galore.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
class GaLore(BaseOptimizer):
    """AdamW with low rank gradient projection.

    Add `rank`, `update_proj_gap`, `scale`, and `projection_type` to parameter groups
    containing matrix weights to enable projection.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        eps: float = 1e-6,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'GaLore'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size: float = group['lr'] * bias_correction2_sq / bias_correction1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                if 'rank' in group and p.dim() == 2:
                    if 'projector' not in state:
                        state['projector'] = GaLoreProjector(
                            rank=group['rank'],
                            update_proj_gap=group['update_proj_gap'],
                            scale=group['scale'],
                            projection_type=group['projection_type'],
                        )

                    grad = state['projector'].project(grad, group['step'])

                if 'exp_avg' not in state:
                    state['exp_avg'] = torch.zeros_like(grad)
                    state['exp_avg_sq'] = torch.zeros_like(grad)

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']
                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                de_nom = exp_avg_sq.sqrt().add_(group['eps'])

                norm_grad = exp_avg / de_nom

                if 'rank' in group and p.dim() == 2:
                    norm_grad = state['projector'].project_back(norm_grad)

                p.add_(norm_grad, alpha=-step_size)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=True,
                    fixed_decay=False,
                )

        return loss

get_supported_optimizers(filters=None)

List registered optimizer names in alphabetical order.

Parameters:

Name Type Description Default
filters str | list[str] | None

Wildcard pattern or list of patterns, such as '*adam*'. None selects all names.

None

Returns:

Type Description
list[str]

list[str]: Matching names in lowercase, without duplicates.

Source code in pytorch_optimizer/optimizer/__init__.py
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
def get_supported_optimizers(filters: str | list[str] | None = None) -> list[str]:
    """List registered optimizer names in alphabetical order.

    Args:
        filters: Wildcard pattern or list of patterns, such as `'*adam*'`. `None` selects all names.

    Returns:
        list[str]: Matching names in lowercase, without duplicates.

    """
    if filters is None:
        return sorted(OPTIMIZERS.keys())

    include_filters: Sequence[str] = filters if isinstance(filters, (tuple, list)) else [filters]

    filtered_list: set[str] = set()
    for include_filter in include_filters:
        filtered_list.update(fnmatch.filter(OPTIMIZERS.keys(), include_filter))

    return sorted(filtered_list)

Grams

Bases: BaseOptimizer

Adaptive updates combining gradient signs with momentum magnitudes.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
eps float

Term added to the denominator to improve numerical stability.

1e-06
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/grams.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
class Grams(BaseOptimizer):
    """Adaptive updates combining gradient signs with momentum magnitudes.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        eps: float = 1e-6,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Grams'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                update = (exp_avg / bias_correction1) / (exp_avg_sq.sqrt() / bias_correction2_sq).add_(group['eps'])
                update.abs_().mul_(grad.sign())

                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=False,
                )

                p.add_(update, alpha=-group['lr'])

        return loss

Gravity

Bases: BaseOptimizer

Kinematic optimization with a gradient dependent velocity update.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.01
alpha float

Alpha controls the V initialization.

0.01
beta float

Beta will be used to compute running average of V.

0.9
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/gravity.py
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
class Gravity(BaseOptimizer):
    """Kinematic optimization with a gradient dependent velocity update.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        alpha: Alpha controls the V initialization.
        beta: Beta will be used to compute running average of V.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-2,
        alpha: float = 0.01,
        beta: float = 0.9,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(alpha, 'alpha', 0.0, 1.0)
        self.validate_range(beta, 'beta', 0.0, 1.0, range_type='[]')

        self.maximize = maximize

        defaults: Defaults = {'lr': lr, 'alpha': alpha, 'beta': beta}

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Gravity'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['v'] = torch.empty_like(p).normal_(mean=0.0, std=group['alpha'] / group['lr'])

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta_t: float = (group['beta'] * group['step'] + 1) / (group['step'] + 2)

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                v = state['v']

                p, grad, v = self.view_as_real(p, grad, v)

                m = 1.0 / grad.abs().max()
                zeta = grad / (1.0 + (grad / m) ** 2)

                v.lerp_(zeta, weight=1.0 - beta_t)

                p.add_(v, alpha=-group['lr'])

        return loss

GrokFastAdamW

Bases: BaseOptimizer

AdamW with amplification of slow gradient components.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.0001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.99)
grokfast bool

Whether to use grokfast.

True
grokfast_alpha float

Momentum hyperparameter of the EMA.

0.98
grokfast_lamb float

Amplifying factor hyperparameter of the filter.

2.0
grokfast_after_step int

Warmup step for grokfast.

0
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
normalize_lr bool

Divide the learning rate by 1 + grokfast_lamb when Grokfast is enabled.

True
eps float

Term added to the denominator to improve numerical stability.

1e-08
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/grokfast.py
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
class GrokFastAdamW(BaseOptimizer):
    """AdamW with amplification of slow gradient components.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        grokfast: Whether to use grokfast.
        grokfast_alpha: Momentum hyperparameter of the EMA.
        grokfast_lamb: Amplifying factor hyperparameter of the filter.
        grokfast_after_step: Warmup step for grokfast.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        normalize_lr: Divide the learning rate by `1 + grokfast_lamb` when Grokfast is enabled.
        eps: Term added to the denominator to improve numerical stability.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-4,
        betas: Betas = (0.9, 0.99),
        grokfast: bool = True,
        grokfast_alpha: float = 0.98,
        grokfast_lamb: float = 2.0,
        grokfast_after_step: int = 0,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        normalize_lr: bool = True,
        eps: float = 1e-8,
        foreach: bool | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(grokfast_alpha, 'grokfast_alpha', 0.0, 1.0)
        self.validate_non_negative(eps, 'eps')

        self.foreach = foreach
        self.maximize = maximize

        if grokfast and normalize_lr:
            lr /= 1.0 + grokfast_lamb

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'grokfast': grokfast,
            'grokfast_alpha': grokfast_alpha,
            'grokfast_lamb': grokfast_lamb,
            'grokfast_after_step': grokfast_after_step,
            'foreach': foreach,
            'eps': eps,
        }
        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'GrokFastAdamW'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                if group['grokfast'] and group['grokfast_lamb'] > 0.0:
                    grok_exp_avg = grad.clone()
                    self.maximize_gradient(grok_exp_avg, maximize=self.maximize)
                    state['grok_exp_avg'] = grok_exp_avg

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        if group.get('foreach') is False:
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        grok_exp_avgs: list[torch.Tensor],
        should_grokfast: bool,
    ) -> None:
        beta1, beta2 = group['betas']

        bias_correction1: float = self.debias(beta1, group['step'])
        bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=group['lr'],
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        if should_grokfast:
            torch._foreach_lerp_(grok_exp_avgs, grads, weight=1.0 - group['grokfast_alpha'])
            torch._foreach_add_(grads, grok_exp_avgs, alpha=group['grokfast_lamb'])

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - beta1)
        torch._foreach_mul_(exp_avg_sqs, beta2)
        torch._foreach_addcmul_(exp_avg_sqs, grads, grads, value=1.0 - beta2)

        de_noms = torch._foreach_sqrt(exp_avg_sqs)
        torch._foreach_div_(de_noms, bias_correction2_sq)
        torch._foreach_clamp_min_(de_noms, group['eps'])

        updates = torch._foreach_div(exp_avgs, bias_correction1)
        torch._foreach_div_(updates, de_noms)

        torch._foreach_add_(params, updates, alpha=-group['lr'])

    def _step_per_param(self, group: ParamGroup, should_grokfast: bool) -> None:
        beta1, beta2 = group['betas']

        bias_correction1: float = self.debias(beta1, group['step'])
        bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            exp_avg, exp_avg_sq, grok_exp_avg = (
                state['exp_avg'],
                state['exp_avg_sq'],
                state.get('grok_exp_avg', None),
            )

            p, grad, exp_avg, exp_avg_sq, grok_exp_avg = self.view_as_real(p, grad, exp_avg, exp_avg_sq, grok_exp_avg)

            self.apply_weight_decay(
                p=p,
                grad=grad,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=group['fixed_decay'],
            )

            if should_grokfast:
                grok_exp_avg.lerp_(grad, weight=1.0 - group['grokfast_alpha'])
                grad.add_(grok_exp_avg, alpha=group['grokfast_lamb'])

            exp_avg.lerp_(grad, weight=1.0 - beta1)
            exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

            de_nom = exp_avg_sq.sqrt().div_(bias_correction2_sq).clamp_(min=group['eps'])

            update = exp_avg.div(bias_correction1).div_(de_nom)

            p.add_(update, alpha=-group['lr'])

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            should_grokfast: bool = (
                group['grokfast'] and group['step'] > group['grokfast_after_step'] and group['grokfast_lamb'] > 0.0
            )

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(
                    group,
                    self.state,
                    state_keys=['exp_avg', 'exp_avg_sq', 'grok_exp_avg'],
                )
                if params:
                    self._step_foreach(
                        group,
                        params,
                        grads,
                        state_dict['exp_avg'],
                        state_dict['exp_avg_sq'],
                        state_dict['grok_exp_avg'],
                        should_grokfast,
                    )
            else:
                self._step_per_param(group, should_grokfast)

        return loss

GSAM

Bases: BaseOptimizer

Sharpness-aware minimization with surrogate gap gradient decomposition.

Use set_closure() to supply the loss and batch before each step. Advance the learning rate scheduler and call update_rho_t() to update the perturbation radius.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
base_optimizer Optimizer

Existing optimizer instance for parameter updates.

required
model Module

Model used for the forward passes.

required
rho_scheduler ProportionScheduler

Scheduler that supplies the perturbation radius.

required
alpha float

Weight of the surrogate gap gradient component.

0.4
adaptive bool

Scale perturbations by the squared parameter values.

False
perturb_eps float

Stability constant for the perturbation norm.

1e-12
**kwargs dict

Additional parameter group options.

{}
Source code in pytorch_optimizer/optimizer/sam.py
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
class GSAM(BaseOptimizer):  # pragma: no cover
    """Sharpness-aware minimization with surrogate gap gradient decomposition.

    Use `set_closure()` to supply the loss and batch before each step. Advance the learning
    rate scheduler and call `update_rho_t()` to update the perturbation radius.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        base_optimizer: Existing optimizer instance for parameter updates.
        model: Model used for the forward passes.
        rho_scheduler (ProportionScheduler): Scheduler that supplies the perturbation radius.
        alpha: Weight of the surrogate gap gradient component.
        adaptive: Scale perturbations by the squared parameter values.
        perturb_eps: Stability constant for the perturbation norm.
        **kwargs (dict): Additional parameter group options.

    """

    def __init__(
        self,
        params: ParamsT,
        base_optimizer: Optimizer,
        model: nn.Module,
        rho_scheduler,
        alpha: float = 0.4,
        adaptive: bool = False,
        perturb_eps: float = 1e-12,
        **kwargs,
    ):
        self.validate_range(alpha, 'alpha', 0.0, 1.0)

        self.model = model
        self.rho_scheduler = rho_scheduler
        self.alpha = alpha
        self.adaptive = adaptive
        self.perturb_eps = perturb_eps

        self.rho_t: float = 0.0
        self.forward_backward_func: Callable | None = None

        if hasattr(ReduceOp, 'AVG'):
            self.grad_reduce = ReduceOp.AVG
            self.manual_average: bool = False
        else:
            self.grad_reduce = ReduceOp.SUM
            self.manual_average: bool = True

        self.base_optimizer = base_optimizer
        self.param_groups = self.base_optimizer.param_groups

        defaults: Defaults = {'adaptive': adaptive, **kwargs}

        super().__init__(params, defaults)

        self.update_rho_t()

    def __str__(self) -> str:
        return 'GSAM'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        pass

    @torch.no_grad()
    def update_rho_t(self) -> float:
        self.rho_t = self.rho_scheduler.step()
        return self.rho_t

    @torch.no_grad()
    def perturb_weights(self, rho: float):
        grad_norm = self.grad_norm(weight_adaptive=self.adaptive)
        for group in self.param_groups:
            scale = rho / (grad_norm + self.perturb_eps)

            for p in group['params']:
                if p.grad is None:
                    continue

                self.state[p]['old_g'] = p.grad.clone()

                e_w = (torch.pow(p, 2) if self.adaptive else 1.0) * p.grad * scale.to(p)

                p.add_(e_w)

                self.state[p]['e_w'] = e_w

    @torch.no_grad()
    def un_perturb(self):
        for group in self.param_groups:
            for p in group['params']:
                if 'e_w' in self.state[p]:
                    p.sub_(self.state[p]['e_w'])

    @torch.no_grad()
    def gradient_decompose(self, alpha: float = 0.0):
        inner_prod = 0.0
        for group in self.param_groups:
            for p in group['params']:
                if p.grad is None:
                    continue

                inner_prod += torch.sum(self.state[p]['old_g'] * p.grad)

        new_grad_norm = self.grad_norm(by=None)
        old_grad_norm = self.grad_norm(by='old_g')

        cosine = inner_prod / (new_grad_norm * old_grad_norm + self.perturb_eps)

        for group in self.param_groups:
            for p in group['params']:
                if p.grad is None:
                    continue

                vertical = self.state[p]['old_g'] - cosine * old_grad_norm * p.grad / (
                    new_grad_norm + self.perturb_eps
                )
                p.grad.add_(vertical, alpha=-alpha)

    @torch.no_grad()
    def sync_grad(self):
        if is_initialized():
            for group in self.param_groups:
                for p in group['params']:
                    if p.grad is None:
                        continue

                    all_reduce(p.grad, op=self.grad_reduce)
                    if self.manual_average:
                        p.grad.div_(float(get_world_size()))

    @torch.no_grad()
    def grad_norm(self, by: str | None = None, weight_adaptive: bool = False) -> torch.Tensor:
        if not by and not weight_adaptive:
            return get_global_gradient_norm(self.param_groups).sqrt_().squeeze(0)

        return torch.norm(
            torch.stack(
                [
                    ((torch.abs(p) if weight_adaptive else 1.0) * (p.grad if not by else self.state[p][by])).norm(p=2)
                    for group in self.param_groups
                    for p in group['params']
                    if p.grad is not None
                ]
            ),
            p=2,
        )

    def maybe_no_sync(self):
        if is_initialized() and hasattr(self.model, 'no_sync'):
            return self.model.no_sync()  # ty: ignore[call-non-callable]
        return ExitStack()

    @torch.no_grad()
    def set_closure(self, loss_fn: nn.Module, inputs: torch.Tensor, targets: torch.Tensor, **kwargs) -> None:
        """Store a forward backward closure for the current batch.

        The closure clears gradients, evaluates the model and loss, and runs backpropagation.

        Args:
            loss_fn: Callable accepting model predictions and targets.
            inputs: Model inputs for the current batch.
            targets: Target values for the current batch.
            **kwargs (dict): Additional arguments for the loss function.

        """

        def get_grad() -> tuple[Any, torch.Tensor]:
            self.base_optimizer.zero_grad()

            with torch.enable_grad():
                outputs = self.model(inputs)
                loss = loss_fn(outputs, targets, **kwargs)

            loss.backward()

            return outputs, loss.detach()

        self.forward_backward_func = get_grad

    @torch.no_grad()
    def step(self, closure: Closure = None) -> tuple[Any, torch.Tensor]:
        get_grad = cast(Callable[[], tuple[Any, torch.Tensor]], closure or self.forward_backward_func)

        with self.maybe_no_sync():
            outputs, loss = get_grad()

            self.perturb_weights(rho=self.rho_t)

            disable_running_stats(self.model)

            get_grad()

            self.gradient_decompose(self.alpha)

            self.un_perturb()

        self.sync_grad()

        self.base_optimizer.step()

        enable_running_stats(self.model)

        return outputs, loss

    def state_dict(self) -> dict:
        state = super().state_dict()
        state['base_optimizer'] = self.base_optimizer.state_dict()
        return state

    def load_state_dict(self, state_dict: dict):
        super().load_state_dict(state_dict)
        if 'base_optimizer' in state_dict:
            self.base_optimizer.load_state_dict(state_dict['base_optimizer'])
            self.param_groups = self.base_optimizer.param_groups
        else:
            self.base_optimizer.param_groups = self.param_groups

set_closure(loss_fn, inputs, targets, **kwargs)

Store a forward backward closure for the current batch.

The closure clears gradients, evaluates the model and loss, and runs backpropagation.

Parameters:

Name Type Description Default
loss_fn Module

Callable accepting model predictions and targets.

required
inputs Tensor

Model inputs for the current batch.

required
targets Tensor

Target values for the current batch.

required
**kwargs dict

Additional arguments for the loss function.

{}
Source code in pytorch_optimizer/optimizer/sam.py
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
@torch.no_grad()
def set_closure(self, loss_fn: nn.Module, inputs: torch.Tensor, targets: torch.Tensor, **kwargs) -> None:
    """Store a forward backward closure for the current batch.

    The closure clears gradients, evaluates the model and loss, and runs backpropagation.

    Args:
        loss_fn: Callable accepting model predictions and targets.
        inputs: Model inputs for the current batch.
        targets: Target values for the current batch.
        **kwargs (dict): Additional arguments for the loss function.

    """

    def get_grad() -> tuple[Any, torch.Tensor]:
        self.base_optimizer.zero_grad()

        with torch.enable_grad():
            outputs = self.model(inputs)
            loss = loss_fn(outputs, targets, **kwargs)

        loss.backward()

        return outputs, loss.detach()

    self.forward_backward_func = get_grad

Kate

Bases: BaseOptimizer

Scale invariant AdaGrad-style updates without square root normalization.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
delta float

Delta parameter, typically 0.0 or 1e-8.

0.0
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Epsilon value for numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/kate.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
class Kate(BaseOptimizer):
    """Scale invariant AdaGrad-style updates without square root normalization.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        delta: Delta parameter, typically 0.0 or 1e-8.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Epsilon value for numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        delta: float = 0.0,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(delta, 'delta', 0.0, 1.0, '[)')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'delta': delta,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Kate'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['m'] = torch.zeros_like(p)
                state['b'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                m, b = state['m'], state['b']

                p, grad, m, b = self.view_as_real(p, grad, m, b)

                self.apply_weight_decay(
                    p=p,
                    grad=p.grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                grad_p2 = grad.pow(2)

                b.mul_(b).add_(grad_p2).add_(group['eps'])
                m.mul_(m).add_(grad_p2, alpha=group['delta']).add_(grad_p2 / b).sqrt_()

                update = m.mul(grad).div_(b)

                p.add_(update, alpha=-group['lr'])

                b.sqrt_()

        return loss

Kron

Bases: BaseOptimizer

Preconditioned SGD with Kronecker factored preconditioners.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
momentum float

Momentum factor.

0.9
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
pre_conditioner_update_probability float | Callable[[int], Tensor] | None

Update frequency as a fraction or a callable of the step index. None uses the default decay schedule.

None
max_size_triangular int

Largest dimension that can use a triangular preconditioner.

8192
min_ndim_triangular int

Minimum tensor dimensionality for triangular preconditioners.

2
memory_save_mode MEMORY_SAVE_MODE_TYPE | None

Diagonal storage policy: None, 'one_diag', 'smart_one_diag', or 'all_diag'.

None
momentum_into_precondition_update bool

Use momentum instead of raw gradients when updating preconditioners.

True
mu_dtype dtype | None

Dtype of the momentum accumulator.

None
precondition_dtype dtype | None

Dtype of the preconditioner.

float32
balance_prob float

Probability of performing balancing.

0.01
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/psgd.py
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
class Kron(BaseOptimizer):
    """Preconditioned SGD with Kronecker factored preconditioners.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Momentum factor.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        pre_conditioner_update_probability: Update frequency as a fraction or a callable of the step index. `None`
            uses the default decay schedule.
        max_size_triangular: Largest dimension that can use a triangular preconditioner.
        min_ndim_triangular: Minimum tensor dimensionality for triangular preconditioners.
        memory_save_mode: Diagonal storage policy: `None`, `'one_diag'`, `'smart_one_diag'`, or `'all_diag'`.
        momentum_into_precondition_update: Use momentum instead of raw gradients when updating preconditioners.
        mu_dtype: Dtype of the momentum accumulator.
        precondition_dtype: Dtype of the preconditioner.
        balance_prob: Probability of performing balancing.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        momentum: float = 0.9,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        pre_conditioner_update_probability: float | Callable[[int], torch.Tensor] | None = None,
        max_size_triangular: int = 8192,
        min_ndim_triangular: int = 2,
        memory_save_mode: MEMORY_SAVE_MODE_TYPE | None = None,
        momentum_into_precondition_update: bool = True,
        mu_dtype: torch.dtype | None = None,
        precondition_dtype: torch.dtype | None = torch.float32,
        balance_prob: float = 0.01,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0)
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.balance_prob: float = balance_prob
        self.eps: float = torch.finfo(torch.bfloat16).tiny
        self.maximize = maximize

        defaults = {
            'lr': lr,
            'momentum': momentum,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'pre_conditioner_update_probability': pre_conditioner_update_probability,
            'max_size_triangular': max_size_triangular,
            'min_ndim_triangular': min_ndim_triangular,
            'memory_save_mode': memory_save_mode,
            'momentum_into_precondition_update': momentum_into_precondition_update,
            'precondition_lr': 1e-1,
            'precondition_init_scale': 1.0,
            'mu_dtype': mu_dtype,
            'precondition_dtype': precondition_dtype,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Kron'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        first_group = self.param_groups[0]
        update_prob = first_group['pre_conditioner_update_probability']
        if update_prob is None:
            update_prob = precondition_update_prob_schedule()(first_group.get('step', 0))
        if callable(update_prob):
            update_prob = cast(torch.Tensor, update_prob(first_group.get('step', 0)))

        update_counter = first_group.get('update_counter', 0) + 1
        do_update: bool = update_counter >= 1 / update_prob
        first_group['update_counter'] = 0 if do_update else update_counter

        balance: bool = np.random.random() < self.balance_prob and do_update

        for group in self.param_groups:
            if 'step' in group:
                group['step'] += 1
            else:
                group['step'] = 1

            bias_correction1: float = self.debias(group['momentum'], group['step'])

            mu_dtype, precondition_dtype = group['mu_dtype'], group['precondition_dtype']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad
                if grad.is_sparse:
                    raise NoSparseGradientError(str(self))

                if torch.is_complex(p):
                    raise NoComplexParameterError(str(self))

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                if len(state) == 0:
                    state['momentum_buffer'] = torch.zeros_like(p, dtype=mu_dtype or p.dtype)
                    state['Q'], state['expressions'] = initialize_q_expressions(
                        p,
                        group['precondition_init_scale'],
                        group['max_size_triangular'],
                        group['min_ndim_triangular'],
                        group['memory_save_mode'],
                        dtype=precondition_dtype,
                    )

                momentum_buffer = state['momentum_buffer']
                momentum_buffer.mul_(group['momentum']).add_(grad, alpha=1.0 - group['momentum'])

                if mu_dtype is not None:
                    momentum_buffer = momentum_buffer.to(dtype=mu_dtype, non_blocking=True)

                de_biased_momentum = (momentum_buffer / bias_correction1).to(
                    dtype=precondition_dtype, non_blocking=True
                )

                if grad.dim() > 1 and balance:
                    balance_q(state['Q'])

                if do_update:
                    update_precondition(
                        state['Q'],
                        state['expressions'],
                        torch.randn_like(de_biased_momentum, dtype=precondition_dtype),
                        de_biased_momentum if group['momentum_into_precondition_update'] else grad,
                        group['precondition_lr'],
                        self.eps,
                    )

                precondition_grad = get_precondition_grad(state['Q'], state['expressions'], de_biased_momentum).to(
                    dtype=p.dtype, non_blocking=True
                )

                precondition_grad.mul_(torch.clamp(1.1 / (precondition_grad.square().mean().sqrt() + 1e-6), max=1.0))

                if group['weight_decay'] != 0 and p.dim() >= 2:
                    precondition_grad.add_(p, alpha=group['weight_decay'])

                p.add_(precondition_grad, alpha=-group['lr'])

        return loss

Lamb

Bases: BaseOptimizer

Adam updates with a trust ratio for each parameter tensor.

The default update follows version 3 of the paper without bias correction.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
rectify bool

Perform the rectified update similar to RAdam.

False
degenerated_to_sgd bool

Use an SGD update before the moving average reaches the rectification threshold.

False
n_sma_threshold int

Minimum effective simple moving average length for rectification.

5
grad_averaging bool

Scale new gradient contributions by 1 - beta1.

True
max_grad_norm float

Reference norm for gradient scaling when pre_norm=True. 0 disables scaling.

1.0
adam bool

Use a trust ratio of 1 for all parameters.

False
pre_norm bool

Divide gradients by a scaling factor derived from their global norm.

False
eps float

Term added to the denominator to improve numerical stability.

1e-06
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/lamb.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
class Lamb(BaseOptimizer):
    """Adam updates with a trust ratio for each parameter tensor.

    The default update follows version 3 of the paper without bias correction.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        rectify: Perform the rectified update similar to RAdam.
        degenerated_to_sgd: Use an SGD update before the moving average reaches the rectification threshold.
        n_sma_threshold: Minimum effective simple moving average length for rectification.
        grad_averaging: Scale new gradient contributions by `1 - beta1`.
        max_grad_norm: Reference norm for gradient scaling when `pre_norm=True`. `0` disables scaling.
        adam: Use a trust ratio of 1 for all parameters.
        pre_norm: Divide gradients by a scaling factor derived from their global norm.
        eps: Term added to the denominator to improve numerical stability.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        maximize: Maximize the objective instead of minimizing it.

    """

    clamp: float = 10.0

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        rectify: bool = False,
        degenerated_to_sgd: bool = False,
        n_sma_threshold: int = 5,
        grad_averaging: bool = True,
        max_grad_norm: float = 1.0,
        adam: bool = False,
        pre_norm: bool = False,
        eps: float = 1e-6,
        foreach: bool | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(max_grad_norm, 'max_grad_norm')
        self.validate_non_negative(eps, 'eps')

        self.degenerated_to_sgd = degenerated_to_sgd
        self.n_sma_threshold = n_sma_threshold
        self.pre_norm = pre_norm
        self.foreach = foreach
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'rectify': rectify,
            'grad_averaging': grad_averaging,
            'max_grad_norm': max_grad_norm,
            'adam': adam,
            'eps': eps,
            'foreach': foreach,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Lamb'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

                if group.get('adanorm'):
                    state['exp_grad_adanorm'] = torch.zeros((1,), dtype=p.dtype, device=p.device)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        """Check tensor compatibility and options for batched updates.

        Disable batched updates when using AdaNorm or rectification.
        """
        if group.get('foreach') is False:
            return False

        if group.get('adanorm') or group.get('rectify'):
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        grad_norm: torch.Tensor | float,
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        step_size: float,
    ) -> None:
        beta1, beta2 = group['betas']
        eps = group['eps']
        beta3: float = 1.0 - beta1 if group['grad_averaging'] else 1.0

        if self.maximize:
            torch._foreach_neg_(grads)

        if self.pre_norm:
            if isinstance(grad_norm, torch.Tensor):
                grad_norm = grad_norm.reshape(())

            torch._foreach_mul_(grads, grad_norm)

        if group['weight_decouple']:
            self.apply_weight_decay_foreach(
                params=params,
                grads=grads,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=True,
                fixed_decay=group['fixed_decay'],
            )

        torch._foreach_mul_(exp_avgs, beta1)
        torch._foreach_add_(exp_avgs, grads, alpha=beta3)

        torch._foreach_mul_(exp_avg_sqs, beta2)
        torch._foreach_addcmul_(exp_avg_sqs, grads, grads, value=1.0 - beta2)

        updates = torch._foreach_sqrt(exp_avg_sqs)
        torch._foreach_add_(updates, eps)
        torch._foreach_reciprocal_(updates)
        torch._foreach_mul_(updates, exp_avgs)

        if not group['weight_decouple'] and group['weight_decay'] > 0.0:
            torch._foreach_add_(updates, params, alpha=group['weight_decay'])

        weight_norms = torch._foreach_norm(params)
        torch._foreach_clamp_max_(weight_norms, self.clamp)

        p_norms = torch._foreach_norm(updates)

        trust_ratios = torch._foreach_div(weight_norms, torch._foreach_add(p_norms, eps))
        trust_ratios = [
            torch.where((wn != 0) & (pn != 0), ratio, torch.ones_like(ratio))
            for wn, pn, ratio in zip(weight_norms, p_norms, trust_ratios)
        ]

        for p, wn, pn, trust_ratio in zip(params, weight_norms, p_norms, trust_ratios):
            state = self.state[p]
            state['weight_norm'] = wn
            state['adam_norm'] = pn
            state['trust_ratio'] = trust_ratio

        if not group['adam']:
            torch._foreach_mul_(updates, trust_ratios)

        torch._foreach_add_(params, updates, alpha=-step_size)

    @torch.no_grad()
    def get_global_gradient_norm(self) -> torch.Tensor | float:
        if self.defaults['max_grad_norm'] == 0.0:
            return 1.0

        global_grad_norm = get_global_gradient_norm(self.param_groups)
        global_grad_norm.sqrt_().add_(self.defaults['eps'])

        return torch.clamp(self.defaults['max_grad_norm'] / global_grad_norm, max=1.0)

    def update(
        self,
        p: torch.Tensor,
        group: ParamGroup,
        grad_norm: torch.Tensor | float,
        n_sma: float,
        step_size: float,
        beta1: float,
        beta2: float,
        beta3: float,
    ) -> None:
        grad = p.grad
        if grad is None:
            return

        if self.pre_norm:
            grad.mul_(grad_norm)

        self.maximize_gradient(grad, maximize=self.maximize)

        state = self.state[p]

        exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

        p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

        s_grad = self.get_adanorm_gradient(
            grad=grad,
            adanorm=group.get('adanorm', False),
            exp_grad_norm=state.get('exp_grad_adanorm', None),
            r=group.get('adanorm_r', None),
        )

        exp_avg.mul_(beta1).add_(s_grad, alpha=beta3)
        exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

        self.apply_weight_decay(
            p=p,
            grad=None,
            lr=group['lr'],
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        if group['rectify'] and step_size <= 0:
            return

        de_nom: torch.Tensor | None = None

        if group['rectify']:
            update = p.clone()
            if n_sma >= self.n_sma_threshold:
                de_nom = exp_avg_sq.sqrt().add_(group['eps'])
                update.addcdiv_(exp_avg, de_nom, value=-step_size)
            else:
                update.add_(exp_avg, alpha=-step_size)
        else:
            update = exp_avg / exp_avg_sq.sqrt().add_(group['eps'])
            if not group['weight_decouple'] and group['weight_decay'] > 0.0:
                update.add_(p, alpha=group['weight_decay'])

        weight_norm = torch.linalg.norm(p).clamp_(min=0, max=self.clamp)
        p_norm = torch.linalg.norm(update)
        trust_ratio: float = 1.0 if weight_norm == 0 or p_norm == 0 else weight_norm / (p_norm + group['eps'])

        state['weight_norm'] = weight_norm
        state['adam_norm'] = p_norm
        state['trust_ratio'] = trust_ratio

        if group['adam']:
            trust_ratio = 1.0

        if group['rectify']:
            if n_sma >= self.n_sma_threshold:
                p.addcdiv_(exp_avg, de_nom, value=-step_size * trust_ratio)
            else:
                p.add_(exp_avg, alpha=-step_size * trust_ratio)
        else:
            p.add_(update, alpha=-step_size * trust_ratio)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        grad_norm = 1.0
        if self.pre_norm:
            grad_norm = self.get_global_gradient_norm()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            beta3: float = 1.0 - beta1 if group['grad_averaging'] else 1.0
            bias_correction1: float = self.debias(beta1, group['step'])

            step_size, n_sma = self.get_rectify_step_size(
                is_rectify=group['rectify'],
                step=group['step'],
                lr=group['lr'],
                beta2=beta2,
                n_sma_threshold=self.n_sma_threshold,
                degenerated_to_sgd=self.degenerated_to_sgd,
            )

            step_size = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=step_size,
                bias_correction1=bias_correction1,
            )

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_avg_sq']
                )
                if params:
                    self._step_foreach(
                        group, params, grads, grad_norm, state_dict['exp_avg'], state_dict['exp_avg_sq'], step_size
                    )
            else:
                for p in group['params']:
                    self.update(p, group, grad_norm, n_sma, step_size, beta1, beta2, beta3)

        return loss

LaProp

Bases: BaseOptimizer

Adaptive updates with momentum of preconditioned gradients.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.0004
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
centered bool

Subtract the squared gradient mean from the second moment estimate.

False
steps_before_using_centered int

Number of steps to accumulate gradient means before applying centered updates.

10
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
eps float

Epsilon value for numerical stability.

1e-15
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/laprop.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
class LaProp(BaseOptimizer):
    """Adaptive updates with momentum of preconditioned gradients.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        centered: Subtract the squared gradient mean from the second moment estimate.
        steps_before_using_centered: Number of steps to accumulate gradient means before applying centered updates.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        eps: Epsilon value for numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 4e-4,
        betas: Betas = (0.9, 0.999),
        centered: bool = False,
        steps_before_using_centered: int = 10,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        ams_bound: bool = False,
        eps: float = 1e-15,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.steps_before_using_centered: int = steps_before_using_centered
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'centered': centered,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'ams_bound': ams_bound,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'LaProp'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['exp_avg_lr_1'] = 0.0
                state['exp_avg_lr_2'] = 0.0

                if group['centered']:
                    state['exp_mean_avg_beta2'] = torch.zeros_like(p)

                if group['ams_bound']:
                    state['max_exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq, exp_mean_avg_beta2 = (
                    state['exp_avg'],
                    state['exp_avg_sq'],
                    state.get('exp_mean_avg_beta2', None),
                )

                p, grad, exp_avg, exp_avg_sq, exp_mean_avg_beta2 = self.view_as_real(
                    p, grad, exp_avg, exp_avg_sq, exp_mean_avg_beta2
                )

                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                state['exp_avg_lr_1'] = state['exp_avg_lr_1'] * beta1 + (1.0 - beta1) * group['lr']
                state['exp_avg_lr_2'] = state['exp_avg_lr_2'] * beta2 + (1.0 - beta2)

                bias_correction1: float = state['exp_avg_lr_1'] / group['lr'] if group['lr'] != 0.0 else 1.0
                step_size: float = 1.0 / bias_correction1

                second_moment = exp_avg_sq
                if group['centered']:
                    exp_mean_avg_beta2.lerp_(grad, weight=1.0 - beta2)
                    if group['step'] > self.steps_before_using_centered:
                        second_moment = exp_avg_sq - exp_mean_avg_beta2.pow(2)

                de_nom = self.apply_ams_bound(
                    ams_bound=group['ams_bound'],
                    exp_avg_sq=second_moment,
                    max_exp_avg_sq=state.get('max_exp_avg_sq', None),
                    eps=group['eps'],
                )
                de_nom.div_(bias_correction2_sq)

                exp_avg.mul_(beta1).addcdiv_(grad, de_nom, value=(1.0 - beta1) * group['lr'])

                if group.get('cautious'):
                    update = exp_avg.clone()
                    self.apply_cautious(update, grad)
                else:
                    update = exp_avg

                p.add_(update, alpha=-step_size)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

        return loss

LARS

Bases: BaseOptimizer

SGD with learning rates scaled by parameter and gradient norms.

Scalars and vectors use ordinary SGD without weight decay or rate scaling.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
weight_decay float

Weight decay coefficient.

0.0
momentum float

Momentum factor.

0.9
dampening float

Dampening factor for momentum.

0.0
trust_coefficient float

Trust coefficient.

0.001
nesterov bool

Use Nesterov momentum.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/lars.py
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
class LARS(BaseOptimizer):
    """SGD with learning rates scaled by parameter and gradient norms.

    Scalars and vectors use ordinary SGD without weight decay or rate scaling.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        weight_decay: Weight decay coefficient.
        momentum: Momentum factor.
        dampening: Dampening factor for momentum.
        trust_coefficient: Trust coefficient.
        nesterov: Use Nesterov momentum.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        weight_decay: float = 0.0,
        momentum: float = 0.9,
        dampening: float = 0.0,
        trust_coefficient: float = 1e-3,
        nesterov: bool = False,
        foreach: bool | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(momentum, 'momentum', 0.0, 1.0)
        self.validate_range(dampening, 'dampening', 0.0, 1.0)
        self.validate_non_negative(trust_coefficient, 'trust_coefficient')

        self.foreach = foreach
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'weight_decay': weight_decay,
            'momentum': momentum,
            'dampening': dampening,
            'trust_coefficient': trust_coefficient,
            'nesterov': nesterov,
            'foreach': foreach,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Lars'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if group['momentum'] > 0.0:
                state = self.state[p]

                if 'momentum_buffer' not in state:
                    state['momentum_buffer'] = torch.zeros_like(p)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        """Check tensor compatibility and options for batched updates.

        Disable batched updates when using Nesterov momentum.
        """
        if group.get('foreach') is False:
            return False

        if group.get('nesterov'):
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        momentum_buffers: list[torch.Tensor],
    ) -> None:
        if self.maximize:
            torch._foreach_neg_(grads)

        masks = [p.ndim > 1 for p in params]
        masked_params = [p for p, m in zip(params, masks) if m]
        masked_grads = [g for g, m in zip(grads, masks) if m]

        if masked_params:
            param_norms = torch._foreach_norm(masked_params)
            grad_norms = torch._foreach_norm(masked_grads)

            trust_ratios = []
            for pn, gn in zip(param_norms, grad_norms):
                one = torch.ones_like(pn)
                denominator = gn + group['weight_decay'] * pn
                trust_ratio = torch.where(
                    pn > 0.0,
                    torch.where(denominator > 0.0, (group['trust_coefficient'] * pn / denominator), one),
                    one,
                )
                trust_ratios.append(trust_ratio)

            torch._foreach_add_(masked_grads, masked_params, alpha=group['weight_decay'])
            torch._foreach_mul_(masked_grads, trust_ratios)

        if group['momentum'] > 0.0:
            torch._foreach_mul_(momentum_buffers, group['momentum'])
            torch._foreach_add_(momentum_buffers, grads, alpha=1.0 - group['dampening'])
            torch._foreach_copy_(grads, momentum_buffers)

        torch._foreach_add_(params, grads, alpha=-group['lr'])

    def _step_per_param(self, group: ParamGroup) -> None:
        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            if p.ndim > 1:
                param_norm = torch.linalg.norm(p)
                update_norm = torch.linalg.norm(grad)

                one = torch.ones_like(param_norm)
                denominator = update_norm + group['weight_decay'] * param_norm

                trust_ratio = torch.where(
                    param_norm > 0.0,
                    torch.where(denominator > 0.0, (group['trust_coefficient'] * param_norm / denominator), one),
                    one,
                )

                grad.add_(p, alpha=group['weight_decay'])
                grad.mul_(trust_ratio)

            if group['momentum'] > 0.0:
                mb = state['momentum_buffer']
                mb.mul_(group['momentum']).add_(grad, alpha=1.0 - group['dampening'])

                if group['nesterov']:
                    grad.add_(mb, alpha=group['momentum'])
                else:
                    grad.copy_(mb)

            p.add_(grad, alpha=-group['lr'])

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if self._can_use_foreach(group) and group['momentum'] > 0.0:
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['momentum_buffer']
                )
                if params:
                    self._step_foreach(group, params, grads, state_dict['momentum_buffer'])
            else:
                self._step_per_param(group)

        return loss

Lion

Bases: BaseOptimizer

Sign based updates from interpolated gradient momentum.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float | Tensor

Learning rate. A scalar tensor avoids recompilation when the rate changes.

0.0001
betas Betas

Decay rates for update interpolation and gradient momentum.

(0.9, 0.99)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/lion.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
class Lion(BaseOptimizer):
    """Sign based updates from interpolated gradient momentum.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate. A scalar tensor avoids recompilation when the rate changes.
        betas: Decay rates for update interpolation and gradient momentum.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        maximize: Maximize the objective instead of minimizing it.

    """

    _supports_compiled_foreach = True

    def __init__(
        self,
        params: ParamsT,
        lr: float | torch.Tensor = 1e-4,
        betas: Betas = (0.9, 0.99),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        foreach: bool | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'foreach': foreach,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Lion'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)

                if group.get('adanorm'):
                    state['exp_grad_adanorm'] = torch.zeros((1,), dtype=grad.dtype, device=grad.device)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        """Check tensor compatibility and options for batched updates.

        Disable batched updates when using gradient centralization, AdaNorm, or cautious updates.
        """
        if group.get('foreach') is False:
            return False

        if group.get('use_gc') or group.get('adanorm') or group.get('cautious'):
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
    ) -> None:
        beta1, beta2 = group['betas']
        lr = group['lr']

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=lr,
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        updates = torch._foreach_lerp(exp_avgs, grads, weight=1.0 - beta1)
        torch._foreach_sign_(updates)

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - beta2)

        foreach_add_(params, updates, alpha=-lr)

    def _step_per_param(self, group: ParamGroup) -> None:
        beta1, beta2 = group['betas']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            exp_avg = state['exp_avg']

            p, grad, exp_avg = self.view_as_real(p, grad, exp_avg)

            if group.get('use_gc'):
                centralize_gradient(grad, gc_conv_only=False)

            self.apply_weight_decay(
                p=p,
                grad=grad,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=group['fixed_decay'],
            )

            s_grad = self.get_adanorm_gradient(
                grad=grad,
                adanorm=group.get('adanorm', False),
                exp_grad_norm=state.get('exp_grad_adanorm', None),
                r=group.get('adanorm_r', None),
            )

            update = exp_avg.clone()

            update.lerp_(grad, weight=1.0 - beta1).sign_()
            exp_avg.lerp_(s_grad, weight=1.0 - beta2)

            if group.get('cautious'):
                self.apply_cautious(update, grad)

            if isinstance(group['lr'], torch.Tensor):
                p.add_(update * -group['lr'])
            else:
                p.add_(update, alpha=-group['lr'])

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(group, self.state, state_keys=['exp_avg'])
                for tensors in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(group, tensors['params'], tensors['grads'], tensors['exp_avg'])
            else:
                self._step_per_param(group)

        return loss

load_ao_optimizer(optimizer)

Return an optimizer class from TorchAO.

Parameters:

Name Type Description Default
optimizer str

Lowercase optimizer name, including the integration prefix.

required

Returns:

Name Type Description
OptimizerType OptimizerType

Optimizer class from the optional integration.

Raises:

Type Description
NotImplementedError

The optimizer name is unsupported.

Source code in pytorch_optimizer/optimizer/__init__.py
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
def load_ao_optimizer(optimizer: str) -> OptimizerType:  # pragma: no cover
    """Return an optimizer class from TorchAO.

    Args:
        optimizer: Lowercase optimizer name, including the integration prefix.

    Returns:
        OptimizerType: Optimizer class from the optional integration.

    Raises:
        NotImplementedError: The optimizer name is unsupported.

    """
    from torchao.prototype import low_bit_optim  # noqa: PLC0415

    if 'adamw8bit' in optimizer:
        return low_bit_optim.AdamW8bit
    if 'adamw4bit' in optimizer:
        return low_bit_optim.AdamW4bit
    if 'adamwfp8' in optimizer:
        return low_bit_optim.AdamWFp8

    raise NotImplementedError(f'not implemented optimizer {optimizer}')

load_bnb_optimizer(optimizer)

Return an optimizer class from bitsandbytes.

Parameters:

Name Type Description Default
optimizer str

Lowercase optimizer name, including the integration prefix.

required

Returns:

Name Type Description
OptimizerType OptimizerType

Optimizer class from the optional integration.

Raises:

Type Description
NotImplementedError

The optimizer name is unsupported.

Source code in pytorch_optimizer/optimizer/__init__.py
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
def load_bnb_optimizer(optimizer: str) -> OptimizerType:  # pragma: no cover
    """Return an optimizer class from bitsandbytes.

    Args:
        optimizer: Lowercase optimizer name, including the integration prefix.

    Returns:
        OptimizerType: Optimizer class from the optional integration.

    Raises:
        NotImplementedError: The optimizer name is unsupported.

    """
    from bitsandbytes import optim  # noqa: PLC0415

    for name, cls_name in BNB_OPTIMIZERS:
        if name in optimizer:
            return getattr(optim, cls_name)

    raise NotImplementedError(f'not implemented optimizer {optimizer}')

load_optimizer(optimizer)

Return an optimizer class by name.

Names are case insensitive. Use the bnb, q_galore, or torchao prefix for optional integrations, which require their dependencies and CUDA.

Parameters:

Name Type Description Default
optimizer str

Registered optimizer name.

required

Returns:

Name Type Description
OptimizerType OptimizerType

Optimizer class to instantiate with parameters and options.

Raises:

Type Description
ImportError

An optional integration or CUDA is unavailable.

NotImplementedError

The optimizer name is unsupported.

Source code in pytorch_optimizer/optimizer/__init__.py
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
def load_optimizer(optimizer: str) -> OptimizerType:
    """Return an optimizer class by name.

    Names are case insensitive. Use the `bnb`, `q_galore`, or `torchao` prefix for
    optional integrations, which require their dependencies and CUDA.

    Args:
        optimizer: Registered optimizer name.

    Returns:
        OptimizerType: Optimizer class to instantiate with parameters and options.

    Raises:
        ImportError: An optional integration or CUDA is unavailable.
        NotImplementedError: The optimizer name is unsupported.

    """
    optimizer_name: str = optimizer.lower()

    if optimizer_name.startswith('bnb'):
        if HAS_BNB and torch.cuda.is_available():
            return load_bnb_optimizer(optimizer_name)  # pragma: no cover
        raise ImportError(f'bitsandbytes and CUDA required for the optimizer {optimizer_name}')
    if optimizer_name.startswith('q_galore'):
        if HAS_Q_GALORE and torch.cuda.is_available():
            return load_q_galore_optimizer(optimizer_name)  # pragma: no cover
        raise ImportError(f'bitsandbytes, q-galore-torch, and CUDA required for the optimizer {optimizer_name}')
    if optimizer_name.startswith('torchao'):
        if HAS_TORCHAO and torch.cuda.is_available():
            return load_ao_optimizer(optimizer_name)  # pragma: no cover
        raise ImportError(
            f'torchao required for the optimizer {optimizer_name}. '
            'usage: https://github.com/pytorch/ao/tree/main/torchao/prototype/low_bit_optim#usage'
        )
    if optimizer_name not in OPTIMIZERS:
        raise NotImplementedError(f'not implemented optimizer : {optimizer_name}')

    return OPTIMIZERS[optimizer_name]

load_q_galore_optimizer(optimizer)

Return an optimizer class from Q-GaLore.

Parameters:

Name Type Description Default
optimizer str

Lowercase optimizer name, including the integration prefix.

required

Returns:

Name Type Description
OptimizerType OptimizerType

Optimizer class from the optional integration.

Raises:

Type Description
NotImplementedError

The optimizer name is unsupported.

Source code in pytorch_optimizer/optimizer/__init__.py
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
def load_q_galore_optimizer(optimizer: str) -> OptimizerType:  # pragma: no cover
    """Return an optimizer class from Q-GaLore.

    Args:
        optimizer: Lowercase optimizer name, including the integration prefix.

    Returns:
        OptimizerType: Optimizer class from the optional integration.

    Raises:
        NotImplementedError: The optimizer name is unsupported.

    """
    import q_galore_torch  # noqa: PLC0415

    if 'adamw8bit' in optimizer:
        return q_galore_torch.QGaLoreAdamW8bit

    raise NotImplementedError(f'not implemented optimizer {optimizer}')

LOMO

Bases: BaseOptimizer

SGD updates fused into backward to reduce optimizer memory.

Reference: https://github.com/OpenLMLab/LOMO/blob/main/src/lomo.py Check usage: https://github.com/OpenLMLab/LOMO/blob/main/lomo/src/lomo_trainer.py

Parameters:

Name Type Description Default
model Module

PyTorch model.

required
lr float

Learning rate.

0.001
clip_grad_norm float | None

Gradient norm clipping value.

None
clip_grad_value float | None

Gradient value clipping threshold.

None
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/lomo.py
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
class LOMO(BaseOptimizer):
    """SGD updates fused into backward to reduce optimizer memory.

    Reference: https://github.com/OpenLMLab/LOMO/blob/main/src/lomo.py
    Check usage: https://github.com/OpenLMLab/LOMO/blob/main/lomo/src/lomo_trainer.py

    Args:
        model: PyTorch model.
        lr: Learning rate.
        clip_grad_norm: Gradient norm clipping value.
        clip_grad_value: Gradient value clipping threshold.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        model: nn.Module,
        lr: float = 1e-3,
        clip_grad_norm: float | None = None,
        clip_grad_value: float | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_non_negative(clip_grad_norm, 'clip_grad_norm')
        self.validate_non_negative(clip_grad_value, 'clip_grad_value')

        self.model = model
        self.lr = lr
        self.clip_grad_norm = clip_grad_norm
        self.clip_grad_value = clip_grad_value
        self.maximize = maximize

        self.local_rank: int = int(os.environ.get('LOCAL_RANK', '0'))

        self.gather_norm: bool = False
        self.grad_norms: list[torch.Tensor] = []
        self.clip_coef: float | torch.Tensor | None = None

        p0: torch.Tensor = next(iter(self.model.parameters()))

        self.grad_func: Callable[[Any], Any] = (
            self.fuse_update_zero3() if hasattr(p0, 'ds_tensor') else self.fuse_update()
        )

        self.loss_scaler: DynamicLossScaler | None = None
        if p0.dtype == torch.float16:
            if clip_grad_norm is None:
                raise ValueError('loss scaling is recommended to be used with grad norm to get better performance.')

            self.loss_scaler = DynamicLossScaler(init_scale=2 ** 16)  # fmt: skip

        for _, p in self.model.named_parameters():
            if p.requires_grad:
                p.register_hook(self.grad_func)

        defaults: Defaults = {'lr': lr}

        super().__init__(self.model.parameters(), defaults)

    def __str__(self) -> str:
        return 'LOMO'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

    def fuse_update(self) -> Callable[[Any], Any]:
        @torch.no_grad()
        def func(x: Any) -> Any:
            for _, p in self.model.named_parameters():
                if not p.requires_grad or p.grad is None:
                    continue

                if (self.loss_scaler and self.loss_scaler.has_overflow_serial) or has_overflow(p.grad):
                    p.grad = None
                    if self.loss_scaler is not None:
                        self.loss_scaler.has_overflow_serial = True
                    break

                grad_fp32 = p.grad.to(torch.float32)
                p.grad = None

                if self.loss_scaler:
                    grad_fp32.div_(self.loss_scaler.loss_scale)

                if self.gather_norm:
                    self.grad_norms.append(torch.norm(grad_fp32, 2.0))
                else:
                    if self.clip_grad_value is not None and self.clip_grad_value > 0.0:
                        grad_fp32.clamp_(min=-self.clip_grad_value, max=self.clip_grad_value)
                    if self.clip_grad_norm is not None and self.clip_grad_norm > 0.0 and self.clip_coef is not None:
                        grad_fp32.mul_(self.clip_coef)

                    self.maximize_gradient(grad_fp32, maximize=self.maximize)

                    p_fp32 = p.to(torch.float32)
                    p_fp32.add_(grad_fp32, alpha=-self.lr)
                    p.copy_(p_fp32)

            return x

        return func

    def fuse_update_zero3(self) -> Callable[[Any], Any]:  # pragma: no cover
        @torch.no_grad()
        def func(x: torch.Tensor) -> torch.Tensor:
            for _, p in self.model.named_parameters():
                if p.grad is None:
                    continue

                all_reduce(p.grad, op=ReduceOp.AVG, async_op=False)

                if (self.loss_scaler and self.loss_scaler.has_overflow_serial) or has_overflow(p.grad):
                    p.grad = None
                    if self.loss_scaler is not None:
                        self.loss_scaler.has_overflow_serial = True
                    break

                grad_fp32 = p.grad.to(torch.float32)
                p.grad = None

                param_fp32 = p.ds_tensor.to(torch.float32)
                if self.loss_scaler:
                    grad_fp32.div_(self.loss_scaler.loss_scale)

                if self.gather_norm:
                    self.grad_norms.append(torch.norm(grad_fp32, 2.0))
                else:
                    one_dim_grad_fp32 = grad_fp32.view(-1)

                    partition_size: int = p.ds_tensor.numel()
                    start: int = partition_size * self.local_rank
                    end: int = min(start + partition_size, grad_fp32.numel())

                    partitioned_grad_fp32 = one_dim_grad_fp32.narrow(0, start, end - start)

                    if self.clip_grad_value is not None:
                        partitioned_grad_fp32.clamp_(min=-self.clip_grad_value, max=self.clip_grad_value)

                    if self.clip_grad_norm is not None and self.clip_grad_norm > 0 and self.clip_coef is not None:
                        partitioned_grad_fp32.mul_(self.clip_coef)

                    self.maximize_gradient(partitioned_grad_fp32, maximize=self.maximize)

                    partitioned_p = param_fp32.narrow(0, 0, end - start)
                    partitioned_p.add_(partitioned_grad_fp32, alpha=-self.lr)

                    p.ds_tensor[: end - start] = partitioned_p  # fmt: skip

            return x

        return func

    def fused_backward(self, loss, lr: float):
        self.lr = lr

        if self.clip_grad_norm is not None and self.clip_grad_norm > 0.0 and self.clip_coef is None:
            raise ValueError(
                'clip_grad_norm is not None, but clip_coef is None. '
                'Please call optimizer.grad_norm() before optimizer.fused_backward().'
            )

        if self.loss_scaler:
            loss = loss * self.loss_scaler.loss_scale

        loss.backward()

        self.grad_func(0)

    def grad_norm(self, loss):
        self.gather_norm = True
        self.grad_norms = []

        if self.loss_scaler:
            self.loss_scaler.has_overflow_serial = False
            loss = loss * self.loss_scaler.loss_scale

        loss.backward(retain_graph=True)

        self.grad_func(0)

        if self.loss_scaler and self.loss_scaler.has_overflow_serial:
            self.loss_scaler.update_scale(overflow=True)

            with torch.no_grad():
                for _, p in self.model.named_parameters():
                    p.grad = None
            return

        with torch.no_grad():
            grad_norms = torch.stack(self.grad_norms)

            total_norm = torch.norm(grad_norms, 2.0)
            self.clip_coef = torch.clamp(float(self.clip_grad_norm) / (total_norm + 1e-6), max=1.0)

        self.gather_norm = False

Lookahead

Bases: BaseOptimizer

Wrap an optimizer with periodic interpolation toward slow weights.

Parameters:

Name Type Description Default
optimizer OptimizerInstanceOrClass

Base optimizer.

required
k int

Number of base optimizer steps between slow weight updates.

5
alpha float

Interpolation factor from slow weights toward fast weights.

0.5
pullback_momentum str

Momentum handling at interpolation: 'none', 'reset', or 'pullback'.

'none'
Source code in pytorch_optimizer/optimizer/lookahead.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
class Lookahead(BaseOptimizer):
    """Wrap an optimizer with periodic interpolation toward slow weights.

    Args:
        optimizer: Base optimizer.
        k: Number of base optimizer steps between slow weight updates.
        alpha: Interpolation factor from slow weights toward fast weights.
        pullback_momentum: Momentum handling at interpolation: `'none'`, `'reset'`, or `'pullback'`.

    """

    def __init__(
        self,
        optimizer: OptimizerInstanceOrClass,
        k: int = 5,
        alpha: float = 0.5,
        pullback_momentum: str = 'none',
        **kwargs,
    ) -> None:
        self.validate_positive(k, 'k')
        self.validate_range(alpha, 'alpha', 0.0, 1.0)
        self.validate_options(pullback_momentum, 'pullback_momentum', ['none', 'reset', 'pullback'])

        self.optimizer: Optimizer = self.load_optimizer(optimizer, **kwargs)

        self._optimizer_step_pre_hooks: dict[int, Callable] = OrderedDict()
        self._optimizer_step_post_hooks: dict[int, Callable] = OrderedDict()
        self._patch_step_function()

        self.alpha = alpha
        self.k = k
        self.pullback_momentum = pullback_momentum

        self.state: State = defaultdict(dict)

        for group in self.param_groups:
            self.init_group(group)

        self.defaults: Defaults = {
            'lookahead_alpha': alpha,
            'lookahead_k': k,
            'lookahead_pullback_momentum': pullback_momentum,
            **self.optimizer.defaults,
        }

    @property
    def param_groups(self):
        return self.optimizer.param_groups

    def __getstate__(self):
        return {
            'state': self.state,
            'optimizer': self.optimizer,
            'alpha': self.alpha,
            'k': self.k,
            'pullback_momentum': self.pullback_momentum,
        }

    @torch.no_grad()
    def zero_grad(self, set_to_none: bool = True) -> None:
        self.optimizer.zero_grad(set_to_none=set_to_none)

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        group.setdefault('counter', 0)
        for p in group.get('params', []):
            state = self.state[p]
            if 'slow_params' not in state:
                state['slow_params'] = p.detach().clone()
                if self.pullback_momentum == 'pullback':
                    state['slow_momentum'] = torch.zeros_like(p)

    def add_param_group(self, param_group: ParamGroup) -> None:
        self.optimizer.add_param_group(param_group)
        self.init_group(self.param_groups[-1])

    def backup_and_load_cache(self) -> None:
        """Back up fast weights and load slow weights for evaluation."""
        for group in self.param_groups:
            for p in group['params']:
                state = self.state[p]
                state['backup_params'] = p.detach().clone()
                p.data.copy_(state['slow_params'])

    def clear_and_load_backup(self) -> None:
        """Restore fast weights after evaluating slow weights."""
        for group in self.param_groups:
            for p in group['params']:
                state = self.state[p]
                p.data.copy_(state['backup_params'])
                del state['backup_params']

    def state_dict(self) -> State:
        lookahead_state: State = {
            (group_index, parameter_index): dict(self.state[p])
            for group_index, group in enumerate(self.param_groups)
            for parameter_index, p in enumerate(group['params'])
            if p in self.state
        }
        return {'lookahead_state': lookahead_state, 'base_optimizer': self.optimizer.state_dict()}

    def load_state_dict(self, state: State) -> None:
        """Restore optimizer state and slow weights from a checkpoint."""
        saved_state = state['lookahead_state']
        restored_state: State = {}
        for group_index, group in enumerate(self.param_groups):
            for parameter_index, p in enumerate(group['params']):
                key = (group_index, parameter_index)
                if key in saved_state:
                    restored_state[p] = dict(saved_state[key])
                elif p in saved_state:
                    restored_state[p] = dict(saved_state[p])
        parameter_count = sum(len(group['params']) for group in self.param_groups)
        if len(restored_state) != len(saved_state) or len(restored_state) != parameter_count:
            raise ValueError('lookahead state does not match the current parameters')

        self.optimizer.load_state_dict(state['base_optimizer'])
        for p, parameter_state in restored_state.items():
            for key, value in parameter_state.items():
                if isinstance(value, torch.Tensor):
                    parameter_state[key] = value.to(device=p.device, dtype=p.dtype).clone()
        self.state = defaultdict(dict, restored_state)

    @torch.no_grad()
    def update(self, group: dict):
        for p in group['params']:
            if p.grad is None:
                continue

            state = self.state[p]

            slow = state['slow_params']

            p.lerp_(slow, weight=1.0 - self.alpha)
            slow.copy_(p)

            if self.pullback_momentum == 'pullback':
                if 'momentum_buffer' not in self.optimizer.state[p]:
                    self.optimizer.state[p]['momentum_buffer'] = torch.zeros_like(p)

                internal_momentum = self.optimizer.state[p]['momentum_buffer']
                internal_momentum.lerp_(state['slow_momentum'], weight=1.0 - self.alpha)
                state['slow_momentum'].copy_(internal_momentum)
            elif self.pullback_momentum == 'reset':
                self.optimizer.state[p]['momentum_buffer'] = torch.zeros_like(p)

    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = self.optimizer.step(closure)
        for group in self.param_groups:
            group['counter'] += 1
            if group['counter'] >= self.k:
                group['counter'] = 0
                self.update(group)
        return loss

backup_and_load_cache()

Back up fast weights and load slow weights for evaluation.

Source code in pytorch_optimizer/optimizer/lookahead.py
86
87
88
89
90
91
92
def backup_and_load_cache(self) -> None:
    """Back up fast weights and load slow weights for evaluation."""
    for group in self.param_groups:
        for p in group['params']:
            state = self.state[p]
            state['backup_params'] = p.detach().clone()
            p.data.copy_(state['slow_params'])

clear_and_load_backup()

Restore fast weights after evaluating slow weights.

Source code in pytorch_optimizer/optimizer/lookahead.py
 94
 95
 96
 97
 98
 99
100
def clear_and_load_backup(self) -> None:
    """Restore fast weights after evaluating slow weights."""
    for group in self.param_groups:
        for p in group['params']:
            state = self.state[p]
            p.data.copy_(state['backup_params'])
            del state['backup_params']

load_state_dict(state)

Restore optimizer state and slow weights from a checkpoint.

Source code in pytorch_optimizer/optimizer/lookahead.py
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
def load_state_dict(self, state: State) -> None:
    """Restore optimizer state and slow weights from a checkpoint."""
    saved_state = state['lookahead_state']
    restored_state: State = {}
    for group_index, group in enumerate(self.param_groups):
        for parameter_index, p in enumerate(group['params']):
            key = (group_index, parameter_index)
            if key in saved_state:
                restored_state[p] = dict(saved_state[key])
            elif p in saved_state:
                restored_state[p] = dict(saved_state[p])
    parameter_count = sum(len(group['params']) for group in self.param_groups)
    if len(restored_state) != len(saved_state) or len(restored_state) != parameter_count:
        raise ValueError('lookahead state does not match the current parameters')

    self.optimizer.load_state_dict(state['base_optimizer'])
    for p, parameter_state in restored_state.items():
        for key, value in parameter_state.items():
            if isinstance(value, torch.Tensor):
                parameter_state[key] = value.to(device=p.device, dtype=p.dtype).clone()
    self.state = defaultdict(dict, restored_state)

LookSAM

Bases: BaseOptimizer

Sharpness-aware minimization with periodic perturbation updates.

Compute gradients at the current weights before calling step(). The closure must recompute the loss and gradients at the perturbed weights.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
base_optimizer OptimizerType

Optimizer class to instantiate for the parameter update.

required
rho float

Radius of the neighborhood used to perturb parameters.

0.1
k int

Number of steps between full sharpness gradient updates.

10
alpha float

Weight of the reused orthogonal sharpness gradient.

0.7
use_gc bool

Centralize gradients before perturbing parameters.

False
adaptive bool

Scale perturbations by the squared parameter values.

False
perturb_eps float

Stability constant for the perturbation norm.

1e-12
**kwargs dict

Options for the base optimizer.

{}

Examples:

optimizer = LookSAM(model.parameters(), torch.optim.AdamW, lr=1e-3)
for inputs, targets in data:
    optimizer.zero_grad()

    def closure():
        optimizer.zero_grad()
        loss = loss_fn(model(inputs), targets)
        loss.backward()
        return loss

    closure()
    optimizer.step(closure)
Source code in pytorch_optimizer/optimizer/sam.py
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
class LookSAM(BaseOptimizer):
    """Sharpness-aware minimization with periodic perturbation updates.

    Compute gradients at the current weights before calling `step()`. The closure
    must recompute the loss and gradients at the perturbed weights.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        base_optimizer: Optimizer class to instantiate for the parameter update.
        rho: Radius of the neighborhood used to perturb parameters.
        k: Number of steps between full sharpness gradient updates.
        alpha: Weight of the reused orthogonal sharpness gradient.
        use_gc: Centralize gradients before perturbing parameters.
        adaptive: Scale perturbations by the squared parameter values.
        perturb_eps: Stability constant for the perturbation norm.
        **kwargs (dict): Options for the base optimizer.

    Examples:
        ```python
        optimizer = LookSAM(model.parameters(), torch.optim.AdamW, lr=1e-3)
        for inputs, targets in data:
            optimizer.zero_grad()

            def closure():
                optimizer.zero_grad()
                loss = loss_fn(model(inputs), targets)
                loss.backward()
                return loss

            closure()
            optimizer.step(closure)
        ```

    """

    def __init__(
        self,
        params: ParamsT,
        base_optimizer: OptimizerType,
        rho: float = 0.1,
        k: int = 10,
        alpha: float = 0.7,
        adaptive: bool = False,
        use_gc: bool = False,
        perturb_eps: float = 1e-12,
        **kwargs,
    ):
        self.validate_non_negative(rho, 'rho')
        self.validate_positive(k, 'k')
        self.validate_range(alpha, 'alpha', 0.0, 1.0, '()')
        self.validate_non_negative(perturb_eps, 'perturb_eps')

        self.k = k
        self.alpha = alpha
        self.use_gc = use_gc
        self.perturb_eps = perturb_eps

        defaults: Defaults = {'rho': rho, 'adaptive': adaptive}
        defaults.update(kwargs)

        super().__init__(params, defaults)

        self.base_optimizer: Optimizer = base_optimizer(self.param_groups, **kwargs)
        self.param_groups = self.base_optimizer.param_groups

    def __str__(self) -> str:
        return 'LookSAM'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        pass

    def get_step(self):
        return (
            self.param_groups[0]['step']
            if 'step' in self.param_groups[0]
            else next(iter(self.base_optimizer.state.values()))['step'] if self.base_optimizer.state else 0
        )

    @torch.no_grad()
    def first_step(self, zero_grad: bool = False) -> None:
        if self.get_step() % self.k != 0:
            return

        grad_norm = get_global_gradient_norm(self.param_groups, weight_adaptive=True)
        grad_norm.sqrt_().squeeze_(0).add_(self.perturb_eps)

        for group in self.param_groups:
            scale = group['rho'] / grad_norm

            for p in group['params']:
                self.state[p].pop('old_grad_p', None)
                if p.grad is None:
                    continue

                grad = p.grad
                if self.use_gc:
                    centralize_gradient(grad, gc_conv_only=False)

                self.state[p]['old_p'] = p.clone()
                self.state[p]['old_grad_p'] = grad.clone()

                e_w = (torch.pow(p, 2) if group['adaptive'] else 1.0) * grad * scale.to(p)

                p.add_(e_w)

        if zero_grad:
            self.zero_grad()

    @torch.no_grad()
    def second_step(self, zero_grad: bool = False):
        step = self.get_step()

        for group in self.param_groups:
            for p in group['params']:
                if 'old_p' in self.state[p]:
                    p.copy_(self.state[p].pop('old_p'))
                old_grad_p = self.state[p].pop('old_grad_p', None)
                if p.grad is None:
                    continue

                grad = p.grad
                grad_norm = grad.norm(p=2)

                if step % self.k == 0 and old_grad_p is not None:
                    g_grad_norm = old_grad_p / old_grad_p.norm(p=2).clamp_min(self.perturb_eps)
                    g_s_grad_norm = grad / grad_norm.clamp_min(self.perturb_eps)

                    self.state[p]['gv'] = torch.sub(
                        grad, grad_norm * torch.sum(g_grad_norm * g_s_grad_norm) * g_grad_norm
                    )
                elif step % self.k != 0 and 'gv' in self.state[p]:
                    gv = self.state[p]['gv']
                    grad.add_(grad_norm / (gv.norm(p=2) + 1e-8) * gv, alpha=self.alpha)

        self.base_optimizer.step()

        if zero_grad:
            self.zero_grad()

    @torch.no_grad()
    def step(self, closure: Closure = None):
        """Perturb weights, recompute gradients, and apply the base optimizer update.

        Args:
            closure: Callable that clears gradients and recomputes the loss and gradients. Compute the initial
                gradients before calling this method.

        Raises:
            NoClosureError: No closure is supplied.

        """
        if closure is None:
            raise NoClosureError(str(self))

        self.first_step(zero_grad=True)

        with torch.enable_grad():
            closure()

        self.second_step()

    def state_dict(self) -> dict:
        state = super().state_dict()
        state['base_optimizer'] = self.base_optimizer.state_dict()
        return state

    def load_state_dict(self, state_dict: dict):
        super().load_state_dict(state_dict)
        if 'base_optimizer' in state_dict:
            self.base_optimizer.load_state_dict(state_dict['base_optimizer'])
            self.param_groups = self.base_optimizer.param_groups
        else:
            self.base_optimizer.param_groups = self.param_groups

step(closure=None)

Perturb weights, recompute gradients, and apply the base optimizer update.

Parameters:

Name Type Description Default
closure Closure

Callable that clears gradients and recomputes the loss and gradients. Compute the initial gradients before calling this method.

None

Raises:

Type Description
NoClosureError

No closure is supplied.

Source code in pytorch_optimizer/optimizer/sam.py
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
@torch.no_grad()
def step(self, closure: Closure = None):
    """Perturb weights, recompute gradients, and apply the base optimizer update.

    Args:
        closure: Callable that clears gradients and recomputes the loss and gradients. Compute the initial
            gradients before calling this method.

    Raises:
        NoClosureError: No closure is supplied.

    """
    if closure is None:
        raise NoClosureError(str(self))

    self.first_step(zero_grad=True)

    with torch.enable_grad():
        closure()

    self.second_step()

LoRARite

Bases: BaseOptimizer

LoRA factor optimization with matrix preconditioning and basis corrections.

This optimizer expects LoRA factors in alternating order, such as lora_a_1, lora_b_1, lora_a_2, lora_b_2. Unpaired parameters and pairs with missing gradients are skipped, matching common fine tuning workflows where only part of the model may receive gradients on a given step.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Coefficients used for first moment and matrix second moment estimates.

(0.9, 0.999)
eps float

Term added to the denominator to improve numerical stability.

1e-06
relative_epsilon bool

Scale the root epsilon by the largest matrix second moment eigenvalue.

False
clip_unmagnified_grad float

Global clipping threshold for unmagnified LoRA gradients. Disabled when 0.

1.0
update_capping float

Per update RMS capping threshold after preconditioning. Disabled when 0.

0.0
update_skipping float

Skip unmagnified updates whose RMS is above this threshold. Disabled when 0.

1.0
weight_decay float

Weight decay coefficient.

0.0
apply_escape bool

Apply the RITE escape correction when rotating second moment bases.

False
lora_l_dim int

LoRA rank dimension for left factors.

0
lora_r_dim int

LoRA rank dimension for right factors.

-1
maybe_inf_to_nan bool

Convert infinite update statistics to NaN before threshold checks.

True
balance_param bool

Balance the norms of each LoRA factor pair after applying the update.

False
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/lora_rite.py
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
class LoRARite(BaseOptimizer):
    """LoRA factor optimization with matrix preconditioning and basis corrections.

    This optimizer expects LoRA factors in alternating order, such as `lora_a_1, lora_b_1, lora_a_2, lora_b_2`.
    Unpaired parameters and pairs with missing gradients are skipped, matching common fine tuning workflows where only
    part of the model may receive gradients on a given step.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Coefficients used for first moment and matrix second moment estimates.
        eps: Term added to the denominator to improve numerical stability.
        relative_epsilon: Scale the root epsilon by the largest matrix second moment eigenvalue.
        clip_unmagnified_grad: Global clipping threshold for unmagnified LoRA gradients. Disabled when 0.
        update_capping: Per update RMS capping threshold after preconditioning. Disabled when 0.
        update_skipping: Skip unmagnified updates whose RMS is above this threshold. Disabled when 0.
        weight_decay: Weight decay coefficient.
        apply_escape: Apply the RITE escape correction when rotating second moment bases.
        lora_l_dim: LoRA rank dimension for left factors.
        lora_r_dim: LoRA rank dimension for right factors.
        maybe_inf_to_nan: Convert infinite update statistics to NaN before threshold checks.
        balance_param: Balance the norms of each LoRA factor pair after applying the update.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        eps: float = 1e-6,
        relative_epsilon: bool = False,
        clip_unmagnified_grad: float = 1.0,
        update_capping: float = 0.0,
        update_skipping: float = 1.0,
        weight_decay: float = 0.0,
        apply_escape: bool = False,
        lora_l_dim: int = 0,
        lora_r_dim: int = -1,
        maybe_inf_to_nan: bool = True,
        balance_param: bool = False,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(eps, 'eps')
        self.validate_non_negative(clip_unmagnified_grad, 'clip_unmagnified_grad')
        self.validate_non_negative(update_capping, 'update_capping')
        self.validate_non_negative(update_skipping, 'update_skipping')
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.helper = LoRARiteHelper(maybe_inf_to_nan=maybe_inf_to_nan)
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'eps': eps,
            'eps_root': eps**2,
            'relative_epsilon': relative_epsilon,
            'clip_unmagnified_grad': clip_unmagnified_grad,
            'update_capping': update_capping,
            'update_skipping': update_skipping,
            'weight_decay': weight_decay,
            'apply_escape': apply_escape,
            'lora_l_dim': lora_l_dim,
            'lora_r_dim': lora_r_dim,
            'balance_param': balance_param,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'LoRARite'

    @staticmethod
    def iter_lora_pairs(group: ParamGroup) -> list[tuple[torch.Tensor, torch.Tensor]]:
        params = list(group['params'])
        return list(zip(params[::2], params[1::2]))

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for param in group['params']:
            if param.grad is None:
                continue

            if param.grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(param):
                raise NoComplexParameterError(str(self))

    def init_pair_state(
        self, group: ParamGroup, state: dict[str, Any], param_left: torch.Tensor, param_right: torch.Tensor
    ) -> None:
        if 'step' in state:
            return

        param_left, _ = self.helper.move_lora_dim_to_last(param_left, group['lora_l_dim'])
        param_right, _ = self.helper.move_lora_dim_to_last(param_right, group['lora_r_dim'])

        state['step'] = 0
        state['v_l'] = self.helper.create_preconditioner(param_left)
        state['v_r'] = self.helper.create_preconditioner(param_right)
        state['m_l'] = torch.zeros_like(param_left)
        state['m_r'] = torch.zeros_like(param_right)
        state['basis_l'] = torch.zeros_like(param_left)
        state['basis_r'] = torch.zeros_like(param_right)
        state['escape_l'] = torch.zeros((), dtype=param_left.dtype, device=param_left.device)
        state['escape_r'] = torch.zeros((), dtype=param_right.dtype, device=param_right.device)

    def build_pair_info(
        self,
        group: ParamGroup,
        param_left: torch.Tensor,
        param_right: torch.Tensor,
    ) -> dict[str, Any]:
        helper = self.helper
        state = self.state[param_left]
        self.init_pair_state(group, state, param_left, param_right)

        param_left_2d, _ = helper.move_lora_dim_to_last(param_left, group['lora_l_dim'])
        param_right_2d, _ = helper.move_lora_dim_to_last(param_right, group['lora_r_dim'])

        basis_left, rotate_left = helper.get_rotation_and_basis(param_left_2d)
        basis_right, rotate_right = helper.get_rotation_and_basis(param_right_2d)
        rotate_inv_left = torch.linalg.pinv(rotate_left)
        rotate_inv_right = torch.linalg.pinv(rotate_right)

        projection_left = basis_right.mT @ state['basis_r']
        projection_right = basis_left.mT @ state['basis_l']

        grad_left = helper.inf_to_nan(param_left.grad.detach())
        grad_right = helper.inf_to_nan(param_right.grad.detach())
        if self.maximize:
            grad_left = grad_left.neg()
            grad_right = grad_right.neg()

        grad_left, _ = helper.move_lora_dim_to_last(grad_left, group['lora_l_dim'])
        grad_right, _ = helper.move_lora_dim_to_last(grad_right, group['lora_r_dim'])

        update_left = helper.get_unmagnified_grad(grad_left, rotate_inv_right)
        update_right = helper.get_unmagnified_grad(grad_right, rotate_inv_left)

        if group['update_skipping'] > 0.0:
            update_left = helper.skip_update(update_left, group['update_skipping'])
            update_right = helper.skip_update(update_right, group['update_skipping'])

        state['basis_l'] = basis_left
        state['basis_r'] = basis_right
        state['rotate_inv_l'] = rotate_inv_left
        state['rotate_inv_r'] = rotate_inv_right
        state['update_l'] = update_left
        state['update_r'] = update_right
        state['projection_l'] = projection_left
        state['projection_r'] = projection_right

        return state

    def apply_pair_update(
        self,
        group: ParamGroup,
        param_left: torch.Tensor,
        param_right: torch.Tensor,
        grad_norm: torch.Tensor,
    ) -> None:
        helper = self.helper
        state = self.state[param_left]
        update_left = state.pop('update_l')
        update_right = state.pop('update_r')
        rotate_inv_left = state.pop('rotate_inv_l')
        rotate_inv_right = state.pop('rotate_inv_r')
        projection_left = state.pop('projection_l')
        projection_right = state.pop('projection_r')
        beta1, beta2 = group['betas']

        param_left_2d, _ = helper.move_lora_dim_to_last(param_left, group['lora_l_dim'])
        param_right_2d, _ = helper.move_lora_dim_to_last(param_right, group['lora_r_dim'])

        if group['clip_unmagnified_grad'] > 0.0 and grad_norm > group['clip_unmagnified_grad']:
            scale = group['clip_unmagnified_grad'] / grad_norm
            update_left = update_left * scale
            update_right = update_right * scale

        second_left = helper.compute_second_moment(update_left)
        second_right = helper.compute_second_moment(update_right)

        transformed_v_left = helper.transform_second_moment_to_new_basis(state['v_l'], projection_left)
        transformed_v_right = helper.transform_second_moment_to_new_basis(state['v_r'], projection_right)

        if group['apply_escape']:
            escape_left = helper.get_unmagnified_rotate_second_escape(transformed_v_left, state['v_l'])
            escape_right = helper.get_unmagnified_rotate_second_escape(transformed_v_right, state['v_r'])
            escape_left = helper.update_second_escape(
                state['step'], torch.zeros_like(escape_left), escape_left + state['escape_l'], beta2
            )
            escape_right = helper.update_second_escape(
                state['step'], torch.zeros_like(escape_right), escape_right + state['escape_r'], beta2
            )
        else:
            escape_left = torch.zeros((), dtype=param_left_2d.dtype, device=param_left_2d.device)
            escape_right = torch.zeros((), dtype=param_right_2d.dtype, device=param_right_2d.device)

        v_left = helper.update_second_moment(state['step'], second_left, transformed_v_left, beta2)
        v_right = helper.update_second_moment(state['step'], second_right, transformed_v_right, beta2)

        update_left = helper.get_preconditioned_update(
            update_left,
            v_left,
            escape_left,
            group['eps'],
            group['eps_root'],
            group['relative_epsilon'],
            group['apply_escape'],
        )
        update_right = helper.get_preconditioned_update(
            update_right,
            v_right,
            escape_right,
            group['eps'],
            group['eps_root'],
            group['relative_epsilon'],
            group['apply_escape'],
        )

        m_left = helper.transform_first_moment_to_new_basis(state['m_l'], projection_left)
        m_right = helper.transform_first_moment_to_new_basis(state['m_r'], projection_right)
        m_left = helper.update_first_moment(state['step'], update_left, m_left, beta1)
        m_right = helper.update_first_moment(state['step'], update_right, m_right, beta1)

        if group['update_capping'] > 0.0:
            m_left = helper.clip_update(m_left, group['update_capping'])
            m_right = helper.clip_update(m_right, group['update_capping'])

        update_left = helper.rotate_update(m_left, rotate_inv_right)
        update_right = helper.rotate_update(m_right, rotate_inv_left)

        if group['weight_decay'] > 0.0:
            update_left = update_left.add(param_left_2d, alpha=group['weight_decay'])
            update_right = update_right.add(param_right_2d, alpha=group['weight_decay'])

        update_left = update_left.mul(-group['lr'])
        update_right = update_right.mul(-group['lr'])

        if group['balance_param']:
            left_norm = torch.linalg.norm(param_left_2d + update_left).add_(1e-6)
            right_norm = torch.linalg.norm(param_right_2d + update_right).add_(1e-6)
            balanced_norm = torch.sqrt(left_norm * right_norm)
            update_left = update_left * (balanced_norm / left_norm) + param_left_2d * (balanced_norm / left_norm - 1.0)
            update_right = update_right * (balanced_norm / right_norm) + param_right_2d * (
                balanced_norm / right_norm - 1.0
            )

        param_left.add_(helper.restore_param_shape(update_left, param_left, group['lora_l_dim']).to(param_left.dtype))
        param_right.add_(
            helper.restore_param_shape(update_right, param_right, group['lora_r_dim']).to(param_right.dtype)
        )

        state['step'] += 1
        state['v_l'] = v_left
        state['v_r'] = v_right
        state['m_l'] = m_left
        state['m_r'] = m_right
        state['escape_l'] = escape_left
        state['escape_r'] = escape_right

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        pair_infos: list[tuple[ParamGroup, torch.Tensor, torch.Tensor]] = []
        grad_norm_sq: torch.Tensor | None = None

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            for param_left, param_right in self.iter_lora_pairs(group):
                if param_left.grad is None or param_right.grad is None:
                    continue

                state = self.build_pair_info(group, param_left, param_right)
                update_left, update_right = state['update_l'], state['update_r']
                if grad_norm_sq is None:
                    grad_norm_sq = update_left.new_zeros(())

                grad_norm_sq.add_(torch.linalg.norm(update_left).pow(2).to(grad_norm_sq.device))
                grad_norm_sq.add_(torch.linalg.norm(update_right).pow(2).to(grad_norm_sq.device))
                pair_infos.append((group, param_left, param_right))

        grad_norm = torch.sqrt(grad_norm_sq) if grad_norm_sq is not None else torch.zeros(())
        for group, param_left, param_right in pair_infos:
            self.apply_pair_update(group, param_left, param_right, grad_norm)

        return loss

MADGRAD

Bases: BaseOptimizer

Momentumized adaptive dual averaged gradient descent.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
momentum float

Interpolation factor toward the previous parameter values. 0 disables momentum.

0.9
eps float

Term added to the denominator to improve numerical stability.

1e-06
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/madgrad.py
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
class MADGRAD(BaseOptimizer):
    """Momentumized adaptive dual averaged gradient descent.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Interpolation factor toward the previous parameter values. `0` disables momentum.
        eps: Term added to the denominator to improve numerical stability.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        momentum: float = 0.9,
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        eps: float = 1e-6,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(momentum, 'momentum', 0.0, 1.0)
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'momentum': momentum,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'MADGRAD'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if group['momentum'] > 0.0 and grad.is_sparse:
                raise NoSparseGradientError(str(self), note='momentum > 0.0')

            if group['weight_decay'] > 0.0 and not group['weight_decouple'] and grad.is_sparse:
                raise NoSparseGradientError(str(self), note='weight_decay')

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if 'grad_sum_sq' not in state:
                state['grad_sum_sq'] = torch.zeros_like(p)
                state['s'] = torch.zeros_like(p)

                if group['momentum'] > 0.0:
                    state['x0'] = p.clone()

    @staticmethod
    def compute_rms(grad_sum_sq: torch.Tensor, eps: float) -> torch.Tensor:
        """Compute the cube root accumulator, treating zero denominators as inactive coordinates."""
        rms = grad_sum_sq.pow(1.0 / 3.0).add_(eps)
        if eps == 0.0:
            rms[rms == 0] = float('inf')

        return rms

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        if 'k' not in self.state:
            self.state['k'] = torch.tensor([0], dtype=torch.long, requires_grad=False)

        for group in self.param_groups:
            self.init_group(group)

            weight_decay, momentum, eps = group['weight_decay'], group['momentum'], group['eps']
            lr: float = group['lr'] + eps if group['lr'] != 0.0 else 0.0

            _lambda = lr * math.pow(self.state['k'] + 1, 0.5)

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                grad_sum_sq, s = state['grad_sum_sq'], state['s']
                if weight_decay > 0.0 and not group['weight_decouple']:
                    grad.add_(p, alpha=weight_decay)

                if grad.is_sparse:
                    grad = grad.coalesce()

                    p_masked = p.sparse_mask(grad)
                    grad_sum_sq_masked = grad_sum_sq.sparse_mask(grad)
                    s_masked = s.sparse_mask(grad)

                    rms_masked_values = self.compute_rms(grad_sum_sq_masked._values(), eps)
                    x0_masked_values = p_masked._values().addcdiv(s_masked._values(), rms_masked_values, value=1)

                    grad_sq = grad * grad
                    grad_sum_sq.add_(grad_sq, alpha=_lambda)
                    grad_sum_sq_masked.add_(grad_sq, alpha=_lambda)

                    rms_masked_values = self.compute_rms(grad_sum_sq_masked._values(), eps)

                    s.add_(grad, alpha=_lambda)
                    s_masked._values().add_(grad._values(), alpha=_lambda)

                    p_kp1_masked_values = x0_masked_values.addcdiv(s_masked._values(), rms_masked_values, value=-1)

                    p_masked._values().add_(p_kp1_masked_values, alpha=-1)
                    p.data.add_(p_masked, alpha=-1)
                else:
                    if momentum == 0.0:
                        rms = self.compute_rms(grad_sum_sq, eps)
                        x0 = p.addcdiv(s, rms, value=1)
                    else:
                        x0 = state['x0']

                    grad_sum_sq.addcmul_(grad, grad, value=_lambda)
                    rms = self.compute_rms(grad_sum_sq, eps)

                    s.add_(grad, alpha=_lambda)

                    p_old: torch.Tensor | None = None
                    if weight_decay > 0.0 and group['weight_decouple']:
                        p_old = p.clone()

                    if momentum == 0.0:
                        p.copy_(x0.addcdiv(s, rms, value=-1))
                    else:
                        z = x0.addcdiv(s, rms, value=-1)
                        p.lerp_(z, weight=1.0 - momentum)

                    if weight_decay > 0.0 and group['weight_decouple']:
                        p.add_(p_old, alpha=-lr * weight_decay)

        self.state['k'].add_(1)

        return loss

compute_rms(grad_sum_sq, eps) staticmethod

Compute the cube root accumulator, treating zero denominators as inactive coordinates.

Source code in pytorch_optimizer/optimizer/madgrad.py
87
88
89
90
91
92
93
94
@staticmethod
def compute_rms(grad_sum_sq: torch.Tensor, eps: float) -> torch.Tensor:
    """Compute the cube root accumulator, treating zero denominators as inactive coordinates."""
    rms = grad_sum_sq.pow(1.0 / 3.0).add_(eps)
    if eps == 0.0:
        rms[rms == 0] = float('inf')

    return rms

Magma

Bases: BaseOptimizer

Momentum aligned gradient masking wrapper for PyTorch optimizers.

Magma applies a block wise Bernoulli mask to the updates generated by a base optimizer. Surviving updates are scaled by an exponential moving average of the cosine similarity between the current gradient and the first moment estimate. The base optimizer's state is updated densely, including when a parameter update is masked.

Parameters:

Name Type Description Default
optimizer OptimizerInstanceOrClass | ParamsT

Base optimizer instance/class, or parameters to optimize with AdamW.

required
mask_prob float

Probability of keeping an update.

0.5
tau float

Temperature used by the alignment sigmoid.

2.0
momentum_beta float

EMA coefficient for the fallback momentum.

0.9
alignment_ema float

EMA coefficient for the alignment score.

0.9
moment_key str | None

First moment key in the base optimizer's state. 'auto' checks 'exp_avg' and 'momentum_buffer'. None uses Magma's fallback momentum.

'auto'
exclude set[Tensor] | None

Parameters that bypass masking.

None

Magma reads the first moment from the base optimizer when it is available, so it adds no additional momentum state for optimizers such as Adam. For optimizers without a first moment buffer, it maintains an EMA controlled by momentum_beta.

Reference

https://arxiv.org/abs/2602.15322

Source code in pytorch_optimizer/optimizer/magma.py
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
class Magma(BaseOptimizer):
    """Momentum aligned gradient masking wrapper for PyTorch optimizers.

    Magma applies a block wise Bernoulli mask to the updates generated by a
    base optimizer. Surviving updates are scaled by an exponential moving
    average of the cosine similarity between the current gradient and the
    first moment estimate. The base optimizer's state is updated densely,
    including when a parameter update is masked.

    Args:
        optimizer: Base optimizer instance/class, or parameters to optimize with AdamW.
        mask_prob: Probability of keeping an update.
        tau: Temperature used by the alignment sigmoid.
        momentum_beta: EMA coefficient for the fallback momentum.
        alignment_ema: EMA coefficient for the alignment score.
        moment_key: First moment key in the base optimizer's state. `'auto'` checks `'exp_avg'` and
            `'momentum_buffer'`. `None` uses Magma's fallback momentum.
        exclude: Parameters that bypass masking.

    Magma reads the first moment from the base optimizer when it is available,
    so it adds no additional momentum state for optimizers such as Adam. For
    optimizers without a first moment buffer, it maintains an EMA controlled
    by `momentum_beta`.

    Reference:
        https://arxiv.org/abs/2602.15322

    """

    def __init__(
        self,
        optimizer: OptimizerInstanceOrClass | ParamsT,
        mask_prob: float = 0.5,
        tau: float = 2.0,
        momentum_beta: float = 0.9,
        alignment_ema: float = 0.9,
        moment_key: str | None = 'auto',
        exclude: set[Tensor] | None = None,
        **kwargs,
    ) -> None:
        self.validate_range(mask_prob, 'mask_prob', 0.0, 1.0, range_type='[]')
        self.validate_positive(tau, 'tau')
        self.validate_range(momentum_beta, 'momentum_beta', 0.0, 1.0, range_type='[]')
        self.validate_range(alignment_ema, 'alignment_ema', 0.0, 1.0, range_type='[]')

        self._optimizer_step_pre_hooks: dict[int, Callable] = OrderedDict()
        self._optimizer_step_post_hooks: dict[int, Callable] = OrderedDict()
        self._patch_step_function()

        if isinstance(optimizer, Optimizer):
            self.optimizer = optimizer
        elif isinstance(optimizer, type) and issubclass(optimizer, Optimizer):
            self.optimizer = self.load_optimizer(optimizer, **kwargs)
        else:
            self.validate_learning_rate(kwargs.get('lr', 1e-3))
            for option in ('alpha', 'k', 'num_iterations', 'pullback_momentum', 'use_muon'):
                kwargs.pop(option, None)
            self.optimizer = self.load_optimizer(AdamW, params=optimizer, **kwargs)

        self.mask_prob = mask_prob
        self.tau = tau
        self.momentum_beta = momentum_beta
        self.alignment_ema = alignment_ema
        self.moment_key = moment_key
        self._exclude_ids: set[int] = {id(parameter) for parameter in (exclude or set())}
        self._state: dict[int, dict[str, Tensor]] = {}

        self.defaults: Defaults = self.optimizer.defaults

    def __str__(self) -> str:
        return 'Magma'

    @property
    def param_groups(self):
        return self.optimizer.param_groups

    @property
    def state(self) -> State:
        return self.optimizer.state

    def add_param_group(self, param_group: ParamGroup) -> None:
        self.optimizer.add_param_group(param_group)

    def state_dict(self) -> State:
        id_to_key: dict[int, tuple[int, int]] = {
            id(parameter): (group_index, parameter_index)
            for group_index, group in enumerate(self.param_groups)
            for parameter_index, parameter in enumerate(group['params'])
        }

        magma_state: dict[tuple[int, int], dict[str, Tensor]] = {}
        for parameter_id, parameter_state in self._state.items():
            key = id_to_key.get(parameter_id)
            if key is not None:
                magma_state[key] = {name: value.clone() for name, value in parameter_state.items()}

        return {
            'base': self.optimizer.state_dict(),
            'magma_state': magma_state,
            'mask_prob': self.mask_prob,
            'tau': self.tau,
            'momentum_beta': self.momentum_beta,
            'alignment_ema': self.alignment_ema,
            'moment_key': self.moment_key,
        }

    def load_state_dict(self, state_dict: State) -> None:
        self.optimizer.load_state_dict(state_dict['base'])

        self.mask_prob = state_dict['mask_prob']
        self.tau = state_dict['tau']
        self.momentum_beta = state_dict['momentum_beta']
        self.alignment_ema = state_dict['alignment_ema']
        if 'moment_key' in state_dict:
            self.moment_key = state_dict['moment_key']

        key_to_parameter: dict[tuple[int, int], Tensor] = {
            (group_index, parameter_index): parameter
            for group_index, group in enumerate(self.param_groups)
            for parameter_index, parameter in enumerate(group['params'])
        }

        self._state = {}
        for key, parameter_state in state_dict.get('magma_state', {}).items():
            parameter = key_to_parameter.get(key)
            if parameter is not None:
                self._state[id(parameter)] = {
                    name: value.to(
                        device=parameter.device, dtype=torch.float32 if name == 'alignment' else parameter.dtype
                    ).clone()
                    for name, value in parameter_state.items()
                }

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

    def _get_first_moment(self, parameter: Tensor) -> Tensor | None:
        if self.moment_key is None:
            return None

        parameter_state = self.optimizer.state.get(parameter)
        if not parameter_state:
            return None

        if self.moment_key != 'auto':
            moment = parameter_state.get(self.moment_key)
            return moment if isinstance(moment, Tensor) else None

        for key in _MOMENT_KEYS:
            moment = parameter_state.get(key)
            if isinstance(moment, Tensor):
                return moment

        return None

    def _capture_gradients(self, saved: list[tuple[Tensor, Tensor | None, Tensor]]) -> None:
        for index, (parameter, _, snapshot) in enumerate(saved):
            if parameter.grad is not None:
                saved[index] = (parameter, parameter.grad.detach().clone(), snapshot)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        saved: list[tuple[Tensor, Tensor | None, Tensor]] = []
        for group in self.param_groups:
            for parameter in group['params']:
                if torch.is_complex(parameter):
                    raise NoComplexParameterError(str(self))
                if id(parameter) in self._exclude_ids:
                    continue
                if closure is not None or parameter.grad is not None:
                    gradient = parameter.grad.detach().clone() if parameter.grad is not None else None
                    saved.append((parameter, gradient, parameter.detach().clone()))

        if closure is not None:

            def magma_closure():
                loss = closure()
                self._capture_gradients(saved)
                return loss

            loss = self.optimizer.step(magma_closure)
        else:
            loss = self.optimizer.step()

        mask_probabilities: dict[torch.device, Tensor] = {}
        for parameter, gradient, snapshot in saved:
            if gradient is None or gradient.is_sparse:
                continue

            parameter_id = id(parameter)
            if parameter_id not in self._state:
                self._state[parameter_id] = {'alignment': torch.tensor(1.0, device=parameter.device)}

            parameter_state = self._state[parameter_id]
            moment = self._get_first_moment(parameter)
            if moment is None:
                if 'momentum' not in parameter_state:
                    parameter_state['momentum'] = torch.zeros_like(parameter)
                parameter_state['momentum'].lerp_(gradient, weight=1.0 - self.momentum_beta)
                moment = parameter_state['momentum']

            cosine = torch.nn.functional.cosine_similarity(
                moment.flatten().unsqueeze(0), gradient.flatten().unsqueeze(0)
            ).squeeze(0)
            alignment_target = torch.sigmoid(cosine.float() / self.tau)
            alignment_state = parameter_state['alignment']
            alignment_state.lerp_(alignment_target, 1.0 - self.alignment_ema)

            mask_probability = mask_probabilities.get(parameter.device)
            if mask_probability is None:
                mask_probability = torch.tensor(self.mask_prob, device=parameter.device)
                mask_probabilities[parameter.device] = mask_probability

            blend = alignment_state * torch.bernoulli(mask_probability)
            parameter.lerp_(snapshot, 1.0 - blend.to(dtype=parameter.dtype))

        return loss

    @torch.no_grad()
    def zero_grad(self, set_to_none: bool = True) -> None:
        self.optimizer.zero_grad(set_to_none=set_to_none)

MARS

Bases: BaseOptimizer

Adaptive updates with variance reduced gradient corrections.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.003
betas Betas

Decay rates for the first and second moments.

(0.95, 0.99)
gamma float

The scaling parameter that controls the strength of gradient correction.

0.025
mars_type MARS_TYPE

Type of MARS. Supported types are adamw, lion, shampoo.

'adamw'
optimize_1d bool

Whether MARS should optimize 1D parameters.

False
lr_1d float

Learning rate for AdamW when optimize_1d is set to False.

0.003
betas_1d Betas

Decay rates for gradient momentum and squared gradients in the 1D AdamW groups.

(0.9, 0.95)
weight_decay float

Weight decay coefficient.

0.0
weight_decay_1d float

Weight decay for 1D parameters.

0.1
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
cautious bool

Mask momentum updates that disagree with the gradient sign.

False
Source code in pytorch_optimizer/optimizer/mars.py
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
class MARS(BaseOptimizer):
    """Adaptive updates with variance reduced gradient corrections.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        gamma: The scaling parameter that controls the strength of gradient correction.
        mars_type: Type of MARS. Supported types are `adamw`, `lion`, `shampoo`.
        optimize_1d: Whether MARS should optimize 1D parameters.
        lr_1d: Learning rate for AdamW when optimize_1d is set to False.
        betas_1d: Decay rates for gradient momentum and squared gradients in the 1D AdamW groups.
        weight_decay: Weight decay coefficient.
        weight_decay_1d: Weight decay for 1D parameters.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.
        cautious: Mask momentum updates that disagree with the gradient sign.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 3e-3,
        betas: Betas = (0.95, 0.99),
        gamma: float = 0.025,
        mars_type: MARS_TYPE = 'adamw',
        optimize_1d: bool = False,
        lr_1d: float = 3e-3,
        betas_1d: Betas = (0.9, 0.95),
        weight_decay: float = 0.0,
        weight_decay_1d: float = 1e-1,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        ams_bound: bool = False,
        cautious: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_learning_rate(lr_1d)
        self.validate_betas(betas)
        self.validate_betas(betas_1d)
        self.validate_options(mars_type, 'mars_type', ['adamw', 'lion', 'shampoo'])
        self.validate_non_negative(gamma, 'gamma')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(weight_decay_1d, 'weight_decay_1d')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'lr_1d': lr_1d,
            'lr_1d_factor': lr_1d / lr,
            'betas': betas,
            'betas_1d': betas_1d,
            'mars_type': mars_type,
            'gamma': gamma,
            'optimize_1d': optimize_1d,
            'weight_decay': weight_decay,
            'weight_decay_1d': weight_decay_1d,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'ams_bound': ams_bound,
            'cautious': cautious,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'MARS'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['last_grad'] = torch.zeros_like(p)

                if group['ams_bound']:
                    state['max_exp_avg_sq'] = torch.zeros_like(p)

    def optimize_mixed(
        self,
        grad: torch.Tensor,
        last_grad: torch.Tensor,
        exp_avg: torch.Tensor,
        exp_avg_sq: torch.Tensor,
        max_exp_avg_sq: torch.Tensor | None,
        betas: tuple[int, int],
        gamma: float,
        mars_type: MARS_TYPE,
        is_grad_2d: bool,
        step: int,
        ams_bound: bool,
        cautious: bool,
        eps: float,
    ) -> torch.Tensor:
        beta1, beta2 = betas

        c_t = (grad - last_grad).mul_(gamma * (beta1 / (1.0 - beta1))).add_(grad)
        c_t_norm = torch.norm(c_t)
        if c_t_norm > 1.0:
            c_t.div_(c_t_norm)

        exp_avg.lerp_(c_t, weight=1.0 - beta1)

        update = exp_avg.clone()
        if cautious:
            self.apply_cautious(update, grad)

        if mars_type == 'adamw' or (mars_type == 'shampoo' and not is_grad_2d):
            exp_avg_sq.mul_(beta2).addcmul_(c_t, c_t, value=1.0 - beta2)

            bias_correction1: float = self.debias(beta1, step)
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, step))

            de_nom = self.apply_ams_bound(ams_bound, exp_avg_sq, max_exp_avg_sq, eps)
            de_nom.div_(bias_correction2_sq).mul_(bias_correction1)

            return update.div_(de_nom)

        if mars_type == 'lion':
            return update.sign_()

        factor: float = math.sqrt(max(1.0, grad.size(0) / grad.size(1)))

        update = update.view(update.size(0), -1)

        return zero_power_via_newton_schulz_5(update.mul_(1.0 / (1.0 - beta1)), eps=eps).mul_(factor).view_as(grad)

    def optimize_1d(
        self,
        grad: torch.Tensor,
        exp_avg: torch.Tensor,
        exp_avg_sq: torch.Tensor,
        max_exp_avg_sq: torch.Tensor | None,
        betas: tuple[int, int],
        step: int,
        ams_bound: bool,
        cautious: bool,
        eps: float,
    ) -> torch.Tensor:
        beta1, beta2 = betas

        bias_correction1: float = self.debias(beta1, step)
        bias_correction2_sq: float = math.sqrt(self.debias(beta2, step))

        exp_avg.lerp_(grad, weight=1.0 - beta1)
        exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

        update = exp_avg.clone()

        if cautious:
            self.apply_cautious(update, grad)

        de_nom = self.apply_ams_bound(ams_bound, exp_avg_sq, max_exp_avg_sq, eps)
        de_nom.div_(bias_correction2_sq).mul_(bias_correction1)

        return update.div_(de_nom)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq, last_grad = state['exp_avg'], state['exp_avg_sq'], state['last_grad']

                p, grad, exp_avg, exp_avg_sq, last_grad = self.view_as_real(p, grad, exp_avg, exp_avg_sq, last_grad)

                is_grad_2d: bool = grad.ndim >= 2
                step_size: float = (
                    group['lr'] if group['optimize_1d'] or is_grad_2d else group['lr'] * group['lr_1d_factor']
                )

                if group['optimize_1d'] or is_grad_2d:
                    update = self.optimize_mixed(
                        grad,
                        last_grad,
                        exp_avg,
                        exp_avg_sq,
                        state.get('max_exp_avg_sq', None),
                        group['betas'],
                        group['gamma'],
                        group['mars_type'],
                        is_grad_2d,
                        group['step'],
                        group['ams_bound'],
                        group.get('cautious'),
                        group['eps'],
                    )
                else:
                    update = self.optimize_1d(
                        grad,
                        exp_avg,
                        exp_avg_sq,
                        state.get('max_exp_avg_sq', None),
                        group['betas_1d'],
                        group['step'],
                        group['ams_bound'],
                        group.get('cautious'),
                        group['eps'],
                    )

                self.apply_weight_decay(
                    p,
                    grad,
                    lr=step_size,
                    weight_decay=(
                        group['weight_decay']
                        if group['optimize_1d'] or is_grad_2d
                        else group['weight_decay_1d']
                    ),
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                p.add_(update, alpha=-step_size)

                last_grad.copy_(torch.view_as_complex(grad) if torch.is_complex(last_grad) else grad)

        return loss

MSVAG

Bases: BaseOptimizer

Adaptive momentum updates with gradient variance damping.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.01
beta float

Moving average (momentum) constant (scalar tensor or float value).

0.9
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/msvag.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
class MSVAG(BaseOptimizer):
    """Adaptive momentum updates with gradient variance damping.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        beta: Moving average (momentum) constant (scalar tensor or float value).
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-2,
        beta: float = 0.9,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(beta, 'beta', 0.0, 1.0, range_type='[]')

        self.maximize = maximize

        defaults: Defaults = {'lr': lr, 'beta': beta}

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'MSVAG'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['s'] = torch.zeros_like(p)

    @staticmethod
    def get_rho(beta_power: float, beta: float) -> float:
        """Compute the finite step variance correction for a moving average."""
        rho: float = (1.0 - beta_power ** 2) * (1.0 - beta) ** 2  # fmt: skip
        rho /= (1.0 - beta ** 2) * (1.0 - beta_power) ** 2  # fmt: skip
        return min(rho, 0.9999)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta: float = group['beta']
            beta_power: float = beta ** group['step']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

                exp_avg.lerp_(grad, weight=1.0 - beta)
                exp_avg_sq.mul_(beta).addcmul_(grad, grad, value=1.0 - beta)

                m = exp_avg.div(1.0 - beta_power)
                v = exp_avg_sq.div(1.0 - beta_power)

                rho: float = self.get_rho(beta_power, beta)

                m_p2 = m.pow(2)
                s = (v - m_p2).div_(1.0 - rho)

                factor = m_p2.div(m_p2 + rho * s)
                torch.nan_to_num(factor, nan=0.0, out=factor)
                factor.clamp_(0.0, 1.0)

                p.add_(m * factor, alpha=-group['lr'])

        return loss

get_rho(beta_power, beta) staticmethod

Compute the finite step variance correction for a moving average.

Source code in pytorch_optimizer/optimizer/msvag.py
58
59
60
61
62
63
@staticmethod
def get_rho(beta_power: float, beta: float) -> float:
    """Compute the finite step variance correction for a moving average."""
    rho: float = (1.0 - beta_power ** 2) * (1.0 - beta) ** 2  # fmt: skip
    rho /= (1.0 - beta ** 2) * (1.0 - beta_power) ** 2  # fmt: skip
    return min(rho, 0.9999)

Muon

Bases: MuonBase

Momentum updates with Newton-Schulz matrix orthogonalization.

Set use_muon=True for hidden weight matrices and use_muon=False for AdamW groups, such as embeddings, classifier heads, biases, and gains. Pass higher dimensional weights directly. The orthogonal update uses a flattened matrix view.

Parameters:

Name Type Description Default
params ParamsT

Parameter group dictionaries with a use_muon flag for each group.

required
lr float

Learning rate.

0.02
momentum float

Momentum factor.

0.95
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
nesterov bool

Use Nesterov momentum.

True
ns_steps int

Number of Newton-Schulz iterations.

5
ns_coeffs NewtonSchulzWeights

Newton-Schulz coefficients or preset name.

'original'
use_adjusted_lr bool

Scale orthogonal updates using the Moonlight shape adjustment.

False
adamw_lr float

Learning rate for parameters in the AdamW groups.

0.0003
adamw_betas Betas

Decay rates for the first and second moments in the AdamW groups.

(0.9, 0.95)
adamw_wd float

Weight decay for parameters in the AdamW groups.

0.0
adamw_eps float

Numerical stability constant for the AdamW groups.

1e-10
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Batch tensor updates and compatible matrix shapes. False disables batching; None enables it.

False

Examples:

from pytorch_optimizer import Muon

hidden_weights = [p for p in model.body.parameters() if p.ndim >= 2]
hidden_gains_biases = [p for p in model.body.parameters() if p.ndim < 2]
non_hidden_params = [*model.head.parameters(), *model.embed.parameters()]

param_groups = [
    dict(params=hidden_weights, lr=0.02, weight_decay=0.01, use_muon=True),
    dict(
        params=hidden_gains_biases + non_hidden_params,
        lr=3e-4,
        betas=(0.9, 0.95),
        weight_decay=0.01,
        use_muon=False,
    ),
]

optimizer = Muon(param_groups)
Source code in pytorch_optimizer/optimizer/muon.py
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
class Muon(MuonBase):
    """Momentum updates with Newton-Schulz matrix orthogonalization.

    Set `use_muon=True` for hidden weight matrices and `use_muon=False` for AdamW groups,
    such as embeddings, classifier heads, biases, and gains. Pass higher dimensional
    weights directly. The orthogonal update uses a flattened matrix view.

    Args:
        params: Parameter group dictionaries with a `use_muon` flag for each group.
        lr: Learning rate.
        momentum: Momentum factor.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        nesterov: Use Nesterov momentum.
        ns_steps: Number of Newton-Schulz iterations.
        ns_coeffs: Newton-Schulz coefficients or preset name.
        use_adjusted_lr: Scale orthogonal updates using the Moonlight shape adjustment.
        adamw_lr: Learning rate for parameters in the AdamW groups.
        adamw_betas: Decay rates for the first and second moments in the AdamW groups.
        adamw_wd: Weight decay for parameters in the AdamW groups.
        adamw_eps: Numerical stability constant for the AdamW groups.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Batch tensor updates and compatible matrix shapes. `False` disables batching; `None` enables it.

    Examples:
        ```python
        from pytorch_optimizer import Muon

        hidden_weights = [p for p in model.body.parameters() if p.ndim >= 2]
        hidden_gains_biases = [p for p in model.body.parameters() if p.ndim < 2]
        non_hidden_params = [*model.head.parameters(), *model.embed.parameters()]

        param_groups = [
            dict(params=hidden_weights, lr=0.02, weight_decay=0.01, use_muon=True),
            dict(
                params=hidden_gains_biases + non_hidden_params,
                lr=3e-4,
                betas=(0.9, 0.95),
                weight_decay=0.01,
                use_muon=False,
            ),
        ]

        optimizer = Muon(param_groups)
        ```

    """

    _muon_state_keys = ('momentum_buffer',)

    def __init__(
        self,
        params: ParamsT,
        lr: float = 2e-2,
        momentum: float = 0.95,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        nesterov: bool = True,
        ns_steps: int = 5,
        ns_coeffs: NewtonSchulzWeights = 'original',
        use_adjusted_lr: bool = False,
        adamw_lr: float = 3e-4,
        adamw_betas: Betas = (0.9, 0.95),
        adamw_wd: float = 0.0,
        adamw_eps: float = 1e-10,
        maximize: bool = False,
        foreach: bool | None = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_learning_rate(adamw_lr)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(momentum, 'momentum', 0.0, 1.0, range_type='[)')
        self.validate_positive(ns_steps, 'ns_steps')
        self.validate_betas(adamw_betas)
        self.validate_non_negative(adamw_wd, 'adamw_wd')
        self.validate_non_negative(adamw_eps, 'adamw_eps')
        ns_coeffs = get_newton_schulz_weights(ns_coeffs)

        self.maximize = maximize
        self.foreach = foreach

        for group in params:
            group = cast(ParamGroup, group)
            if 'use_muon' not in group:
                raise ValueError('`use_muon` must be set.')

            if group['use_muon']:
                group['lr'] = group.get('lr', lr)
                group['momentum'] = group.get('momentum', momentum)
                group['nesterov'] = group.get('nesterov', nesterov)
                group['weight_decay'] = group.get('weight_decay', weight_decay)
                group['ns_steps'] = group.get('ns_steps', ns_steps)
                group['ns_coeffs'] = get_newton_schulz_weights(group.get('ns_coeffs', ns_coeffs))
                group['use_adjusted_lr'] = group.get('use_adjusted_lr', use_adjusted_lr)
            else:
                group['lr'] = group.get('lr', adamw_lr)
                group['betas'] = group.get('betas', adamw_betas)
                group['eps'] = group.get('eps', adamw_eps)
                group['weight_decay'] = group.get('weight_decay', adamw_wd)

            group['weight_decouple'] = group.get('weight_decouple', weight_decouple)

        super().__init__(params, {'foreach': foreach, **kwargs})

    def __str__(self) -> str:
        return 'Muon'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                if group['use_muon']:
                    state['momentum_buffer'] = torch.zeros_like(p)
                else:
                    state['exp_avg'] = torch.zeros_like(p)
                    state['exp_avg_sq'] = torch.zeros_like(p)

    def _step_muon_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        state_dict: dict[str, list[torch.Tensor]],
        bias_correction2: float | torch.Tensor,
    ) -> None:
        updates = self._momentum_updates(group, grads, state_dict['momentum_buffer'])
        self._apply_muon_updates(group, params, grads, updates)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if self.can_use_foreach(group, group.get('foreach', self.foreach)):
                self._step_foreach_group(group)
                continue

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=False,
                )

                if group['use_muon']:
                    buf = state['momentum_buffer']
                    buf.lerp_(grad, weight=1.0 - group['momentum'])

                    update = grad.lerp_(buf, weight=group['momentum']) if group['nesterov'] else buf
                    if update.ndim > 2:
                        update = update.view(len(update), -1)

                    update = zero_power_via_newton_schulz_5(
                        update, num_steps=group['ns_steps'], weights=group['ns_coeffs']
                    )

                    if group.get('cautious'):
                        self.apply_cautious(update.reshape(p.shape), grad)

                    lr = get_adjusted_lr(group['lr'], p.size(), use_adjusted_lr=group['use_adjusted_lr'])

                    p.add_(update.reshape(p.shape), alpha=-lr)
                else:
                    exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                    beta1, beta2 = group['betas']

                    bias_correction1: float = self.debias(beta1, group['step'])
                    bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

                    exp_avg.lerp_(grad, weight=1.0 - beta1)
                    exp_avg_sq.lerp_(grad.square(), weight=1.0 - beta2)

                    de_nom = exp_avg_sq.sqrt().div_(bias_correction2_sq).add_(group['eps'])

                    p.addcdiv_(exp_avg / bias_correction1, de_nom, value=-group['lr'])

        return loss

Nero

Bases: BaseOptimizer

Neuron wise adaptive updates with optional weight constraints.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.01
beta float

Decay rate for squared neuron wise gradient norms.

0.999
constraints bool

Center and normalize weights with more than one dimension after each update.

True
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/nero.py
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
class Nero(BaseOptimizer):
    """Neuron wise adaptive updates with optional weight constraints.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        beta: Decay rate for squared neuron wise gradient norms.
        constraints: Center and normalize weights with more than one dimension after each update.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 0.01,
        beta: float = 0.999,
        constraints: bool = True,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(beta, 'beta', 0.0, 1.0, range_type='[]')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {'lr': lr, 'beta': beta, 'constraints': constraints, 'eps': eps}

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Nero'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                if group['constraints'] and p.dim() > 1:
                    p.sub_(neuron_mean(p))
                    p.div_(neuron_norm(p).add_(group['eps']))

                state['exp_avg_sq'] = torch.zeros_like(neuron_norm(p))

                state['scale'] = neuron_norm(p).mean()
                if state['scale'] == 0.0:
                    state['scale'] = 0.01

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            bias_correction: float = self.debias(group['beta'], group['step'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                grad_norm = neuron_norm(grad)

                exp_avg_sq = state['exp_avg_sq']
                exp_avg_sq.mul_(group['beta']).addcmul_(grad_norm, grad_norm, value=1.0 - group['beta'])

                grad_normed = grad / ((exp_avg_sq / bias_correction).sqrt_().add_(group['eps']))
                torch.nan_to_num(grad_normed, nan=0.0, out=grad_normed)

                p.add_(grad_normed, alpha=-group['lr'] * state['scale'])

                if group['constraints'] and p.dim() > 1:
                    p.sub_(neuron_mean(p))
                    p.div_(neuron_norm(p).add_(group['eps']))

        return loss

NorMuon

Bases: MuonBase

Muon updates with row wise second moment normalization.

Set use_muon=True for hidden weight matrices and use_muon=False for AdamW groups, such as embeddings, classifier heads, biases, and gains. Pass higher dimensional weights directly. The orthogonal update uses a flattened matrix view.

Parameters:

Name Type Description Default
params ParamsT

Parameter group dictionaries with a use_muon flag for each group.

required
lr float

Learning rate.

0.02
momentum float

Momentum factor.

0.95
beta2 float

Decay rate of the row wise second moment of the orthogonalized update.

0.95
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
nesterov bool

Use Nesterov momentum.

True
ns_steps int

Number of Newton-Schulz iterations.

5
ns_coeffs NewtonSchulzWeights

Newton-Schulz coefficients or preset name.

'original'
update_scale str

How to rescale the row normalized update. preserve_norm keeps the Frobenius norm of the orthogonalized update, as the official code does. match_rms gives it the Frobenius norm 0.2 * sqrt(m * n), so the RMS is 0.2 as in Algorithm 1 of the paper.

'preserve_norm'
use_adjusted_lr bool

Apply the Moonlight shape adjustment in preserve_norm mode. Unused in match_rms mode.

True
adamw_lr float

Learning rate for parameters in the AdamW groups.

0.0003
adamw_betas Betas

Decay rates for the first and second moments in the AdamW groups.

(0.9, 0.95)
adamw_wd float

Weight decay for parameters in the AdamW groups.

0.0
adamw_eps float

Numerical stability constant for the AdamW groups.

1e-10
eps float

Term added to the denominator of the row wise normalization.

1e-10
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Batch tensor updates and compatible matrix shapes. False disables batching; None enables it.

False

Examples:

from pytorch_optimizer import NorMuon

hidden_weights = [p for p in model.body.parameters() if p.ndim >= 2]
hidden_gains_biases = [p for p in model.body.parameters() if p.ndim < 2]
non_hidden_params = [*model.head.parameters(), *model.embed.parameters()]

param_groups = [
    dict(params=hidden_weights, lr=0.02, weight_decay=0.01, use_muon=True),
    dict(
        params=hidden_gains_biases + non_hidden_params,
        lr=3e-4,
        betas=(0.9, 0.95),
        weight_decay=0.01,
        use_muon=False,
    ),
]

optimizer = NorMuon(param_groups)
Source code in pytorch_optimizer/optimizer/muon.py
1088
1089
1090
1091
1092
1093
1094
1095
1096
1097
1098
1099
1100
1101
1102
1103
1104
1105
1106
1107
1108
1109
1110
1111
1112
1113
1114
1115
1116
1117
1118
1119
1120
1121
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131
1132
1133
1134
1135
1136
1137
1138
1139
1140
1141
1142
1143
1144
1145
1146
1147
1148
1149
1150
1151
1152
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
1188
1189
1190
1191
1192
1193
1194
1195
1196
1197
1198
1199
1200
1201
1202
1203
1204
1205
1206
1207
1208
1209
1210
1211
1212
1213
1214
1215
1216
1217
1218
1219
1220
1221
1222
1223
1224
1225
1226
1227
1228
1229
1230
1231
1232
1233
1234
1235
1236
1237
1238
1239
1240
1241
1242
1243
1244
1245
1246
1247
1248
1249
1250
1251
1252
1253
1254
1255
1256
1257
1258
1259
1260
1261
1262
1263
1264
1265
1266
1267
1268
1269
1270
1271
1272
1273
1274
1275
1276
1277
1278
1279
1280
1281
1282
1283
1284
1285
1286
1287
1288
1289
1290
1291
1292
1293
1294
1295
1296
1297
1298
1299
1300
1301
1302
1303
1304
1305
1306
1307
1308
1309
1310
1311
1312
1313
1314
1315
1316
1317
1318
1319
1320
1321
1322
1323
1324
1325
1326
1327
1328
1329
1330
1331
1332
1333
1334
1335
1336
1337
1338
1339
1340
1341
1342
1343
1344
1345
class NorMuon(MuonBase):
    """Muon updates with row wise second moment normalization.

    Set `use_muon=True` for hidden weight matrices and `use_muon=False` for AdamW groups,
    such as embeddings, classifier heads, biases, and gains. Pass higher dimensional
    weights directly. The orthogonal update uses a flattened matrix view.

    Args:
        params: Parameter group dictionaries with a `use_muon` flag for each group.
        lr: Learning rate.
        momentum: Momentum factor.
        beta2: Decay rate of the row wise second moment of the orthogonalized update.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        nesterov: Use Nesterov momentum.
        ns_steps: Number of Newton-Schulz iterations.
        ns_coeffs: Newton-Schulz coefficients or preset name.
        update_scale: How to rescale the row normalized update. `preserve_norm` keeps the Frobenius norm of the
            orthogonalized update, as the official code does. `match_rms` gives it the Frobenius norm `0.2 * sqrt(m
            * n)`, so the RMS is 0.2 as in Algorithm 1 of the paper.
        use_adjusted_lr: Apply the Moonlight shape adjustment in `preserve_norm` mode. Unused in `match_rms` mode.
        adamw_lr: Learning rate for parameters in the AdamW groups.
        adamw_betas: Decay rates for the first and second moments in the AdamW groups.
        adamw_wd: Weight decay for parameters in the AdamW groups.
        adamw_eps: Numerical stability constant for the AdamW groups.
        eps: Term added to the denominator of the row wise normalization.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Batch tensor updates and compatible matrix shapes. `False` disables batching; `None` enables it.

    Examples:
        ```python
        from pytorch_optimizer import NorMuon

        hidden_weights = [p for p in model.body.parameters() if p.ndim >= 2]
        hidden_gains_biases = [p for p in model.body.parameters() if p.ndim < 2]
        non_hidden_params = [*model.head.parameters(), *model.embed.parameters()]

        param_groups = [
            dict(params=hidden_weights, lr=0.02, weight_decay=0.01, use_muon=True),
            dict(
                params=hidden_gains_biases + non_hidden_params,
                lr=3e-4,
                betas=(0.9, 0.95),
                weight_decay=0.01,
                use_muon=False,
            ),
        ]

        optimizer = NorMuon(param_groups)
        ```

    """

    _muon_state_keys = ('momentum_buffer', 'second_momentum_buffer')

    def __init__(
        self,
        params: ParamsT,
        lr: float = 2e-2,
        momentum: float = 0.95,
        beta2: float = 0.95,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        nesterov: bool = True,
        ns_steps: int = 5,
        ns_coeffs: NewtonSchulzWeights = 'original',
        update_scale: str = 'preserve_norm',
        use_adjusted_lr: bool = True,
        adamw_lr: float = 3e-4,
        adamw_betas: Betas = (0.9, 0.95),
        adamw_wd: float = 0.0,
        adamw_eps: float = 1e-10,
        eps: float = 1e-10,
        maximize: bool = False,
        foreach: bool | None = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_learning_rate(adamw_lr)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(momentum, 'momentum', 0.0, 1.0, range_type='[)')
        self.validate_range(beta2, 'beta2', 0.0, 1.0, range_type='[)')
        self.validate_positive(ns_steps, 'ns_steps')
        self.validate_options(update_scale, 'update_scale', ['preserve_norm', 'match_rms'])
        self.validate_betas(adamw_betas)
        self.validate_non_negative(adamw_wd, 'adamw_wd')
        self.validate_non_negative(adamw_eps, 'adamw_eps')
        self.validate_non_negative(eps, 'eps')
        ns_coeffs = get_newton_schulz_weights(ns_coeffs)

        self.maximize = maximize
        self.foreach = foreach

        for group in params:
            group = cast(ParamGroup, group)
            if 'use_muon' not in group:
                raise ValueError('`use_muon` must be set.')

            if group['use_muon']:
                group['lr'] = group.get('lr', lr)
                group['momentum'] = group.get('momentum', momentum)
                group['beta2'] = group.get('beta2', beta2)
                group['nesterov'] = group.get('nesterov', nesterov)
                group['weight_decay'] = group.get('weight_decay', weight_decay)
                group['ns_steps'] = group.get('ns_steps', ns_steps)
                group['ns_coeffs'] = get_newton_schulz_weights(group.get('ns_coeffs', ns_coeffs))
                group['update_scale'] = group.get('update_scale', update_scale)
                group['use_adjusted_lr'] = group.get('use_adjusted_lr', use_adjusted_lr)
                group['eps'] = group.get('eps', eps)
            else:
                group['lr'] = group.get('lr', adamw_lr)
                group['betas'] = group.get('betas', adamw_betas)
                group['eps'] = group.get('eps', adamw_eps)
                group['weight_decay'] = group.get('weight_decay', adamw_wd)

            group['weight_decouple'] = group.get('weight_decouple', weight_decouple)

        super().__init__(params, {'foreach': foreach, **kwargs})

    def __str__(self) -> str:
        return 'NorMuon'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                if group['use_muon']:
                    state['momentum_buffer'] = torch.zeros_like(p)
                    state['second_momentum_buffer'] = p.new_zeros(p.size(0), 1)
                else:
                    state['exp_avg'] = torch.zeros_like(p)
                    state['exp_avg_sq'] = torch.zeros_like(p)

    def _step_muon_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        state_dict: dict[str, list[torch.Tensor]],
        bias_correction2: float | torch.Tensor,
    ) -> None:
        updates = self._momentum_updates(group, grads, state_dict['momentum_buffer'])
        updates = [update.to(grads[0].dtype) for update in updates]
        original_norms = torch._foreach_norm(updates, ord=2)

        second_moments = state_dict['second_momentum_buffer']
        row_means = [update.square().mean(dim=-1, keepdim=True) for update in updates]
        torch._foreach_lerp_(second_moments, row_means, weight=1.0 - group['beta2'])

        de_noms = torch._foreach_sqrt(second_moments)
        torch._foreach_add_(de_noms, group['eps'])
        torch._foreach_div_(updates, de_noms)

        norms = torch._foreach_norm(updates, ord=2)
        torch._foreach_add_(norms, group['eps'])

        if group['update_scale'] == 'preserve_norm':
            torch._foreach_mul_(updates, torch._foreach_div(original_norms, norms))
            lr = get_adjusted_lr(group['lr'], params[0].shape, use_adjusted_lr=group['use_adjusted_lr'])
        else:
            scales = [0.2 * math.sqrt(update.numel()) / norm for update, norm in zip(updates, norms)]
            torch._foreach_mul_(updates, scales)
            lr = group['lr']

        updates = [update.reshape(p.shape) for p, update in zip(params, updates)]
        foreach_add_(params, updates, alpha=-lr)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if self.can_use_foreach(group, group.get('foreach', self.foreach)):
                self._step_foreach_group(group)
                continue

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=False,
                )

                if group['use_muon']:
                    buf = state['momentum_buffer']
                    buf.lerp_(grad, weight=1.0 - group['momentum'])

                    update = grad.lerp_(buf, weight=group['momentum']) if group['nesterov'] else buf
                    update = update.reshape(len(update), -1)

                    update = zero_power_via_newton_schulz_5(
                        update, num_steps=group['ns_steps'], weights=group['ns_coeffs']
                    ).to(grad.dtype)

                    original_norm = update.norm()

                    v_mean = update.square().mean(dim=-1, keepdim=True)
                    second_momentum = state['second_momentum_buffer']
                    second_momentum.lerp_(v_mean, weight=1.0 - group['beta2'])

                    update.div_(second_momentum.sqrt().add_(group['eps']))

                    if group['update_scale'] == 'preserve_norm':
                        update.mul_(original_norm / update.norm().add_(group['eps']))
                        lr = get_adjusted_lr(group['lr'], p.size(), use_adjusted_lr=group['use_adjusted_lr'])
                    else:
                        update.mul_(0.2 * math.sqrt(update.numel()) / update.norm().add_(group['eps']))
                        lr = group['lr']

                    p.add_(update.reshape(p.shape), alpha=-lr)
                else:
                    exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                    beta1, beta2 = group['betas']

                    bias_correction1: float = self.debias(beta1, group['step'])
                    bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

                    exp_avg.lerp_(grad, weight=1.0 - beta1)
                    exp_avg_sq.lerp_(grad.square(), weight=1.0 - beta2)

                    de_nom = exp_avg_sq.sqrt().div_(bias_correction2_sq).add_(group['eps'])

                    p.addcdiv_(exp_avg / bias_correction1, de_nom, value=-group['lr'])

        return loss

NovoGrad

Bases: BaseOptimizer

Adaptive updates with layer wise squared gradient norms.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.95, 0.98)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
grad_averaging bool

Scale new gradient contributions by 1 - beta1.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/novograd.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
class NovoGrad(BaseOptimizer):
    """Adaptive updates with layer wise squared gradient norms.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        grad_averaging: Scale new gradient contributions by `1 - beta1`.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.95, 0.98),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        grad_averaging: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'grad_averaging': grad_averaging,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'NovoGrad'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            step_size: float = group['lr']
            if group.get('adam_debias', False):
                step_size *= math.sqrt(self.debias(beta2, group['step']))

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]
                first_update = len(state) == 0
                grad_p2 = grad.pow(2).sum()

                if first_update:
                    state['grads_ema'] = grad_p2
                else:
                    state['grads_ema'].lerp_(grad_p2, weight=1.0 - beta2)

                de_nom = state['grads_ema'].sqrt().add_(group['eps'])
                grad.div_(de_nom)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                if first_update:
                    state['moments'] = grad.clone()
                else:
                    if group['grad_averaging']:
                        grad.mul_(1.0 - beta1)

                    state['moments'].mul_(beta1).add_(grad)

                p.add_(state['moments'], alpha=-step_size)

        return loss

OrthoGrad

Bases: BaseOptimizer

Wrap an optimizer with gradients orthogonal to the current parameters.

Parameters:

Name Type Description Default
optimizer OptimizerInstanceOrClass

Base optimizer.

required
Source code in pytorch_optimizer/optimizer/orthograd.py
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
class OrthoGrad(BaseOptimizer):
    """Wrap an optimizer with gradients orthogonal to the current parameters.

    Args:
        optimizer: Base optimizer.

    """

    def __init__(self, optimizer: OptimizerInstanceOrClass, **kwargs) -> None:
        self._optimizer_step_pre_hooks: dict[int, Callable] = OrderedDict()
        self._optimizer_step_post_hooks: dict[int, Callable] = OrderedDict()
        self._patch_step_function()
        self.eps: float = 1e-30

        self.optimizer: Optimizer = self.load_optimizer(optimizer, **kwargs)

        self.defaults: Defaults = self.optimizer.defaults

    def __str__(self) -> str:
        return 'OrthoGrad'

    @property
    def param_groups(self):
        return self.optimizer.param_groups

    @property
    def state(self) -> State:
        return self.optimizer.state

    def state_dict(self) -> State:
        return self.optimizer.state_dict()

    def load_state_dict(self, state_dict: State) -> None:
        self.optimizer.load_state_dict(state_dict)

    @torch.no_grad()
    def zero_grad(self, set_to_none: bool = True) -> None:
        self.optimizer.zero_grad(set_to_none=set_to_none)

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

    @torch.no_grad()
    def apply_orthogonal_gradients(self, params) -> None:
        super().apply_orthogonal_gradients(params, eps=self.eps)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        if closure is None:
            for group in self.param_groups:
                self.apply_orthogonal_gradients(group['params'])
            return self.optimizer.step()

        def orthogonal_closure():
            with torch.enable_grad():
                loss = closure()
            for group in self.param_groups:
                self.apply_orthogonal_gradients(group['params'])
            return loss

        return self.optimizer.step(orthogonal_closure)

PAdam

Bases: BaseOptimizer

Partially adaptive Adam with a configurable second moment exponent.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.1
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
partial float

Partially adaptive parameter.

0.25
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
Source code in pytorch_optimizer/optimizer/padam.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
class PAdam(BaseOptimizer):
    """Partially adaptive Adam with a configurable second moment exponent.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        partial: Partially adaptive parameter.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.

    """

    _supports_compiled_foreach = True

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-1,
        betas: Betas = (0.9, 0.999),
        partial: float = 0.25,
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        foreach: bool | None = None,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_range(partial, 'partial', 0.0, 1.0, range_type='(]')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'partial': partial,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            'foreach': foreach,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'PAdam'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        step_size: float | torch.Tensor,
    ) -> None:
        beta1, beta2 = group['betas']
        exponent = group['partial'] * 2

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=group['lr'],
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - beta1)

        torch._foreach_mul_(exp_avg_sqs, beta2)
        torch._foreach_addcmul_(exp_avg_sqs, grads, grads, value=1.0 - beta2)

        de_noms = torch._foreach_sqrt(exp_avg_sqs)
        torch._foreach_add_(de_noms, group['eps'])

        if exponent != 1.0:
            torch._foreach_pow_(de_noms, exponent)

        foreach_addcdiv_(params, exp_avgs, de_noms, value=-step_size)

    def _step_per_param(self, group: ParamGroup, step_size: float) -> None:
        beta1, beta2 = group['betas']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

            p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

            self.apply_weight_decay(
                p,
                grad=grad,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=group['fixed_decay'],
            )

            exp_avg.lerp_(grad, weight=1.0 - beta1)

            exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

            de_nom = exp_avg_sq.sqrt().add_(group['eps'])

            exponent = group['partial'] * 2
            if exponent != 1.0:
                de_nom.pow_(exponent)

            p.addcdiv_(exp_avg, de_nom, value=-step_size)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size: float = group['lr'] * bias_correction2_sq / bias_correction1

            if self.can_use_foreach(group, group.get('foreach')):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_avg_sq']
                )

                for tensors in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(
                        group,
                        tensors['params'],
                        tensors['grads'],
                        tensors['exp_avg'],
                        tensors['exp_avg_sq'],
                        step_size,
                    )
            else:
                self._step_per_param(group, step_size)

        return loss

PCGrad

Bases: BaseOptimizer

Wrap an optimizer with gradient projection for conflicting task objectives.

Learning rate schedulers and checkpoints use the wrapped optimizer's parameter groups and state. Checkpoint hooks receive the wrapped optimizer.

Parameters:

Name Type Description Default
optimizer Optimizer

Optimizer instance.

required
reduction str

Reduction method for gradients.

'mean'
Source code in pytorch_optimizer/optimizer/pcgrad.py
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
class PCGrad(BaseOptimizer):
    """Wrap an optimizer with gradient projection for conflicting task objectives.

    Learning rate schedulers and checkpoints use the wrapped optimizer's parameter groups and state.
    Checkpoint hooks receive the wrapped optimizer.

    Args:
        optimizer: Optimizer instance.
        reduction: Reduction method for gradients.

    """

    def __init__(self, optimizer: Optimizer, reduction: str = 'mean'):
        self.validate_options(reduction, 'reduction', ['mean', 'sum'])

        self.optimizer = optimizer
        self.reduction = reduction

        self._optimizer_step_pre_hooks: dict[int, Callable] = OrderedDict()
        self._optimizer_step_post_hooks: dict[int, Callable] = OrderedDict()
        self.defaults: Defaults = self.optimizer.defaults
        self._patch_step_function()

    @property
    def param_groups(self):
        return self.optimizer.param_groups

    @property
    def state(self) -> State:
        return self.optimizer.state

    def add_param_group(self, param_group: ParamGroup) -> None:
        self.optimizer.add_param_group(param_group)

    def state_dict(self) -> State:
        return self.optimizer.state_dict()

    def load_state_dict(self, state_dict: State) -> None:
        self.optimizer.load_state_dict(state_dict)

    def register_state_dict_pre_hook(self, hook: Callable, prepend: bool = False):
        return self.optimizer.register_state_dict_pre_hook(hook, prepend=prepend)

    def register_state_dict_post_hook(self, hook: Callable, prepend: bool = False):
        return self.optimizer.register_state_dict_post_hook(hook, prepend=prepend)

    def register_load_state_dict_pre_hook(self, hook: Callable, prepend: bool = False):
        return self.optimizer.register_load_state_dict_pre_hook(hook, prepend=prepend)

    def register_load_state_dict_post_hook(self, hook: Callable, prepend: bool = False):
        return self.optimizer.register_load_state_dict_post_hook(hook, prepend=prepend)

    @torch.no_grad()
    def init_group(self):
        self.zero_grad()

    def zero_grad(self, set_to_none: bool = True) -> None:
        self.optimizer.zero_grad(set_to_none=set_to_none)

    def step(self, closure: Closure = None) -> Loss:
        return self.optimizer.step(closure)

    def set_grad(self, grads: list[torch.Tensor], has_grads: list[torch.Tensor] | None = None) -> None:
        idx: int = 0
        for group in self.optimizer.param_groups:
            for p in group['params']:
                p.grad = grads[idx] if has_grads is None or torch.any(has_grads[idx]) else None
                idx += 1

    def retrieve_grad(self) -> tuple[list[torch.Tensor], list[int], list[torch.Tensor]]:
        """Collect gradients, shapes, and masks for parameters with gradients."""
        grad, shape, has_grad = [], [], []
        for group in self.optimizer.param_groups:
            for p in group['params']:
                if p.grad is None:
                    shape.append(p.shape)
                    grad.append(torch.zeros_like(p, device=p.device))
                    has_grad.append(torch.zeros_like(p, device=p.device))
                    continue

                shape.append(p.grad.shape)
                grad.append(p.grad.clone())
                has_grad.append(torch.ones_like(p, device=p.device))

        return grad, shape, has_grad

    def pack_grad(self, objectives: Iterable) -> tuple[list[torch.Tensor], list[list[int]], list[torch.Tensor]]:
        """Compute and flatten gradients for each task loss.

        Args:
            objectives: Scalar task loss tensors to backpropagate.

        """
        grads, shapes, has_grads = [], [], []
        for objective in objectives:
            self.optimizer.zero_grad(set_to_none=True)
            objective.backward(retain_graph=True)

            grad, shape, has_grad = self.retrieve_grad()

            grads.append(flatten_grad(grad))
            has_grads.append(flatten_grad(has_grad))
            shapes.append(shape)

        return grads, shapes, has_grads

    def project_conflicting(self, grads: list[torch.Tensor], has_grads: list[torch.Tensor]) -> torch.Tensor:
        """Remove conflicting task gradient components and combine task gradients.

        Args:
            grads: A list of the gradient of the parameters.
            has_grads: A list of masks representing whether the parameter has gradient.

        """
        shared: torch.Tensor = torch.stack(has_grads).prod(0).bool()

        pc_grad: list[torch.Tensor] = deepcopy(grads)
        for i, g_i in enumerate(pc_grad):
            random.shuffle(grads)
            for g_j in grads:
                g_i_g_j: torch.Tensor = torch.dot(g_i, g_j)
                if g_i_g_j < 0:
                    pc_grad[i] -= g_i_g_j * g_j / (g_j.norm() ** 2)

        merged_grad: torch.Tensor = torch.zeros_like(grads[0])

        shared_pc_gradients: torch.Tensor = torch.stack([g[shared] for g in pc_grad])
        if self.reduction == 'mean':
            merged_grad[shared] = shared_pc_gradients.mean(dim=0)
        else:
            merged_grad[shared] = shared_pc_gradients.sum(dim=0)

        merged_grad[~shared] = torch.stack([g[~shared] for g in pc_grad]).sum(dim=0)

        return merged_grad

    def pc_backward(self, objectives: Iterable[nn.Module]) -> None:
        """Set parameter gradients after projecting conflicting task gradients.

        Args:
            objectives: Scalar task loss tensors to backpropagate.

        """
        grads, shapes, has_grads = self.pack_grad(objectives)

        pc_grad = self.project_conflicting(grads, has_grads)
        pc_grad = un_flatten_grad(pc_grad, shapes[0])
        has_grad = un_flatten_grad(torch.stack(has_grads).sum(dim=0), shapes[0])

        self.set_grad(pc_grad, has_grad)

pack_grad(objectives)

Compute and flatten gradients for each task loss.

Parameters:

Name Type Description Default
objectives Iterable

Scalar task loss tensors to backpropagate.

required
Source code in pytorch_optimizer/optimizer/pcgrad.py
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
def pack_grad(self, objectives: Iterable) -> tuple[list[torch.Tensor], list[list[int]], list[torch.Tensor]]:
    """Compute and flatten gradients for each task loss.

    Args:
        objectives: Scalar task loss tensors to backpropagate.

    """
    grads, shapes, has_grads = [], [], []
    for objective in objectives:
        self.optimizer.zero_grad(set_to_none=True)
        objective.backward(retain_graph=True)

        grad, shape, has_grad = self.retrieve_grad()

        grads.append(flatten_grad(grad))
        has_grads.append(flatten_grad(has_grad))
        shapes.append(shape)

    return grads, shapes, has_grads

pc_backward(objectives)

Set parameter gradients after projecting conflicting task gradients.

Parameters:

Name Type Description Default
objectives Iterable[Module]

Scalar task loss tensors to backpropagate.

required
Source code in pytorch_optimizer/optimizer/pcgrad.py
167
168
169
170
171
172
173
174
175
176
177
178
179
180
def pc_backward(self, objectives: Iterable[nn.Module]) -> None:
    """Set parameter gradients after projecting conflicting task gradients.

    Args:
        objectives: Scalar task loss tensors to backpropagate.

    """
    grads, shapes, has_grads = self.pack_grad(objectives)

    pc_grad = self.project_conflicting(grads, has_grads)
    pc_grad = un_flatten_grad(pc_grad, shapes[0])
    has_grad = un_flatten_grad(torch.stack(has_grads).sum(dim=0), shapes[0])

    self.set_grad(pc_grad, has_grad)

project_conflicting(grads, has_grads)

Remove conflicting task gradient components and combine task gradients.

Parameters:

Name Type Description Default
grads list[Tensor]

A list of the gradient of the parameters.

required
has_grads list[Tensor]

A list of masks representing whether the parameter has gradient.

required
Source code in pytorch_optimizer/optimizer/pcgrad.py
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
def project_conflicting(self, grads: list[torch.Tensor], has_grads: list[torch.Tensor]) -> torch.Tensor:
    """Remove conflicting task gradient components and combine task gradients.

    Args:
        grads: A list of the gradient of the parameters.
        has_grads: A list of masks representing whether the parameter has gradient.

    """
    shared: torch.Tensor = torch.stack(has_grads).prod(0).bool()

    pc_grad: list[torch.Tensor] = deepcopy(grads)
    for i, g_i in enumerate(pc_grad):
        random.shuffle(grads)
        for g_j in grads:
            g_i_g_j: torch.Tensor = torch.dot(g_i, g_j)
            if g_i_g_j < 0:
                pc_grad[i] -= g_i_g_j * g_j / (g_j.norm() ** 2)

    merged_grad: torch.Tensor = torch.zeros_like(grads[0])

    shared_pc_gradients: torch.Tensor = torch.stack([g[shared] for g in pc_grad])
    if self.reduction == 'mean':
        merged_grad[shared] = shared_pc_gradients.mean(dim=0)
    else:
        merged_grad[shared] = shared_pc_gradients.sum(dim=0)

    merged_grad[~shared] = torch.stack([g[~shared] for g in pc_grad]).sum(dim=0)

    return merged_grad

retrieve_grad()

Collect gradients, shapes, and masks for parameters with gradients.

Source code in pytorch_optimizer/optimizer/pcgrad.py
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
def retrieve_grad(self) -> tuple[list[torch.Tensor], list[int], list[torch.Tensor]]:
    """Collect gradients, shapes, and masks for parameters with gradients."""
    grad, shape, has_grad = [], [], []
    for group in self.optimizer.param_groups:
        for p in group['params']:
            if p.grad is None:
                shape.append(p.shape)
                grad.append(torch.zeros_like(p, device=p.device))
                has_grad.append(torch.zeros_like(p, device=p.device))
                continue

            shape.append(p.grad.shape)
            grad.append(p.grad.clone())
            has_grad.append(torch.ones_like(p, device=p.device))

    return grad, shape, has_grad

PID

Bases: BaseOptimizer

SGD with proportional, integral, and derivative update terms.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
momentum float

Momentum factor.

0.0
dampening float

Dampening factor for momentum.

0.0
derivative float

Weight of the gradient difference term.

10.0
integral float

Weight of the accumulated gradient term.

5.0
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/pid.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
class PID(BaseOptimizer):
    """SGD with proportional, integral, and derivative update terms.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Momentum factor.
        dampening: Dampening factor for momentum.
        derivative: Weight of the gradient difference term.
        integral: Weight of the accumulated gradient term.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        momentum: float = 0.0,
        dampening: float = 0.0,
        derivative: float = 10.0,
        integral: float = 5.0,
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0)
        self.validate_non_negative(derivative, 'derivative')
        self.validate_non_negative(integral, 'integral')
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'momentum': momentum,
            'dampening': dampening,
            'derivative': derivative,
            'integral': integral,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'PID'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0 and group['momentum'] > 0.0:
                state['grad_buffer'] = torch.zeros_like(p)
                state['i_buffer'] = torch.zeros_like(p)
                state['d_buffer'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                g_buf, i_buf, d_buf = (
                    state.get('grad_buffer', None),
                    state.get('i_buffer', None),
                    state.get('d_buffer', None),
                )

                p, grad, g_buf, i_buf, d_buf = self.view_as_real(p, grad, g_buf, i_buf, d_buf)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                if group['momentum'] > 0.0:
                    i_buf.mul_(group['momentum']).add_(grad, alpha=1.0 - group['dampening'])
                    d_buf.mul_(group['momentum'])

                    if group['step'] > 1:
                        d_buf.add_(grad - g_buf, alpha=1.0 - group['momentum'])

                    g_buf.copy_(grad)

                    grad.add_(i_buf, alpha=group['integral']).add_(d_buf, alpha=group['derivative'])

                p.add_(grad, alpha=-group['lr'])

        return loss

PNM

Bases: BaseOptimizer

SGD with alternating positive and negative momentum.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Momentum decay rate and positive negative momentum mixing coefficient.

(0.9, 1.0)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/pnm.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
class PNM(BaseOptimizer):
    """SGD with alternating positive and negative momentum.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Momentum decay rate and positive negative momentum mixing coefficient.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 1.0),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas, beta_range_type='[]')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'PNM'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['pos_momentum'] = torch.zeros_like(p)
                state['neg_momentum'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            beta1_p2: float = beta1 ** 2  # fmt: skip
            noise_norm: float = math.sqrt((1 + beta2) ** 2 + beta2 ** 2)  # fmt: skip

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                if group['step'] % 2 == 1:
                    pos_momentum, neg_momentum = state['pos_momentum'], state['neg_momentum']
                else:
                    neg_momentum, pos_momentum = state['pos_momentum'], state['neg_momentum']

                p, grad, pos_momentum, neg_momentum = self.view_as_real(p, grad, pos_momentum, neg_momentum)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                pos_momentum.lerp_(grad, weight=1.0 - beta1_p2)

                delta_p = pos_momentum.mul(1.0 + beta2).add_(neg_momentum, alpha=-beta2).mul_(1.0 / noise_norm)

                p.add_(delta_p, alpha=-group['lr'])

        return loss

Prodigy

Bases: BaseOptimizer

Adam with distance adaptive step sizes.

Leave LR set to 1 unless you encounter instability.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

1.0
betas Betas

Decay rates for the gradient mean and squared gradients.

(0.9, 0.999)
beta3 float | None

Decay rate for the distance estimate. None uses the square root of beta2.

None
d0 float

Initial estimate of the distance to the optimum.

1e-06
d_coef float

Coefficient in the expression for the estimate of d.

1.0
growth_rate float

Maximum multiplicative growth of the distance estimate per step.

float('inf')
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
bias_correction bool

Apply bias correction to the moment estimates.

False
safeguard_warmup bool

Exclude the learning rate from the distance estimate denominator during warmup.

False
eps float | None

Term added to the denominator to improve numerical stability. when eps is None, use atan2 rather than epsilon and division for parameter updates.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/prodigy.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
class Prodigy(BaseOptimizer):
    """Adam with distance adaptive step sizes.

    Leave LR set to 1 unless you encounter instability.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the gradient mean and squared gradients.
        beta3: Decay rate for the distance estimate. `None` uses the square root of `beta2`.
        d0: Initial estimate of the distance to the optimum.
        d_coef: Coefficient in the expression for the estimate of d.
        growth_rate: Maximum multiplicative growth of the distance estimate per step.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        bias_correction: Apply bias correction to the moment estimates.
        safeguard_warmup: Exclude the learning rate from the distance estimate denominator during warmup.
        eps: Term added to the denominator to improve numerical stability. when eps is None, use atan2 rather than
            epsilon and division for parameter updates.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1.0,
        betas: Betas = (0.9, 0.999),
        beta3: float | None = None,
        d0: float = 1e-6,
        d_coef: float = 1.0,
        growth_rate: float = float('inf'),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        bias_correction: bool = False,
        safeguard_warmup: bool = False,
        eps: float | None = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas((*betas, beta3))
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'beta3': beta3,
            'd': d0,
            'd0': d0,
            'd_max': d0,
            'd_coef': d_coef,
            'growth_rate': growth_rate,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'bias_correction': bias_correction,
            'safeguard_warmup': safeguard_warmup,
            'step': 1,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Prodigy'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['s'] = torch.zeros_like(p)
                state['p0'] = p.clone()
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        group = self.param_groups[0]
        device = group['params'][0].device

        d_de_nom = torch.tensor([0.0], device=device)

        beta1, beta2 = group['betas']
        beta3: float = group['beta3'] if group['beta3'] is not None else math.sqrt(beta2)

        bias_correction1: float = self.debias(beta1, group['step'])
        bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))
        bias_correction: float = (bias_correction1 / bias_correction2_sq) if group['bias_correction'] else 1.0

        d, d0 = group['d'], group['d0']
        d_lr: float = d * group['lr'] / bias_correction

        if 'd_numerator' not in group:
            group['d_numerator'] = torch.tensor([0.0], device=device)
        elif group['d_numerator'].device != device:
            group['d_numerator'] = group['d_numerator'].to(device)  # pragma: no cover

        d_numerator = group['d_numerator']
        d_numerator.mul_(beta3)

        for group in self.param_groups:
            self.init_group(group)

            group['step'] += 1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                p0, exp_avg, exp_avg_sq = state['p0'], state['exp_avg'], state['exp_avg_sq']

                d_numerator.add_(torch.dot(grad.flatten(), (p0 - p).flatten()), alpha=(d / d0) * d_lr)

                exp_avg.mul_(beta1).add_(grad, alpha=d * (1.0 - beta1))
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=d * d * (1.0 - beta2))

                s = state['s']
                s.mul_(beta3).add_(grad, alpha=(d / d0) * (d if group['safeguard_warmup'] else d_lr))

                d_de_nom.add_(s.abs().sum())

        if d_de_nom == 0:
            return loss

        d_hat = (group['d_coef'] * d_numerator / d_de_nom).item()
        if d == group['d0']:
            d = max(d, d_hat)

        d_max = max(group['d_max'], d_hat)
        d = min(d_max, d * group['growth_rate'])

        for group in self.param_groups:
            group['d_numerator'] = d_numerator
            group['d_de_nom'] = d_de_nom
            group['d'] = d
            group['d_max'] = d_max
            group['d_hat'] = d_hat

            for p in group['params']:
                if p.grad is None:
                    continue

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                self.apply_weight_decay(
                    p,
                    p.grad,
                    lr=d_lr,
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                de_nom = exp_avg_sq.sqrt()

                if group['eps'] is not None:
                    de_nom.add_(d * group['eps'])
                    p.addcdiv_(exp_avg, de_nom, value=-d_lr)
                else:
                    update = exp_avg.clone().atan2_(de_nom)
                    p.add_(update, alpha=-d_lr)

        return loss

QHAdam

Bases: BaseOptimizer

Adam with quasi-hyperbolic moment averaging.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the gradient mean and squared gradients.

(0.9, 0.999)
nus tuple[float, float]

Weights of the running averages relative to the current gradient and its square.

(1.0, 1.0)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/qhadam.py
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
class QHAdam(BaseOptimizer):
    """Adam with quasi-hyperbolic moment averaging.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the gradient mean and squared gradients.
        nus: Weights of the running averages relative to the current gradient and its square.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        nus: tuple[float, float] = (1.0, 1.0),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_nus(nus)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'nus': nus,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'QHAdam'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['beta1_weight'] = torch.zeros((1,), dtype=torch.float32, device=grad.device)
                state['beta2_weight'] = torch.zeros((1,), dtype=torch.float32, device=grad.device)
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']
            nu1, nu2 = group['nus']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                beta1_weight, beta2_weight = state['beta1_weight'], state['beta2_weight']
                beta1_weight.mul_(beta1).add_(1.0)
                beta2_weight.mul_(beta2).add_(1.0)

                grad_p2 = grad.pow(2)

                exp_avg.lerp_(grad, weight=beta1_weight.reciprocal().to(dtype=grad.dtype))
                exp_avg_sq.lerp_(grad_p2, weight=beta2_weight.reciprocal().to(dtype=grad.dtype))

                avg_grad = exp_avg.mul(nu1)
                if nu1 != 1.0:
                    avg_grad.add_(grad, alpha=1.0 - nu1)

                avg_grad_rms = exp_avg_sq.mul(nu2)
                if nu2 != 1.0:
                    avg_grad_rms.add_(grad_p2, alpha=1.0 - nu2)

                avg_grad_rms.sqrt_().add_(group['eps'])

                p.addcdiv_(avg_grad, avg_grad_rms, value=-group['lr'])

        return loss

QHM

Bases: BaseOptimizer

SGD with quasi-hyperbolic momentum.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
momentum float

Momentum factor.

0.0
nu float

Weight of momentum relative to the current gradient.

1.0
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/qhm.py
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
class QHM(BaseOptimizer):
    """SGD with quasi-hyperbolic momentum.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Momentum factor.
        nu: Weight of momentum relative to the current gradient.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        momentum: float = 0.0,
        nu: float = 1.0,
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_nus(nu)

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'momentum': momentum,
            'nu': nu,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'QHM'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['momentum_buffer'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                buf = state['momentum_buffer']

                p, grad, buf = self.view_as_real(p, grad, buf)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                buf.lerp_(grad, weight=1.0 - group['momentum'])

                p.add_(buf, alpha=-group['lr'] * group['nu'])
                p.add_(grad, alpha=-group['lr'] * (1.0 - group['nu']))

        return loss

RACS

Bases: BaseOptimizer

Row and Column Scaled SGD.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
beta float

Decay rate for row- and column wise squared gradient averages.

0.9
alpha float

Update scaling factor.

0.05
gamma float

Maximum multiplicative growth of the scaled update norm.

1.01
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/racs.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
class RACS(BaseOptimizer):
    """Row and Column Scaled SGD.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        beta: Decay rate for row- and column wise squared gradient averages.
        alpha: Update scaling factor.
        gamma: Maximum multiplicative growth of the scaled update norm.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        beta: float = 0.9,
        alpha: float = 0.05,
        gamma: float = 1.01,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(beta, 'beta', 0.0, 1.0)
        self.validate_range(alpha, 'alpha', 0.0, 1.0)
        self.validate_positive(gamma, 'gamma')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'beta': beta,
            'alpha': alpha,
            'gamma': gamma,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'RACS'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta = group['beta']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad
                if grad.is_sparse:
                    raise NoSparseGradientError(str(self))

                if torch.is_complex(p):
                    raise NoComplexParameterError(str(self))

                state = self.state[p]

                self.maximize_gradient(grad, maximize=self.maximize)

                if grad.ndim < 2:
                    grad = grad.reshape(len(grad), 1)
                elif grad.ndim > 2:
                    grad = grad.reshape(len(grad), -1)

                has_state = 's' in state
                if not has_state:
                    state['s'] = torch.zeros(grad.size(0), dtype=grad.dtype, device=grad.device)
                    state['q'] = torch.ones(grad.size(1), dtype=grad.dtype, device=grad.device)
                    state['theta'] = torch.zeros((), dtype=grad.dtype, device=grad.device)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                s, q = state['s'], state['q']

                grad_p2 = grad.pow(2)
                s.lerp_(grad_p2.mean(dim=1), weight=1.0 - beta)
                q.lerp_(grad_p2.mean(dim=0), weight=1.0 - beta)

                s_sq = s.add(group['eps']).sqrt_().unsqueeze(1)
                q_sq = q.add(group['eps']).sqrt_().unsqueeze(0)

                grad_hat = grad / (s_sq * q_sq)

                grad_hat_norm = torch.norm(grad_hat)
                threshold = (
                    group['gamma'] / max(grad_hat_norm / (state['theta'] + group['eps']), group['gamma'])
                    if has_state
                    else 1.0
                )
                state['theta'] = grad_hat_norm.mul_(threshold)

                p.add_(grad_hat.view_as(p), alpha=-group['lr'] * group['alpha'] * threshold)

        return loss

RAdam

Bases: BaseOptimizer

Rectified Adam.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
n_sma_threshold int

Minimum effective simple moving average length for rectification.

5
degenerated_to_sgd bool

Use an SGD update before the moving average reaches the rectification threshold.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
Source code in pytorch_optimizer/optimizer/radam.py
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
class RAdam(BaseOptimizer):
    """Rectified Adam.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        n_sma_threshold: Minimum effective simple moving average length for rectification.
        degenerated_to_sgd: Use an SGD update before the moving average reaches the rectification threshold.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.

    """

    _supports_compiled_foreach = True

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        n_sma_threshold: int = 5,
        degenerated_to_sgd: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        foreach: bool | None = None,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.n_sma_threshold = n_sma_threshold
        self.degenerated_to_sgd = degenerated_to_sgd
        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            'foreach': foreach,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'RAdam'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)

                if group.get('adanorm'):
                    state['exp_grad_adanorm'] = torch.zeros((1,), dtype=p.dtype, device=p.device)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        return not group.get('adanorm') and self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        step_size: float | torch.Tensor,
        is_rectified: bool,
        apply_update: bool,
    ) -> None:
        beta1, beta2 = group['betas']

        if self.maximize:
            torch._foreach_neg_(grads)

        if group['weight_decouple'] and apply_update:
            self.apply_weight_decay_foreach(
                params=params,
                grads=grads,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=True,
                fixed_decay=group['fixed_decay'],
            )

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - beta1)

        torch._foreach_mul_(exp_avg_sqs, beta2)
        torch._foreach_addcmul_(exp_avg_sqs, grads, grads, value=1.0 - beta2)

        if is_rectified:
            de_noms = torch._foreach_sqrt(exp_avg_sqs)
            torch._foreach_add_(de_noms, group['eps'])

            foreach_addcdiv_(params, exp_avgs, de_noms, value=-step_size)
        elif apply_update:
            foreach_add_(params, exp_avgs, alpha=-step_size)

    def _step_per_param(self, group: ParamGroup, step_size: float, n_sma: float) -> None:
        beta1, beta2 = group['betas']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

            p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

            if step_size > 0 or n_sma >= self.n_sma_threshold:
                self.apply_weight_decay(
                    p=p,
                    grad=None,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

            s_grad = self.get_adanorm_gradient(
                grad=grad,
                adanorm=group.get('adanorm', False),
                exp_grad_norm=state.get('exp_grad_adanorm', None),
                r=group.get('adanorm_r', None),
            )

            exp_avg.lerp_(s_grad, weight=1.0 - beta1)

            exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

            if n_sma >= self.n_sma_threshold:
                de_nom = exp_avg_sq.sqrt().add_(group['eps'])
                p.addcdiv_(exp_avg, de_nom, value=-step_size)
            elif step_size > 0:
                p.add_(exp_avg, alpha=-step_size)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])

            step_size, n_sma = self.get_rectify_step_size(
                is_rectify=True,
                step=group['step'],
                lr=group['lr'],
                beta2=beta2,
                n_sma_threshold=self.n_sma_threshold,
                degenerated_to_sgd=self.degenerated_to_sgd,
            )

            step_size = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=step_size,
                bias_correction1=bias_correction1,
            )

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_avg_sq']
                )

                for tensors in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(
                        group,
                        tensors['params'],
                        tensors['grads'],
                        tensors['exp_avg'],
                        tensors['exp_avg_sq'],
                        step_size,
                        n_sma >= self.n_sma_threshold,
                        n_sma >= self.n_sma_threshold or bool(step_size > 0),
                    )
            else:
                self._step_per_param(group, step_size, n_sma)

        return loss

Ranger

Bases: BaseOptimizer

RAdam with Lookahead and optional gradient centralization.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.95, 0.999)
alpha float

Lookahead interpolation factor from slow weights toward fast weights.

0.5
k int

Number of steps between Lookahead updates.

6
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
n_sma_threshold int

Minimum effective simple moving average length for rectification.

5
degenerated_to_sgd bool

Use an SGD update before the moving average reaches the rectification threshold.

False
use_gc bool

Use Gradient Centralization (both convolution & fc layers).

True
gc_conv_only bool

Use Gradient Centralization (only convolution layer).

False
eps float

Term added to the denominator to improve numerical stability.

1e-05
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/ranger.py
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
class Ranger(BaseOptimizer):
    """RAdam with Lookahead and optional gradient centralization.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        alpha: Lookahead interpolation factor from slow weights toward fast weights.
        k: Number of steps between Lookahead updates.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        n_sma_threshold: Minimum effective simple moving average length for rectification.
        degenerated_to_sgd: Use an SGD update before the moving average reaches the rectification threshold.
        use_gc: Use Gradient Centralization (both convolution & fc layers).
        gc_conv_only: Use Gradient Centralization (only convolution layer).
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.95, 0.999),
        alpha: float = 0.5,
        k: int = 6,
        n_sma_threshold: int = 5,
        degenerated_to_sgd: bool = False,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        use_gc: bool = True,
        gc_conv_only: bool = False,
        eps: float = 1e-5,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_range(alpha, 'alpha', 0.0, 1.0, range_type='[]')
        self.validate_positive(k, 'k')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.n_sma_threshold = n_sma_threshold
        self.degenerated_to_sgd = degenerated_to_sgd
        self.use_gc = use_gc
        self.gc_gradient_threshold: int = 3 if gc_conv_only else 1
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'alpha': alpha,
            'k': k,
            'step_counter': 0,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Ranger'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['slow_buffer'] = p.clone()

                if group.get('adanorm'):
                    state['exp_grad_adanorm'] = torch.zeros((1,), dtype=p.dtype, device=p.device)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])

            step_size, n_sma = self.get_rectify_step_size(
                is_rectify=True,
                step=group['step'],
                lr=group['lr'],
                beta2=beta2,
                n_sma_threshold=self.n_sma_threshold,
                degenerated_to_sgd=self.degenerated_to_sgd,
            )

            step_size = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=step_size,
                bias_correction1=bias_correction1,
            )

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                exp_avg, exp_avg_sq, slow_buffer = state['exp_avg'], state['exp_avg_sq'], state['slow_buffer']

                p, grad, exp_avg, exp_avg_sq, slow_buffer = self.view_as_real(
                    p, grad, exp_avg, exp_avg_sq, slow_buffer
                )

                if self.use_gc and grad.dim() > self.gc_gradient_threshold:
                    centralize_gradient(grad, gc_conv_only=False)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                s_grad = self.get_adanorm_gradient(
                    grad=grad,
                    adanorm=group.get('adanorm', False),
                    exp_grad_norm=state.get('exp_grad_adanorm', None),
                    r=group.get('adanorm_r', None),
                )

                exp_avg.lerp_(s_grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                if n_sma >= self.n_sma_threshold:
                    de_nom = exp_avg_sq.sqrt().add_(group['eps'])
                    p.addcdiv_(exp_avg, de_nom, value=-step_size)
                else:
                    p.add_(exp_avg, alpha=-step_size)

                if group['step'] % group['k'] == 0:
                    slow_buffer.lerp_(p, weight=group['alpha'])
                    p.copy_(slow_buffer)

        return loss

Ranger21

Bases: BaseOptimizer

AdamW with positive negative momentum, gradient clipping, and Lookahead.

Includes gradient centralization and normalization, stable weight decay, norm loss, softplus smoothing, and an optional learning rate schedule with warmup and warmdown.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
num_iterations int

Total training steps for the built in learning rate schedule.

required
lr float

Learning rate.

0.001
beta0 float

Manages the amplitude of the noise introduced by positive negative momentum.

0.9
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
use_softplus bool

Use softplus to smooth.

True
beta_softplus float

Beta parameter for softplus smoothing.

50.0
disable_lr_scheduler bool

Whether to disable learning rate schedule.

False
num_warm_up_iterations int | None

Number of warmup iterations. Ranger21 performs linear learning rate warmup.

None
num_warm_down_iterations int | None

Number of warmdown iterations. Ranger21 performs Explore-exploit learning rate scheduling.

None
warm_down_min_lr float

Learning rate at the end of warmdown.

3e-05
agc_clipping_value float

Maximum gradient-to-parameter norm ratio for adaptive clipping.

0.01
agc_eps float

Lower bound for the parameter norm in adaptive clipping.

0.001
centralize_gradients bool

Use GC both convolution & fc layers.

True
normalize_gradients bool

Use gradient normalization.

True
lookahead_merge_time int

Steps between Lookahead slow weight updates.

5
lookahead_blending_alpha float

Interpolation factor from slow weights toward fast weights.

0.5
weight_decay float

Weight decay coefficient.

0.0001
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
norm_loss_factor float

Coefficient for the unit norm regularization update.

0.0001
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/ranger21.py
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
class Ranger21(BaseOptimizer):
    """AdamW with positive negative momentum, gradient clipping, and Lookahead.

    Includes gradient centralization and normalization, stable weight decay, norm loss,
    softplus smoothing, and an optional learning rate schedule with warmup and warmdown.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        num_iterations: Total training steps for the built in learning rate schedule.
        lr: Learning rate.
        beta0: Manages the amplitude of the noise introduced by positive negative momentum.
        betas: Decay rates for the first and second moments.
        use_softplus: Use softplus to smooth.
        beta_softplus: Beta parameter for softplus smoothing.
        disable_lr_scheduler: Whether to disable learning rate schedule.
        num_warm_up_iterations: Number of warmup iterations. Ranger21 performs linear learning rate warmup.
        num_warm_down_iterations: Number of warmdown iterations. Ranger21 performs Explore-exploit learning rate
            scheduling.
        warm_down_min_lr: Learning rate at the end of warmdown.
        agc_clipping_value: Maximum gradient-to-parameter norm ratio for adaptive clipping.
        agc_eps: Lower bound for the parameter norm in adaptive clipping.
        centralize_gradients: Use GC both convolution & fc layers.
        normalize_gradients: Use gradient normalization.
        lookahead_merge_time: Steps between Lookahead slow weight updates.
        lookahead_blending_alpha: Interpolation factor from slow weights toward fast weights.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        norm_loss_factor: Coefficient for the unit norm regularization update.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(  # pylint: disable=R0913
        self,
        params: ParamsT,
        num_iterations: int,
        lr: float = 1e-3,
        beta0: float = 0.9,
        betas: Betas = (0.9, 0.999),
        use_softplus: bool = True,
        beta_softplus: float = 50.0,
        disable_lr_scheduler: bool = False,
        num_warm_up_iterations: int | None = None,
        num_warm_down_iterations: int | None = None,
        warm_down_min_lr: float = 3e-5,
        agc_clipping_value: float = 1e-2,
        agc_eps: float = 1e-3,
        centralize_gradients: bool = True,
        normalize_gradients: bool = True,
        lookahead_merge_time: int = 5,
        lookahead_blending_alpha: float = 0.5,
        weight_decay: float = 1e-4,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        norm_loss_factor: float = 1e-4,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_learning_rate(warm_down_min_lr)
        self.validate_betas(betas)
        self.validate_range(beta0, 'beta0', 0.0, 1.0, range_type='[)')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(agc_clipping_value, 'agc_clipping_value')
        self.validate_non_negative(eps, 'eps')
        self.validate_non_negative(agc_eps, 'agc_eps')

        self.min_lr = warm_down_min_lr
        self.beta0 = beta0
        self.use_softplus = use_softplus
        self.beta_softplus = beta_softplus
        self.disable_lr_scheduler = disable_lr_scheduler
        self.agc_clipping_value = agc_clipping_value
        self.agc_eps = agc_eps
        self.centralize_gradients = centralize_gradients
        self.normalize_gradients = normalize_gradients
        self.lookahead_merge_time = lookahead_merge_time
        self.lookahead_blending_alpha = lookahead_blending_alpha
        self.norm_loss_factor = norm_loss_factor
        self.maximize = maximize

        self.lookahead_step: int = 0
        self.starting_lr: float = lr
        self.current_lr: float = lr

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            **kwargs,
        }

        super().__init__(params, defaults)

        self.num_warm_up_iterations: int = (
            self.build_warm_up_iterations(num_iterations, betas[1])
            if num_warm_up_iterations is None
            else num_warm_up_iterations
        )
        self.num_warm_down_iterations: int = (
            self.build_warm_down_iterations(num_iterations)
            if num_warm_down_iterations is None
            else num_warm_down_iterations
        )
        self.start_warm_down: int = num_iterations - self.num_warm_down_iterations
        self.warm_down_lr_delta: float = self.starting_lr - self.min_lr

    def __str__(self) -> str:
        return 'Ranger21'

    def state_dict(self) -> dict:
        state = super().state_dict()
        state['lookahead_step'] = self.lookahead_step
        state['starting_lr'] = self.starting_lr
        state['current_lr'] = self.current_lr
        return state

    def load_state_dict(self, state_dict: dict) -> None:
        super().load_state_dict(state_dict)
        self.lookahead_step = state_dict.get('lookahead_step', 0)
        self.starting_lr = state_dict.get('starting_lr', self.starting_lr)
        self.current_lr = state_dict.get('current_lr', self.current_lr)

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['grad_ma'] = torch.zeros_like(p)
                state['variance_ma'] = torch.zeros_like(p)
                state['lookahead_params'] = p.clone()
                state['neg_grad_ma'] = torch.zeros_like(p)
                state['max_variance_ma'] = torch.zeros_like(p)

    @staticmethod
    def build_warm_up_iterations(total_iterations: int, beta2: float, warm_up_pct: float = 0.22) -> int:
        warm_up_iterations: int = math.ceil(2.0 / (1.0 - beta2))  # default un-tuned linear warmup
        beta_pct: float = warm_up_iterations / total_iterations
        return int(warm_up_pct * total_iterations) if beta_pct > 0.45 else warm_up_iterations

    @staticmethod
    def build_warm_down_iterations(total_iterations: int, warm_down_pct: float = 0.72) -> int:
        start_warm_down: int = int(warm_down_pct * total_iterations)
        return total_iterations - start_warm_down

    def warm_up_dampening(self, lr: float, step: int) -> float:
        if step > self.num_warm_up_iterations:
            return lr

        warm_up_current_pct: float = min(1.0, (step / self.num_warm_up_iterations))

        self.current_lr = lr * warm_up_current_pct

        return self.current_lr

    def warm_down(self, lr: float, iteration: int) -> float:
        if iteration < self.start_warm_down:
            return lr

        # start iteration from 1, not 0
        warm_down_iteration: int = max((iteration + 1) - self.start_warm_down, 1)
        warm_down_pct: float = min(warm_down_iteration / (self.num_warm_down_iterations + 1), 1.0)

        self.current_lr = max(self.starting_lr - self.warm_down_lr_delta * warm_down_pct, self.min_lr)

        return self.current_lr

    def _preprocess_gradient(self, grad: torch.Tensor) -> None:
        if self.centralize_gradients:
            centralize_gradient(grad, gc_conv_only=False)
        if self.normalize_gradients:
            normalize_gradient(grad)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        param_size: int = 0
        variance_ma_sum: float = 1.0

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction2: float = self.debias(beta2, group['step'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                param_size += p.numel()

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                grad.copy_(agc(p, grad, self.agc_eps, self.agc_clipping_value))

                self._preprocess_gradient(grad)

                variance_ma = state['variance_ma']
                variance_ma.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)
                variance_ma_sum += (variance_ma / bias_correction2).sum()

        if param_size == 0:
            raise ZeroParameterSizeError

        variance_normalized = math.sqrt(variance_ma_sum / param_size)

        for group in self.param_groups:
            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            noise_norm: float = math.sqrt((1.0 + self.beta0) ** 2 + self.beta0 ** 2)  # fmt: skip

            if self.disable_lr_scheduler:
                lr: float = group['lr']
            else:
                lr: float = self.warm_up_dampening(group['lr'], group['step'])
                lr = self.warm_down(lr, group['step'])

            step_size: float = self.apply_adam_debias(group.get('adam_debias', False), lr, bias_correction1)

            for p in group['params']:
                if p.grad is None:
                    continue

                self.apply_weight_decay(
                    p=p,
                    grad=None,
                    lr=lr,
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                    ratio=1.0 / variance_normalized,
                )

                correction = 2.0 * self.norm_loss_factor * (1.0 - 1.0 / unit_norm(p).add_(group['eps']))
                p.mul_(1.0 - lr * correction)

                state = self.state[p]
                if group['step'] % 2 == 1:
                    grad_ma, neg_grad_ma = state['grad_ma'], state['neg_grad_ma']
                else:
                    grad_ma, neg_grad_ma = state['neg_grad_ma'], state['grad_ma']

                variance_ma = state['variance_ma']
                max_variance_ma = state['max_variance_ma']
                torch.maximum(max_variance_ma, variance_ma, out=max_variance_ma)

                de_nom = (max_variance_ma.sqrt() / bias_correction2_sq).add_(group['eps'])

                if self.use_softplus:
                    de_nom = softplus(de_nom, beta=self.beta_softplus)

                grad = p.grad
                self._preprocess_gradient(grad)

                grad_ma.lerp_(grad, weight=1.0 - beta1 ** 2)  # fmt: skip

                pn_momentum = grad_ma.mul(1.0 + self.beta0).add_(neg_grad_ma, alpha=-self.beta0)
                pn_momentum.mul_(1.0 / noise_norm)
                p.addcdiv_(pn_momentum, de_nom, value=-step_size)

        self.lookahead_process_step()

        return loss

    def lookahead_process_step(self):
        self.lookahead_step += 1
        if self.lookahead_step >= self.lookahead_merge_time:
            self.lookahead_step: int = 0
            for group in self.param_groups:
                for p in group['params']:
                    if p.grad is None:
                        continue

                    state = self.state[p]

                    p.lerp_(state['lookahead_params'], weight=1.0 - self.lookahead_blending_alpha)
                    state['lookahead_params'].copy_(p)

Ranger25

Bases: BaseOptimizer

Adaptive updates combining ADOPT preconditioning, mixed momentum, and Lookahead.

Includes adaptive gradient clipping and cautious weight decay, with optional cautious updates, OrthoGrad, and StableAdamW or Adam atan2 scaling.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for fast normalized gradient momentum, squared gradients, and slow normalized gradient momentum.

(0.9, 0.98, 0.9999)
weight_decay float

Weight decay coefficient.

0.001
alpha float

Weight of slow momentum relative to fast momentum.

5.0
t_alpha_beta3 float | None

Steps to warm up the slow momentum weight and decay rate. None disables warmup.

None
cautious bool

Whether to use the Cautious variant.

True
stable_adamw bool

Whether to use stable AdamW variant.

True
orthograd bool

Whether to use OrthoGrad variant.

True
eps float | None

Term added to the denominator to improve numerical stability. When eps is None and stable_adamw is False, adam-atan2 feature will be used.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
lookahead_merge_time int

Number of steps between Lookahead slow weight updates.

5
lookahead_blending_alpha float

Interpolation factor from slow weights toward fast weights.

0.5
Source code in pytorch_optimizer/optimizer/experimental/ranger25.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
class Ranger25(BaseOptimizer):
    """Adaptive updates combining ADOPT preconditioning, mixed momentum, and Lookahead.

    Includes adaptive gradient clipping and cautious weight decay, with optional
    cautious updates, OrthoGrad, and StableAdamW or Adam atan2 scaling.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for fast normalized gradient momentum, squared gradients, and slow normalized gradient
            momentum.
        weight_decay: Weight decay coefficient.
        alpha: Weight of slow momentum relative to fast momentum.
        t_alpha_beta3: Steps to warm up the slow momentum weight and decay rate. `None` disables warmup.
        cautious: Whether to use the Cautious variant.
        stable_adamw: Whether to use stable AdamW variant.
        orthograd: Whether to use OrthoGrad variant.
        eps: Term added to the denominator to improve numerical stability. When eps is None and stable_adamw is
            False, adam-atan2 feature will be used.
        maximize: Maximize the objective instead of minimizing it.
        lookahead_merge_time: Number of steps between Lookahead slow weight updates.
        lookahead_blending_alpha: Interpolation factor from slow weights toward fast weights.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.98, 0.9999),
        weight_decay: float = 1e-3,
        alpha: float = 5.0,
        t_alpha_beta3: float | None = None,
        lookahead_merge_time: int = 5,
        lookahead_blending_alpha: float = 0.5,
        cautious: bool = True,
        stable_adamw: bool = True,
        orthograd: bool = True,
        eps: float | None = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(alpha, 'alpha')
        self.validate_non_negative(t_alpha_beta3, 't_alpha_beta3')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_positive(lookahead_merge_time, 'lookahead_merge_time')
        self.validate_range(lookahead_blending_alpha, 'lookahead_blending_alpha', 0.0, 1.0, '[]')
        self.validate_non_negative(eps, 'eps')

        self.lookahead_merge_time = lookahead_merge_time
        self.lookahead_blending_alpha = lookahead_blending_alpha
        self.cautious = cautious
        self.stable_adamw: bool = stable_adamw if isinstance(eps, float) else False
        self.orthograd = orthograd
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'alpha': alpha,
            't_alpha_beta3': t_alpha_beta3,
            'eps': eps if (eps is not None) or (eps is None and not stable_adamw) else 1e-8,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Ranger25'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(grad)
                state['exp_avg_sq'] = torch.zeros_like(grad)
                state['exp_avg_slow'] = torch.zeros_like(grad)
                state['slow_momentum'] = p.clone()

    @staticmethod
    def schedule_alpha(t_alpha_beta3: float | None, step: int, alpha: float) -> float:
        return alpha if t_alpha_beta3 is None else min(step * alpha / t_alpha_beta3, alpha)

    @staticmethod
    def schedule_beta3(t_alpha_beta3: float | None, step: int, beta1: float, beta3: float) -> float:
        if t_alpha_beta3 is None:
            return beta3

        log_beta1, log_beta3 = math.log(beta1), math.log(beta3)

        return min(
            math.exp(
                log_beta1 * log_beta3 / ((1.0 - step / t_alpha_beta3) * log_beta3 + (step / t_alpha_beta3) * log_beta1)
            ),
            beta3,
        )

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        if self.orthograd:
            for group in self.param_groups:
                self.apply_orthogonal_gradients(group['params'])

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2, beta3 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size: float | torch.Tensor = group['lr'] / bias_correction1
            clip: float = math.pow(group['step'], 0.25)

            alpha_t: float = self.schedule_alpha(group['t_alpha_beta3'], group['step'], group['alpha'])
            beta3_t: float = self.schedule_beta3(group['t_alpha_beta3'], group['step'], beta1, beta3)

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                grad.copy_(agc(p, grad))

                exp_avg, exp_avg_sq, exp_avg_slow = state['exp_avg'], state['exp_avg_sq'], state['exp_avg_slow']

                normed_grad = grad.div(
                    exp_avg_sq.sqrt().clamp_(min=group['eps'] if group['eps'] is not None else 1e-8)
                ).clamp_(-clip, clip)

                exp_avg.mul_(beta1).add_(normed_grad, alpha=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)
                exp_avg_slow.mul_(beta3_t).add_(normed_grad, alpha=1.0 - beta3_t)

                update = exp_avg.clone()

                self.apply_cautious_weight_decay(p, update, group['lr'], group['weight_decay'])

                if self.cautious:
                    self.apply_cautious(update, grad)

                if self.stable_adamw:
                    param_step_size = step_size / self.get_stable_adamw_rms(grad, exp_avg_sq)
                else:
                    param_step_size = step_size

                update.add_(exp_avg_slow, alpha=alpha_t)

                de_nom = exp_avg_sq.sqrt().div_(bias_correction2_sq)

                if group['eps'] is not None:
                    de_nom.add_(group['eps'])
                    if self.stable_adamw:
                        de_nom = de_nom.to(dtype=param_step_size.dtype).div_(-param_step_size)
                        p.addcdiv_(update, de_nom)
                    else:
                        p.addcdiv_(update, de_nom, value=-param_step_size)
                else:
                    p.add_(update.atan2_(de_nom), alpha=-param_step_size)

                if group['step'] % self.lookahead_merge_time == 0:
                    slow_p = state['slow_momentum']
                    slow_p.lerp_(p, weight=self.lookahead_blending_alpha)
                    p.copy_(slow_p)

        return loss

ROSE

Bases: BaseOptimizer

Gradient updates scaled by row and column ranges.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
weight_decay float

Weight decay coefficient.

0.0001
wd_schedule bool | float

Scale decoupled decay by lr / lr_ref. A float supplies lr_ref. True reads max_lr or initial_lr from the group. False uses lr.

False
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
centralize bool

Subtract the mean of each gradient slice for parameters with two or more dimensions.

True
stabilize bool

Blend local range scaling with the global mean range using coefficient of variation gating.

True
bf16_sr bool

Use stochastic rounding for bfloat16 parameter updates.

True
compute_dtype dtype

Data type for gradient and parameter update arithmetic.

float64
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/rose.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
class ROSE(BaseOptimizer):
    """Gradient updates scaled by row and column ranges.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        weight_decay: Weight decay coefficient.
        wd_schedule: Scale decoupled decay by `lr / lr_ref`. A float supplies `lr_ref`. `True` reads `max_lr` or
            `initial_lr` from the group. `False` uses `lr`.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        centralize: Subtract the mean of each gradient slice for parameters with two or more dimensions.
        stabilize: Blend local range scaling with the global mean range using coefficient of variation gating.
        bf16_sr: Use stochastic rounding for bfloat16 parameter updates.
        compute_dtype: Data type for gradient and parameter update arithmetic.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        weight_decay: float = 1e-4,
        wd_schedule: bool | float = False,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        centralize: bool = True,
        stabilize: bool = True,
        bf16_sr: bool = True,
        compute_dtype: torch.dtype = torch.float64,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.maximize = maximize

        if bf16_sr and compute_dtype not in (torch.float32, torch.float64, None):
            raise ValueError(f'bf16_sr=True has no useful effect when compute_dtype is {compute_dtype}.')

        defaults: Defaults = {
            'lr': lr,
            'weight_decay': weight_decay,
            'wd_schedule': wd_schedule,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'centralize': centralize,
            'stabilize': stabilize,
            'bf16_sr': bf16_sr,
            'compute_dtype': compute_dtype,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'ROSE'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            lr = group['lr']
            weight_decay, wd_schedule = group['weight_decay'], group['wd_schedule']
            compute_dtype = group['compute_dtype']

            if weight_decay and wd_schedule:
                wd_lr = lr / (
                    wd_schedule if isinstance(wd_schedule, float) else group.get('max_lr', group.get('initial_lr'))
                )
            else:
                wd_lr = lr

            for p in group['params']:
                if p.grad is None:
                    continue

                use_bf16_sr = group['bf16_sr'] and p.dtype is torch.bfloat16
                fp32 = use_bf16_sr and not compute_dtype

                grad = p.grad.to(dtype=torch.float32 if fp32 else compute_dtype)
                param = p.to(dtype=torch.float32 if fp32 else compute_dtype)

                self.maximize_gradient(grad, maximize=self.maximize)

                self.apply_weight_decay(
                    param,
                    grad,
                    lr=wd_lr,
                    weight_decay=weight_decay,
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                if grad.ndim == 0:
                    param.add_(grad.sign(), alpha=-lr)
                elif grad.ndim == 1:
                    g_min, g_max = grad.aminmax()
                    de_nom = g_max.abs_().sub_(g_min)

                    de_nom.masked_fill_(de_nom == 0.0, 1.0)
                    param.addcdiv_(grad, de_nom, value=-lr)
                else:
                    active_axes = tuple(range(1, grad.ndim))

                    if group['centralize']:
                        if grad is not p.grad:
                            grad.sub_(grad.mean(dim=active_axes, keepdim=True))
                        else:
                            grad = grad.sub(grad.mean(dim=active_axes, keepdim=True))

                    raw_scale = (
                        grad.amax(dim=active_axes, keepdim=True).abs_().sub_(grad.amin(dim=active_axes, keepdim=True))
                    )

                    if group['stabilize']:
                        std, mean = torch.std_mean(raw_scale, correction=0)

                        trust = mean.div(std.add_(mean).masked_fill_(mean == 0.0, 1.0))

                        de_nom = mean.lerp(raw_scale, trust)
                    else:
                        de_nom = raw_scale

                    de_nom.masked_fill_(de_nom == 0.0, 1.0)
                    param.addcdiv_(grad, de_nom, value=-lr)

                if use_bf16_sr:
                    param = param.to(dtype=torch.float32)

                    copy_stochastic(p, param)
                elif param is not p:
                    p.copy_(param)

        return loss

RotoGrad

Bases: RotateOnly

Balance multitask gradient directions and magnitudes with learned rotations.

Parameters:

Name Type Description Default
backbone Module

Shared model producing the latent representation.

required
heads Sequence[Module]

Task specific models consuming the latent representation.

required
latent_size int

Number of features in the shared representation.

required
*args tuple

Additional positional arguments, accepted without effect.

()
burn_in_period int

Ignored. This variant uses the base default of 20 steps.

20
normalize_losses bool

Normalize losses when computing gradients for the task specific heads.

False
Source code in pytorch_optimizer/optimizer/rotograd.py
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
class RotoGrad(RotateOnly):
    """Balance multitask gradient directions and magnitudes with learned rotations.

    Args:
        backbone: Shared model producing the latent representation.
        heads: Task specific models consuming the latent representation.
        latent_size: Number of features in the shared representation.
        *args (tuple): Additional positional arguments, accepted without effect.
        burn_in_period: Ignored. This variant uses the base default of 20 steps.
        normalize_losses: Normalize losses when computing gradients for the task specific heads.

    """

    num_tasks: int
    backbone: nn.Module
    heads: Sequence[nn.Module]
    rep: torch.Tensor

    def __init__(
        self,
        backbone: nn.Module,
        heads: Sequence[nn.Module],
        latent_size: int,
        *args,
        burn_in_period: int = 20,
        normalize_losses: bool = False,
    ):
        super().__init__(backbone, heads, latent_size, burn_in_period, *args, normalize_losses=normalize_losses)

        self.initial_grads = None
        self.counter: int = 0

    def _rep_grad(self):
        super()._rep_grad()

        grad_norms = [torch.linalg.norm(g, keepdim=True).clamp_min(1e-15) for g in self.original_grads]
        if self.initial_grads is None or self.counter == self.burn_in_period:
            self.initial_grads = grad_norms
            conv_ratios = [torch.ones((1,)) for _ in range(len(self.initial_grads))]
        else:
            conv_ratios = [x / y for x, y in zip(grad_norms, self.initial_grads)]

        self.counter += 1

        alphas = [x / torch.clamp(sum(conv_ratios), 1e-15) for x in conv_ratios]
        weighted_sum_norms = sum(a * g for a, g in zip(alphas, grad_norms))

        return sum(g / n * weighted_sum_norms for g, n in zip(self.original_grads, grad_norms))

SafeFP16Optimizer

Bases: Optimizer

Wrap an optimizer with float32 master weights and dynamic loss scaling.

Supports a single parameter group. Use backward() to scale the loss and clip_main_grads() to check for overflow before updating parameters.

Parameters:

Name Type Description Default
optimizer Optimizer

Base optimizer instance with low precision parameters.

required
aggregate_g_norms bool

Aggregate squared gradient norms across distributed workers.

False
min_loss_scale float

Scale below which persistent overflow raises FloatingPointError.

2 ** -5
Source code in pytorch_optimizer/optimizer/fp16.py
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
class SafeFP16Optimizer(Optimizer):  # pragma: no cover
    """Wrap an optimizer with float32 master weights and dynamic loss scaling.

    Supports a single parameter group. Use `backward()` to scale the loss and
    `clip_main_grads()` to check for overflow before updating parameters.

    Args:
        optimizer: Base optimizer instance with low precision parameters.
        aggregate_g_norms: Aggregate squared gradient norms across distributed workers.
        min_loss_scale: Scale below which persistent overflow raises `FloatingPointError`.

    """

    def __init__(
        self,
        optimizer: Optimizer,
        aggregate_g_norms: bool = False,
        min_loss_scale: float = 2 ** -5,
    ) -> None:  # fmt: skip
        self.optimizer = optimizer
        self.aggregate_g_norms = aggregate_g_norms
        self.min_loss_scale = min_loss_scale

        self.fp16_params = self.get_parameters(optimizer)
        self.fp32_params = self.build_fp32_params(self.fp16_params, flatten=False)

        # we want the optimizer to be tracking the fp32 parameters
        if len(optimizer.param_groups) != 1:
            # future implementers: this should hopefully be a matter of just iterating through the param groups and
            # keeping track of the pointer through the fp32_params
            raise NotImplementedError('Need to implement the parameter group transfer.')

        optimizer.param_groups[0]['params'] = self.fp32_params

        self.scaler: DynamicLossScaler = DynamicLossScaler(2.0 ** 15)  # fmt: skip
        self.needs_sync: bool = True

    @classmethod
    def get_parameters(cls, optimizer: Optimizer) -> list:
        params: list = []
        for group in optimizer.param_groups:
            params += list(group['params'])
        return params

    @classmethod
    def build_fp32_params(cls, parameters: ParamsT, flatten: bool = True) -> torch.Tensor | list[torch.Tensor]:
        parameters = cast(list[torch.Tensor], parameters)

        if flatten:
            total_param_size: int = sum(p.numel() for p in parameters)
            fp32_params = torch.zeros(total_param_size, dtype=torch.float, device=parameters[0].device)

            offset: int = 0
            for p in parameters:
                p_num_el = p.numel()
                fp32_params[offset:offset + p_num_el].copy_(p.view(-1))  # fmt: skip
                offset += p_num_el

            fp32_params = nn.Parameter(fp32_params)
            fp32_params.grad = fp32_params.new(total_param_size)

            return fp32_params

        fp32_params: list[torch.Tensor] = []
        for p in parameters:
            p32 = nn.Parameter(p.float())
            p32.grad = torch.zeros_like(p32)
            fp32_params.append(p32)

        return fp32_params

    def state_dict(self) -> dict:
        """Return the optimizer state dict."""
        state_dict = self.optimizer.state_dict()
        state_dict['fp32_params'] = [p.detach().clone() for p in self.fp32_params]
        if self.scaler is not None:
            state_dict['loss_scaler'] = self.scaler.loss_scale
        return state_dict

    def load_state_dict(self, state_dict: dict):
        """Restore the base optimizer state and loss scale from a checkpoint.

        Args:
            state_dict: Checkpoint state, including saved parameter group options.

        """
        if 'loss_scaler' in state_dict and self.scaler is not None and isinstance(state_dict['loss_scaler'], float):
            self.scaler.loss_scale = state_dict['loss_scaler']
        self.optimizer.load_state_dict(state_dict)
        if 'fp32_params' in state_dict:
            if len(state_dict['fp32_params']) != len(self.fp32_params):
                raise ValueError('master weights do not match the current parameters')
            with torch.no_grad():
                for p, saved in zip(self.fp32_params, state_dict['fp32_params']):
                    p.copy_(saved)

    def backward(self, loss, update_main_grads: bool = False):
        """Scale the loss and compute low precision parameter gradients.

        Args:
            loss (torch.Tensor): Scalar loss tensor to backpropagate.
            update_main_grads: Copy and unscale gradients into the float32 master buffers after backward.

        """
        if self.scaler is not None:
            loss = loss * self.scaler.loss_scale

        loss.backward()

        self.needs_sync = True
        if update_main_grads:
            self.update_main_grads()

    def sync_fp16_grads_to_fp32(self, multiply_grads: float = 1.0) -> None:
        """Copy and unscale low precision gradients into float32 master buffers."""
        if self.needs_sync:
            if self.scaler is not None:
                multiply_grads /= self.scaler.loss_scale

            for p16, p32 in zip(self.fp16_params, self.fp32_params):
                if not p16.requires_grad:
                    continue

                if p16.grad is not None:
                    if p32.grad is None:
                        p32.grad = torch.empty_like(p32)
                    p32.grad.copy_(p16.grad)
                    p32.grad.mul_(multiply_grads)
                else:
                    p32.grad = None

            self.needs_sync = False

    def multiply_grads(self, c: float) -> None:
        """Multiply the float32 master gradients by `c`."""
        if self.needs_sync:
            self.sync_fp16_grads_to_fp32(c)
            return

        for p32 in self.fp32_params:
            if p32.grad is not None:
                p32.grad.mul_(c)

    def update_main_grads(self) -> None:
        self.sync_fp16_grads_to_fp32()

    def clip_main_grads(self, max_norm: float):
        """Clip master gradients and update the loss scale after checking for overflow."""
        self.sync_fp16_grads_to_fp32()

        grad_norm = clip_grad_norm(self.fp32_params, max_norm, sync=self.aggregate_g_norms)

        # detect overflow and adjust loss scale
        if self.scaler is not None:
            overflow: bool = has_overflow(grad_norm)
            prev_scale: float = self.scaler.loss_scale

            self.scaler.update_scale(overflow)

            if overflow:
                self.zero_grad()
                if self.scaler.loss_scale <= self.min_loss_scale:
                    # Use FloatingPointError as an uncommon error
                    # that parent functions can safely catch to stop training.
                    self.scaler.loss_scale = prev_scale

                    raise FloatingPointError(
                        f'Minimum loss scale reached ({self.min_loss_scale}). Your loss is probably exploding. '
                        'Try lowering the learning rate, using gradient clipping or increasing the batch size.\n'
                        f'Overflow: setting loss scale to {self.scaler.loss_scale}'
                    )

        return grad_norm

    def step(self, closure: Closure = None):
        """Perform a single optimization step."""
        self.sync_fp16_grads_to_fp32()
        self.optimizer.step(closure)

        for p16, p32 in zip(self.fp16_params, self.fp32_params):
            if not p16.requires_grad:
                continue
            p16.data.copy_(p32)

    def zero_grad(self) -> None:
        """Clear the gradients of all optimized parameters."""
        for p16 in self.fp16_params:
            p16.grad = None
        for p32 in self.fp32_params:
            p32.grad = None
        self.needs_sync = False

    def get_lr(self) -> float:
        """Return the base optimizer learning rate."""
        return self.optimizer.param_groups[0]['lr']

    def set_lr(self, lr: float):
        """Set the base optimizer learning rate."""
        for group in self.optimizer.param_groups:
            group['lr'] = lr

    @property
    def loss_scale(self) -> float:
        """Return the current dynamic loss scale."""
        return self.scaler.loss_scale

loss_scale property

Return the current dynamic loss scale.

backward(loss, update_main_grads=False)

Scale the loss and compute low precision parameter gradients.

Parameters:

Name Type Description Default
loss Tensor

Scalar loss tensor to backpropagate.

required
update_main_grads bool

Copy and unscale gradients into the float32 master buffers after backward.

False
Source code in pytorch_optimizer/optimizer/fp16.py
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
def backward(self, loss, update_main_grads: bool = False):
    """Scale the loss and compute low precision parameter gradients.

    Args:
        loss (torch.Tensor): Scalar loss tensor to backpropagate.
        update_main_grads: Copy and unscale gradients into the float32 master buffers after backward.

    """
    if self.scaler is not None:
        loss = loss * self.scaler.loss_scale

    loss.backward()

    self.needs_sync = True
    if update_main_grads:
        self.update_main_grads()

clip_main_grads(max_norm)

Clip master gradients and update the loss scale after checking for overflow.

Source code in pytorch_optimizer/optimizer/fp16.py
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
def clip_main_grads(self, max_norm: float):
    """Clip master gradients and update the loss scale after checking for overflow."""
    self.sync_fp16_grads_to_fp32()

    grad_norm = clip_grad_norm(self.fp32_params, max_norm, sync=self.aggregate_g_norms)

    # detect overflow and adjust loss scale
    if self.scaler is not None:
        overflow: bool = has_overflow(grad_norm)
        prev_scale: float = self.scaler.loss_scale

        self.scaler.update_scale(overflow)

        if overflow:
            self.zero_grad()
            if self.scaler.loss_scale <= self.min_loss_scale:
                # Use FloatingPointError as an uncommon error
                # that parent functions can safely catch to stop training.
                self.scaler.loss_scale = prev_scale

                raise FloatingPointError(
                    f'Minimum loss scale reached ({self.min_loss_scale}). Your loss is probably exploding. '
                    'Try lowering the learning rate, using gradient clipping or increasing the batch size.\n'
                    f'Overflow: setting loss scale to {self.scaler.loss_scale}'
                )

    return grad_norm

get_lr()

Return the base optimizer learning rate.

Source code in pytorch_optimizer/optimizer/fp16.py
280
281
282
def get_lr(self) -> float:
    """Return the base optimizer learning rate."""
    return self.optimizer.param_groups[0]['lr']

load_state_dict(state_dict)

Restore the base optimizer state and loss scale from a checkpoint.

Parameters:

Name Type Description Default
state_dict dict

Checkpoint state, including saved parameter group options.

required
Source code in pytorch_optimizer/optimizer/fp16.py
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
def load_state_dict(self, state_dict: dict):
    """Restore the base optimizer state and loss scale from a checkpoint.

    Args:
        state_dict: Checkpoint state, including saved parameter group options.

    """
    if 'loss_scaler' in state_dict and self.scaler is not None and isinstance(state_dict['loss_scaler'], float):
        self.scaler.loss_scale = state_dict['loss_scaler']
    self.optimizer.load_state_dict(state_dict)
    if 'fp32_params' in state_dict:
        if len(state_dict['fp32_params']) != len(self.fp32_params):
            raise ValueError('master weights do not match the current parameters')
        with torch.no_grad():
            for p, saved in zip(self.fp32_params, state_dict['fp32_params']):
                p.copy_(saved)

multiply_grads(c)

Multiply the float32 master gradients by c.

Source code in pytorch_optimizer/optimizer/fp16.py
221
222
223
224
225
226
227
228
229
def multiply_grads(self, c: float) -> None:
    """Multiply the float32 master gradients by `c`."""
    if self.needs_sync:
        self.sync_fp16_grads_to_fp32(c)
        return

    for p32 in self.fp32_params:
        if p32.grad is not None:
            p32.grad.mul_(c)

set_lr(lr)

Set the base optimizer learning rate.

Source code in pytorch_optimizer/optimizer/fp16.py
284
285
286
287
def set_lr(self, lr: float):
    """Set the base optimizer learning rate."""
    for group in self.optimizer.param_groups:
        group['lr'] = lr

state_dict()

Return the optimizer state dict.

Source code in pytorch_optimizer/optimizer/fp16.py
159
160
161
162
163
164
165
def state_dict(self) -> dict:
    """Return the optimizer state dict."""
    state_dict = self.optimizer.state_dict()
    state_dict['fp32_params'] = [p.detach().clone() for p in self.fp32_params]
    if self.scaler is not None:
        state_dict['loss_scaler'] = self.scaler.loss_scale
    return state_dict

step(closure=None)

Perform a single optimization step.

Source code in pytorch_optimizer/optimizer/fp16.py
262
263
264
265
266
267
268
269
270
def step(self, closure: Closure = None):
    """Perform a single optimization step."""
    self.sync_fp16_grads_to_fp32()
    self.optimizer.step(closure)

    for p16, p32 in zip(self.fp16_params, self.fp32_params):
        if not p16.requires_grad:
            continue
        p16.data.copy_(p32)

sync_fp16_grads_to_fp32(multiply_grads=1.0)

Copy and unscale low precision gradients into float32 master buffers.

Source code in pytorch_optimizer/optimizer/fp16.py
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
def sync_fp16_grads_to_fp32(self, multiply_grads: float = 1.0) -> None:
    """Copy and unscale low precision gradients into float32 master buffers."""
    if self.needs_sync:
        if self.scaler is not None:
            multiply_grads /= self.scaler.loss_scale

        for p16, p32 in zip(self.fp16_params, self.fp32_params):
            if not p16.requires_grad:
                continue

            if p16.grad is not None:
                if p32.grad is None:
                    p32.grad = torch.empty_like(p32)
                p32.grad.copy_(p16.grad)
                p32.grad.mul_(multiply_grads)
            else:
                p32.grad = None

        self.needs_sync = False

zero_grad()

Clear the gradients of all optimized parameters.

Source code in pytorch_optimizer/optimizer/fp16.py
272
273
274
275
276
277
278
def zero_grad(self) -> None:
    """Clear the gradients of all optimized parameters."""
    for p16 in self.fp16_params:
        p16.grad = None
    for p32 in self.fp32_params:
        p32.grad = None
    self.needs_sync = False

SAM

Bases: BaseOptimizer

Sharpness-aware minimization with a two pass parameter update.

Compute gradients at the current weights before calling step(). The closure must recompute the loss and gradients at the perturbed weights.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
base_optimizer OptimizerType

Optimizer class to instantiate for the parameter update.

required
rho float

Radius of the neighborhood used to perturb parameters.

0.05
use_gc bool

Centralize gradients before perturbing parameters.

False
adaptive bool

Scale perturbations by the squared parameter values.

False
perturb_eps float

Stability constant for the perturbation norm.

1e-12
**kwargs dict

Options for the base optimizer.

{}

Examples:

optimizer = SAM(model.parameters(), torch.optim.AdamW, lr=1e-3)
for inputs, targets in data:
    optimizer.zero_grad()

    def closure():
        optimizer.zero_grad()
        loss = loss_fn(model(inputs), targets)
        loss.backward()
        return loss

    closure()
    optimizer.step(closure)
Source code in pytorch_optimizer/optimizer/sam.py
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
class SAM(BaseOptimizer):
    """Sharpness-aware minimization with a two pass parameter update.

    Compute gradients at the current weights before calling `step()`. The closure
    must recompute the loss and gradients at the perturbed weights.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        base_optimizer: Optimizer class to instantiate for the parameter update.
        rho: Radius of the neighborhood used to perturb parameters.
        use_gc: Centralize gradients before perturbing parameters.
        adaptive: Scale perturbations by the squared parameter values.
        perturb_eps: Stability constant for the perturbation norm.
        **kwargs (dict): Options for the base optimizer.

    Examples:
        ```python
        optimizer = SAM(model.parameters(), torch.optim.AdamW, lr=1e-3)
        for inputs, targets in data:
            optimizer.zero_grad()

            def closure():
                optimizer.zero_grad()
                loss = loss_fn(model(inputs), targets)
                loss.backward()
                return loss

            closure()
            optimizer.step(closure)
        ```

    """

    def __init__(
        self,
        params: ParamsT,
        base_optimizer: OptimizerType,
        rho: float = 0.05,
        adaptive: bool = False,
        use_gc: bool = False,
        perturb_eps: float = 1e-12,
        **kwargs,
    ):
        self.validate_non_negative(rho, 'rho')
        self.validate_non_negative(perturb_eps, 'perturb_eps')

        self.use_gc = use_gc
        self.perturb_eps = perturb_eps

        defaults: Defaults = {'rho': rho, 'adaptive': adaptive, **kwargs}

        super().__init__(params, defaults)

        self.base_optimizer: Optimizer = base_optimizer(self.param_groups, **kwargs)
        self.param_groups = self.base_optimizer.param_groups
        self.state = self.base_optimizer.state

    def __str__(self) -> str:
        return 'SAM'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

    @torch.no_grad()
    def first_step(self, zero_grad: bool = False):
        grad_norm = get_global_gradient_norm(self.param_groups, weight_adaptive=True)
        grad_norm.sqrt_().squeeze_(0).add_(self.perturb_eps)

        for group in self.param_groups:
            scale = group['rho'] / grad_norm

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad
                if self.use_gc:
                    centralize_gradient(grad, gc_conv_only=False)

                self.state[p]['old_p'] = p.clone()

                e_w = (torch.pow(p, 2) if group['adaptive'] else 1.0) * grad * scale.to(p)

                p.add_(e_w)

        if zero_grad:
            self.zero_grad()

    @torch.no_grad()
    def second_step(self, zero_grad: bool = False):
        for group in self.param_groups:
            for p in group['params']:
                if 'old_p' in self.state[p]:
                    p.copy_(self.state[p].pop('old_p'))

        self.base_optimizer.step()

        if zero_grad:
            self.zero_grad()

    @torch.no_grad()
    def step(self, closure: Closure = None):
        """Perturb weights, recompute gradients, and apply the base optimizer update.

        Args:
            closure: Callable that clears gradients and recomputes the loss and gradients. Compute the initial
                gradients before calling this method.

        Raises:
            NoClosureError: No closure is supplied.

        """
        if closure is None:
            raise NoClosureError(str(self))

        self.first_step(zero_grad=True)

        with torch.enable_grad():
            closure()

        self.second_step()

    def load_state_dict(self, state_dict: dict):
        super().load_state_dict(state_dict)
        self.base_optimizer.param_groups = self.param_groups
        self.base_optimizer.state = self.state  # ty: ignore[invalid-assignment]

step(closure=None)

Perturb weights, recompute gradients, and apply the base optimizer update.

Parameters:

Name Type Description Default
closure Closure

Callable that clears gradients and recomputes the loss and gradients. Compute the initial gradients before calling this method.

None

Raises:

Type Description
NoClosureError

No closure is supplied.

Source code in pytorch_optimizer/optimizer/sam.py
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
@torch.no_grad()
def step(self, closure: Closure = None):
    """Perturb weights, recompute gradients, and apply the base optimizer update.

    Args:
        closure: Callable that clears gradients and recomputes the loss and gradients. Compute the initial
            gradients before calling this method.

    Raises:
        NoClosureError: No closure is supplied.

    """
    if closure is None:
        raise NoClosureError(str(self))

    self.first_step(zero_grad=True)

    with torch.enable_grad():
        closure()

    self.second_step()

SaRA

Bases: BaseOptimizer

AdamW updates on progressively refined masks of small magnitude weights.

Implements the parameter based reference in sjtuplayer/SaRA/optim/adamw2.py. Only weights with initial absolute values below threshold receive AdamW updates, including weight decay. Moments are stored only for these weights. Mask refinement keeps their accumulated moments. This optimizer uses ordinary PyTorch backpropagation rather than the paper's model reparameterization for unstructured backpropagation.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float | None

Learning rate. Defaults to 1e-3 * exp(-350 * threshold) when None.

None
betas Betas

Coefficients used for computing running averages of gradient and its square.

(0.9, 0.999)
threshold float

Strict upper bound on the absolute values of initially trainable weights.

0.001
progressive_iter int

Refine the mask before update progressive_iter + 1. -1 disables refinement.

-1
lambda_rank float

Nuclear norm penalty coefficient. Each step samples one matrix per parameter group with both dimensions greater than 64 and applies the penalty if more than 100 weights are selected.

0.0
weight_decay float

Weight decay coefficient.

0.01
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/sara.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
class SaRA(BaseOptimizer):
    """AdamW updates on progressively refined masks of small magnitude weights.

    Implements the parameter based reference in `sjtuplayer/SaRA/optim/adamw2.py`. Only weights with initial absolute
    values below `threshold` receive AdamW updates, including weight decay. Moments are stored only for these weights.
    Mask refinement keeps their accumulated moments. This optimizer uses ordinary PyTorch backpropagation rather than
    the paper's model reparameterization for unstructured backpropagation.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate. Defaults to `1e-3 * exp(-350 * threshold)` when None.
        betas: Coefficients used for computing running averages of gradient and its square.
        threshold: Strict upper bound on the absolute values of initially trainable weights.
        progressive_iter: Refine the mask before update `progressive_iter + 1`. -1 disables refinement.
        lambda_rank: Nuclear norm penalty coefficient. Each step samples one matrix per parameter group with both
            dimensions greater than 64 and applies the penalty if more than 100 weights are selected.
        weight_decay: Weight decay coefficient.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float | None = None,
        betas: Betas = (0.9, 0.999),
        threshold: float = 1e-3,
        progressive_iter: int = -1,
        lambda_rank: float = 0.0,
        weight_decay: float = 1e-2,
        ams_bound: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(threshold, 'threshold')
        self.validate_boundary(progressive_iter, -1, bound_type='lower')
        self.validate_non_negative(lambda_rank, 'lambda_rank')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        defaults: Defaults = {
            'lr': 1e-3 * math.exp(-350.0 * threshold) if lr is None else lr,
            'betas': betas,
            'threshold': threshold,
            'progressive_iter': progressive_iter,
            'lambda_rank': lambda_rank,
            'weight_decay': weight_decay,
            'ams_bound': ams_bound,
            'eps': eps,
            'maximize': maximize,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SaRA'

    def add_param_group(self, param_group: ParamGroup) -> None:
        super().add_param_group(param_group)

        group = self.param_groups[-1]
        for p in group['params']:
            self.state[p]['mask'] = p.detach().abs() < group['threshold']

    def load_state_dict(self, state_dict: State) -> None:
        super().load_state_dict(state_dict)

        for group in self.param_groups:
            for p in group['params']:
                self.state[p]['mask'] = self.state[p]['mask'].bool()

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            if p.grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]
            if 'exp_avg' not in state:
                state['exp_avg'] = torch.zeros_like(p[state['mask']])
                state['exp_avg_sq'] = torch.zeros_like(state['exp_avg'])

                if group['ams_bound']:
                    state['max_exp_avg_sq'] = torch.zeros_like(state['exp_avg'])

    @torch.no_grad()
    def update_mask(self, group: ParamGroup) -> None:
        """Keep only initially selected weights that are still below the threshold and retain their moments."""
        for p in group['params']:
            state = self.state[p]

            mask = state['mask']
            keep = p[mask].abs() < group['threshold']

            state['mask'] = mask.clone()
            state['mask'][mask] = keep

            for key in ('exp_avg', 'exp_avg_sq', 'max_exp_avg_sq'):
                if key in state:
                    state[key] = state[key][keep]

    @staticmethod
    @torch.no_grad()
    def apply_rank_constraint(
        p: torch.Tensor, grad_mask: torch.Tensor, mask: torch.Tensor, lambda_rank: float
    ) -> None:
        """Add the masked nuclear norm subgradient to the selected gradient entries."""
        dtype = torch.float32 if p.dtype in (torch.float16, torch.bfloat16) else p.dtype

        matrix = torch.where(mask, p, 0.0).to(dtype=dtype)

        u, _, vh = torch.linalg.svd(matrix, full_matrices=False)
        torch.mm(u, vh, out=matrix)

        grad_mask.add_(matrix[mask].to(dtype=grad_mask.dtype), alpha=lambda_rank)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size: float = group['lr'] / bias_correction1

            if group['progressive_iter'] >= 0 and group['step'] == group['progressive_iter'] + 1:
                self.update_mask(group)

            rank_params = [
                p
                for p in group['params']
                if group['lambda_rank'] > 0.0 and p.grad is not None and p.dim() == 2 and min(p.shape) > 64
            ]
            rank_param = random.choice(rank_params) if rank_params else None  # noqa: S311

            for p in group['params']:
                if p.grad is None:
                    continue

                state = self.state[p]

                mask = state['mask']

                grad_mask = p.grad[mask]

                if p is rank_param and mask.sum() > 100:
                    self.apply_rank_constraint(p, grad_mask, mask, group['lambda_rank'])

                self.maximize_gradient(grad_mask, maximize=group['maximize'])

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']
                exp_avg.lerp_(grad_mask, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad_mask, grad_mask, value=1.0 - beta2)

                de_nom = self.apply_ams_bound(
                    ams_bound=group['ams_bound'],
                    exp_avg_sq=exp_avg_sq,
                    max_exp_avg_sq=state.get('max_exp_avg_sq'),
                    eps=0.0,
                    exp_avg_sq_eps=0.0,
                )
                de_nom.div_(bias_correction2_sq).add_(group['eps'])

                p_mask = p[mask]

                self.apply_weight_decay(
                    p=p_mask,
                    grad=grad_mask,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=True,
                    fixed_decay=False,
                )

                p_mask.addcdiv_(exp_avg, de_nom, value=-step_size)
                p[mask] = p_mask

        return loss

apply_rank_constraint(p, grad_mask, mask, lambda_rank) staticmethod

Add the masked nuclear norm subgradient to the selected gradient entries.

Source code in pytorch_optimizer/optimizer/sara.py
126
127
128
129
130
131
132
133
134
135
136
137
138
139
@staticmethod
@torch.no_grad()
def apply_rank_constraint(
    p: torch.Tensor, grad_mask: torch.Tensor, mask: torch.Tensor, lambda_rank: float
) -> None:
    """Add the masked nuclear norm subgradient to the selected gradient entries."""
    dtype = torch.float32 if p.dtype in (torch.float16, torch.bfloat16) else p.dtype

    matrix = torch.where(mask, p, 0.0).to(dtype=dtype)

    u, _, vh = torch.linalg.svd(matrix, full_matrices=False)
    torch.mm(u, vh, out=matrix)

    grad_mask.add_(matrix[mask].to(dtype=grad_mask.dtype), alpha=lambda_rank)

update_mask(group)

Keep only initially selected weights that are still below the threshold and retain their moments.

Source code in pytorch_optimizer/optimizer/sara.py
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
@torch.no_grad()
def update_mask(self, group: ParamGroup) -> None:
    """Keep only initially selected weights that are still below the threshold and retain their moments."""
    for p in group['params']:
        state = self.state[p]

        mask = state['mask']
        keep = p[mask].abs() < group['threshold']

        state['mask'] = mask.clone()
        state['mask'][mask] = keep

        for key in ('exp_avg', 'exp_avg_sq', 'max_exp_avg_sq'):
            if key in state:
                state[key] = state[key][keep]

ScalableShampoo

Bases: BaseOptimizer

Shampoo with blockwise preconditioning and optional gradient grafting.

Compute matrix inverse roots with SVD or coupled Schur-Newton iteration on the parameter device. Grafting uses the update norm of SGD, AdaGrad, or RMSProp.

Reference: https://github.com/google-research/google-research/blob/master/scalable_shampoo/optax/distributed_shampoo.py

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for momentum and preconditioner statistics.

(0.9, 0.999)
moving_average_for_momentum bool

Whether to perform moving average for momentum (beta1).

False
weight_decay float

Weight decay coefficient.

0.0
decoupled_weight_decay bool

Use decoupled weight decay.

False
decoupled_learning_rate bool

Use decoupled learning rate, otherwise coupled with preconditioned gradient.

True
inverse_exponent_override int

Fixed exponent for preconditioner if > 0.

0
start_preconditioning_step int

Step to start preconditioning.

25
preconditioning_compute_steps int

Frequency of preconditioner computation.

1000
statistics_compute_steps int

Frequency of statistics computation.

1
block_size int

Block size for large layers. 1 means AdaGrad (inefficient).

512
skip_preconditioning_rank_lt int

Skip preconditioning for layers with rank below this.

1
no_preconditioning_for_layers_with_dim_gt int

Avoid preconditioning large layers.

8192
shape_interpretation bool

Automatic shape interpretation for tensor dims.

True
graft_type int

Layer wise scale reference from LayerWiseGrafting.

SGD
pre_conditioner_type int

Dimensions to precondition, from PreConditionerType.

ALL
nesterov bool

Use Nesterov momentum.

True
diagonal_eps float

Epsilon for numerical stability in diagonal.

1e-10
matrix_eps float

Epsilon for numerical stability in matrix.

1e-06
use_svd bool

Whether to use SVD for matrix inverse powers (alternative is Schur-Newton).

False
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/shampoo.py
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
class ScalableShampoo(BaseOptimizer):
    """Shampoo with blockwise preconditioning and optional gradient grafting.

    Compute matrix inverse roots with SVD or coupled Schur-Newton iteration on the
    parameter device. Grafting uses the update norm of SGD, AdaGrad, or RMSProp.

    Reference: https://github.com/google-research/google-research/blob/master/scalable_shampoo/optax/distributed_shampoo.py

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for momentum and preconditioner statistics.
        moving_average_for_momentum: Whether to perform moving average for momentum (beta1).
        weight_decay: Weight decay coefficient.
        decoupled_weight_decay: Use decoupled weight decay.
        decoupled_learning_rate: Use decoupled learning rate, otherwise coupled with preconditioned gradient.
        inverse_exponent_override: Fixed exponent for preconditioner if > 0.
        start_preconditioning_step: Step to start preconditioning.
        preconditioning_compute_steps: Frequency of preconditioner computation.
        statistics_compute_steps: Frequency of statistics computation.
        block_size: Block size for large layers. 1 means AdaGrad (inefficient).
        skip_preconditioning_rank_lt: Skip preconditioning for layers with rank below this.
        no_preconditioning_for_layers_with_dim_gt: Avoid preconditioning large layers.
        shape_interpretation: Automatic shape interpretation for tensor dims.
        graft_type: Layer wise scale reference from `LayerWiseGrafting`.
        pre_conditioner_type: Dimensions to precondition, from `PreConditionerType`.
        nesterov: Use Nesterov momentum.
        diagonal_eps: Epsilon for numerical stability in diagonal.
        matrix_eps: Epsilon for numerical stability in matrix.
        use_svd: Whether to use SVD for matrix inverse powers (alternative is Schur-Newton).
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        moving_average_for_momentum: bool = False,
        weight_decay: float = 0.0,
        decoupled_weight_decay: bool = False,
        decoupled_learning_rate: bool = True,
        inverse_exponent_override: int = 0,
        start_preconditioning_step: int = 25,
        preconditioning_compute_steps: int = 1000,
        statistics_compute_steps: int = 1,
        block_size: int = 512,
        skip_preconditioning_rank_lt: int = 1,
        no_preconditioning_for_layers_with_dim_gt: int = 8192,
        shape_interpretation: bool = True,
        graft_type: int = LayerWiseGrafting.SGD,
        pre_conditioner_type: int = PreConditionerType.ALL,
        nesterov: bool = True,
        diagonal_eps: float = 1e-10,
        matrix_eps: float = 1e-6,
        use_svd: bool = False,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_step(start_preconditioning_step, 'start_preconditioning_step')
        self.validate_step(preconditioning_compute_steps, 'preconditioning_compute_steps')
        self.validate_step(statistics_compute_steps, 'statistics_compute_steps')
        self.validate_non_negative(diagonal_eps, 'diagonal_eps')
        self.validate_non_negative(matrix_eps, 'matrix_eps')

        self.inverse_exponent_override = inverse_exponent_override
        self.start_preconditioning_step = start_preconditioning_step
        self.preconditioning_compute_steps = preconditioning_compute_steps
        self.statistics_compute_steps = statistics_compute_steps
        self.block_size = block_size
        self.skip_preconditioning_rank_lt = skip_preconditioning_rank_lt
        self.no_preconditioning_for_layers_with_dim_gt = no_preconditioning_for_layers_with_dim_gt
        self.shape_interpretation = shape_interpretation
        self.graft_type = graft_type
        self.pre_conditioner_type = pre_conditioner_type
        self.diagonal_eps = diagonal_eps
        self.matrix_eps = matrix_eps
        self.use_svd = use_svd
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'decoupled_weight_decay': decoupled_weight_decay,
            'decoupled_learning_rate': decoupled_learning_rate,
            'moving_average_for_momentum': moving_average_for_momentum,
            'nesterov': nesterov,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'ScalableShampoo'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        _, beta2 = group['betas']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['momentum'] = torch.zeros_like(grad)
                state['pre_conditioner'] = PreConditioner(
                    p,
                    beta2,
                    self.inverse_exponent_override,
                    self.block_size,
                    self.skip_preconditioning_rank_lt,
                    self.no_preconditioning_for_layers_with_dim_gt,
                    self.shape_interpretation,
                    self.pre_conditioner_type,
                    self.matrix_eps,
                    self.use_svd,
                )
                state['graft'] = build_graft(p, self.graft_type, self.diagonal_eps)

    def is_precondition_step(self, step: int) -> bool:
        return step >= self.start_preconditioning_step

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            is_precondition_step: bool = self.is_precondition_step(group['step'])
            pre_conditioner_multiplier: float = 1.0 if group['decoupled_learning_rate'] else group['lr']
            momentum_multiplier: float = group['lr'] if group['decoupled_learning_rate'] else 1.0
            w: float = (1.0 - beta1) if group['moving_average_for_momentum'] else 1.0

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                pre_conditioner, graft = state['pre_conditioner'], state['graft']

                graft.add_statistics(grad, beta2)
                if group['step'] % self.statistics_compute_steps == 0:
                    pre_conditioner.add_statistics(grad)
                if group['step'] % self.preconditioning_compute_steps == 0:
                    pre_conditioner.compute_pre_conditioners()

                graft_grad: torch.Tensor = graft.precondition_gradient(grad).mul(pre_conditioner_multiplier)
                shampoo_grad: torch.Tensor = pre_conditioner.preconditioned_grad(grad)
                shampoo_grad = (
                    shampoo_grad.mul(pre_conditioner_multiplier)
                    if len(pre_conditioner.pre_conditioners) > 0
                    else graft_grad.clone()
                )

                if self.graft_type != LayerWiseGrafting.NONE:
                    graft_norm = torch.linalg.norm(graft_grad)
                    shampoo_norm = torch.linalg.norm(shampoo_grad)

                    shampoo_grad.mul_(graft_norm / (shampoo_norm + 1e-16))

                if group['decoupled_weight_decay']:
                    self.apply_weight_decay(
                        p,
                        grad=None,
                        lr=group['lr'],
                        weight_decay=group['weight_decay'],
                        weight_decouple=True,
                        fixed_decay=False,
                    )
                else:
                    for g in (graft_grad, shampoo_grad):
                        self.apply_weight_decay(
                            p,
                            grad=g,
                            lr=group['lr'],
                            weight_decay=group['weight_decay'],
                            weight_decouple=False,
                            fixed_decay=False,
                        )

                if group['moving_average_for_momentum']:
                    state['momentum'].lerp_(shampoo_grad, weight=w)
                else:
                    state['momentum'].mul_(beta1).add_(shampoo_grad)

                graft_momentum = graft.update_momentum(graft_grad, beta1, w)

                momentum_update = state['momentum'] if is_precondition_step else graft_momentum

                if group['nesterov']:
                    wd_update = shampoo_grad if is_precondition_step else graft_grad
                    if group['moving_average_for_momentum']:
                        momentum_update = momentum_update.lerp(wd_update, weight=w)
                    else:
                        momentum_update = momentum_update.mul(beta1).add_(wd_update)

                p.add_(momentum_update, alpha=-momentum_multiplier)

        return loss

ScheduleFreeAdamW

Bases: BaseOptimizer

Schedule free AdamW with weighted parameter averaging.

Call train() before training and eval() before evaluating averaged weights.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.0025
betas Betas

Interpolation coefficient for the training iterate and decay rate for squared gradients.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
r float

Exponent of the polynomial step weighting in parameter averaging.

0.0
weight_lr_power float

Exponent of the maximum learning rate seen so far in parameter averaging. 0 disables rate weighting.

2.0
warmup_steps int

Number of linear learning rate warmup steps.

0
decoupling_c int

Coefficient scaling the parameter averaging weight. 0 uses standard averaging.

0
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
eps float

Term added to denominator for numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/schedulefree.py
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
class ScheduleFreeAdamW(BaseOptimizer):
    """Schedule free AdamW with weighted parameter averaging.

    Call `train()` before training and `eval()` before evaluating averaged weights.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Interpolation coefficient for the training iterate and decay rate for squared gradients.
        weight_decay: Weight decay coefficient.
        r: Exponent of the polynomial step weighting in parameter averaging.
        weight_lr_power: Exponent of the maximum learning rate seen so far in parameter averaging. `0` disables
            rate weighting.
        warmup_steps: Number of linear learning rate warmup steps.
        decoupling_c: Coefficient scaling the parameter averaging weight. `0` uses standard averaging.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        eps: Term added to denominator for numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 2.5e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        r: float = 0.0,
        weight_lr_power: float = 2.0,
        warmup_steps: int = 0,
        decoupling_c: int = 0,
        ams_bound: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_range(betas[0], 'beta1', 0.0, 1.0, range_type='()')
        self.validate_non_negative(decoupling_c, 'decoupling_c')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'r': r,
            'weight_lr_power': weight_lr_power,
            'warmup_steps': warmup_steps,
            'decoupling_c': decoupling_c,
            'ams_bound': ams_bound,
            'eps': eps,
            'train_mode': True,
            'weight_sum': 0.0,
            'lr_max': -1.0,
            'use_palm': kwargs.get('use_palm', False),
        }

        super().__init__(params, defaults)

        self.base_lrs: list[float] = [group['lr'] for group in self.param_groups]

    def __str__(self) -> str:
        return 'ScheduleFreeAdamW'

    def eval(self):
        """Switch to averaged weights for evaluation."""
        for group in self.param_groups:
            beta1, _ = group['betas']
            if group['train_mode']:
                for p in group['params']:
                    state = self.state[p]
                    if 'z' in state:
                        p.data.lerp_(end=state['z'], weight=1.0 - 1.0 / beta1)
                group['train_mode'] = False

    def train(self):
        """Switch to training weights before forward and backward passes."""
        for group in self.param_groups:
            beta1, _ = group['betas']
            if not group['train_mode']:
                for p in group['params']:
                    state = self.state[p]
                    if 'z' in state:
                        p.data.lerp_(end=state['z'], weight=1.0 - beta1)
                group['train_mode'] = True

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['z'] = p.clone()
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            warmup_steps: int = group['warmup_steps']
            schedule: float = group['step'] / warmup_steps if group['step'] < warmup_steps else 1.0

            beta1, beta2 = group['betas']

            bias_correction2: float = self.debias(beta2, group['step'])

            lr: float = group['lr'] * schedule
            lr_max = group['lr_max'] = max(lr, group['lr_max'])

            weight: float = (group['step'] ** group['r']) * (lr_max ** group['weight_lr_power'])
            weight_sum = group['weight_sum'] = group['weight_sum'] + weight

            checkpoint: float = weight / weight_sum if weight_sum != 0.0 else 0.0

            if group['decoupling_c'] > 0:
                checkpoint = min(1.0, checkpoint * (1.0 - beta1) * group['decoupling_c'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                z, exp_avg_sq = state['z'], state['exp_avg_sq']

                p, grad, z, exp_avg_sq = self.view_as_real(p, grad, z, exp_avg_sq)

                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                de_nom = self.apply_ams_bound(
                    ams_bound=group['ams_bound'],
                    exp_avg_sq=exp_avg_sq.div(bias_correction2),
                    max_exp_avg_sq=state.get('max_exp_avg_sq', None),
                    eps=group['eps'],
                )

                grad.div_(de_nom)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=lr,
                    weight_decay=group['weight_decay'],
                    weight_decouple=False,
                    fixed_decay=False,
                )

                p.lerp_(z, weight=checkpoint)
                p.add_(grad, alpha=lr * (beta1 * (1.0 - checkpoint) - 1))

                z.sub_(grad, alpha=lr)

        return loss

eval()

Switch to averaged weights for evaluation.

Source code in pytorch_optimizer/optimizer/schedulefree.py
242
243
244
245
246
247
248
249
250
251
def eval(self):
    """Switch to averaged weights for evaluation."""
    for group in self.param_groups:
        beta1, _ = group['betas']
        if group['train_mode']:
            for p in group['params']:
                state = self.state[p]
                if 'z' in state:
                    p.data.lerp_(end=state['z'], weight=1.0 - 1.0 / beta1)
            group['train_mode'] = False

train()

Switch to training weights before forward and backward passes.

Source code in pytorch_optimizer/optimizer/schedulefree.py
253
254
255
256
257
258
259
260
261
262
def train(self):
    """Switch to training weights before forward and backward passes."""
    for group in self.param_groups:
        beta1, _ = group['betas']
        if not group['train_mode']:
            for p in group['params']:
                state = self.state[p]
                if 'z' in state:
                    p.data.lerp_(end=state['z'], weight=1.0 - beta1)
            group['train_mode'] = True

ScheduleFreeRAdam

Bases: BaseOptimizer

Schedule free RAdam with weighted parameter averaging.

Call train() before training and eval() before evaluating averaged weights.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.0025
betas Betas

Interpolation coefficient for the training iterate and decay rate for squared gradients.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
r float

Exponent of the polynomial step weighting in parameter averaging.

0.0
weight_lr_power float

Exponent of the maximum learning rate seen so far in parameter averaging. 0 disables rate weighting.

2.0
silent_sgd_phase bool

If True, disables updates in the early SGD phase, only updates momentum to stabilize training.

True
eps float

Term added to denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/schedulefree.py
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
class ScheduleFreeRAdam(BaseOptimizer):
    """Schedule free RAdam with weighted parameter averaging.

    Call `train()` before training and `eval()` before evaluating averaged weights.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Interpolation coefficient for the training iterate and decay rate for squared gradients.
        weight_decay: Weight decay coefficient.
        r: Exponent of the polynomial step weighting in parameter averaging.
        weight_lr_power: Exponent of the maximum learning rate seen so far in parameter averaging. `0` disables
            rate weighting.
        silent_sgd_phase: If True, disables updates in the early SGD phase, only updates momentum to stabilize
            training.
        eps: Term added to denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 2.5e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        r: float = 0.0,
        weight_lr_power: float = 2.0,
        silent_sgd_phase: bool = True,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_range(betas[0], 'beta1', 0.0, 1.0, range_type='()')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'silent_sgd_phase': silent_sgd_phase,
            'r': r,
            'weight_lr_power': weight_lr_power,
            'eps': eps,
            'train_mode': True,
            'weight_sum': 0.0,
            'lr_max': -1.0,
            'use_palm': kwargs.get('use_palm', False),
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'ScheduleFreeRAdam'

    def eval(self):
        """Switch to averaged weights for evaluation."""
        for group in self.param_groups:
            beta1, _ = group['betas']
            if group['train_mode']:
                for p in group['params']:
                    state = self.state[p]
                    if 'z' in state:
                        p.data.lerp_(end=state['z'], weight=1.0 - 1.0 / beta1)
                group['train_mode'] = False

    def train(self):
        """Switch to training weights before forward and backward passes."""
        for group in self.param_groups:
            beta1, _ = group['betas']
            if not group['train_mode']:
                for p in group['params']:
                    state = self.state[p]
                    if 'z' in state:
                        p.data.lerp_(end=state['z'], weight=1.0 - beta1)
                group['train_mode'] = True

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['z'] = p.clone()
                state['exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction2: float = self.debias(beta2, group['step'])
            bias_correction2_sq: float = bias_correction2 ** 0.5

            lr, n_sma = self.get_rectify_step_size(
                is_rectify=True,
                step=group['step'],
                lr=group['lr'],
                beta2=beta2,
                n_sma_threshold=4,
                degenerated_to_sgd=False,
            )
            if lr < 0.0:
                lr = float(not group['silent_sgd_phase'])
            elif n_sma > 4.0:
                lr = lr / bias_correction2_sq

            lr_max = group['lr_max'] = max(lr, group['lr_max'])

            weight: float = (group['step'] ** group['r']) * (lr_max ** group['weight_lr_power'])
            weight_sum = group['weight_sum'] = group['weight_sum'] + weight

            checkpoint: float = weight / weight_sum if weight_sum != 0.0 else 0.0

            adaptive_y_lr: float = lr * (beta1 * (1.0 - checkpoint) - 1.0)

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                z, exp_avg_sq = state['z'], state['exp_avg_sq']

                p, grad, z, exp_avg_sq = self.view_as_real(p, grad, z, exp_avg_sq)

                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                if n_sma > 4.0:
                    de_nom = exp_avg_sq.sqrt().div_(bias_correction2_sq).add_(group['eps'])
                    grad.div_(de_nom)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=lr,
                    weight_decay=group['weight_decay'],
                    weight_decouple=False,
                    fixed_decay=False,
                )

                p.lerp_(z, weight=checkpoint)
                p.add_(grad, alpha=adaptive_y_lr)

                z.sub_(grad, alpha=lr)

        return loss

eval()

Switch to averaged weights for evaluation.

Source code in pytorch_optimizer/optimizer/schedulefree.py
413
414
415
416
417
418
419
420
421
422
def eval(self):
    """Switch to averaged weights for evaluation."""
    for group in self.param_groups:
        beta1, _ = group['betas']
        if group['train_mode']:
            for p in group['params']:
                state = self.state[p]
                if 'z' in state:
                    p.data.lerp_(end=state['z'], weight=1.0 - 1.0 / beta1)
            group['train_mode'] = False

train()

Switch to training weights before forward and backward passes.

Source code in pytorch_optimizer/optimizer/schedulefree.py
424
425
426
427
428
429
430
431
432
433
def train(self):
    """Switch to training weights before forward and backward passes."""
    for group in self.param_groups:
        beta1, _ = group['betas']
        if not group['train_mode']:
            for p in group['params']:
                state = self.state[p]
                if 'z' in state:
                    p.data.lerp_(end=state['z'], weight=1.0 - beta1)
            group['train_mode'] = True

ScheduleFreeSGD

Bases: BaseOptimizer

Schedule free SGD with weighted parameter averaging.

Call train() before training and eval() before evaluating averaged weights.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

1.0
momentum float

Interpolation coefficient for the training iterate, strictly between 0 and 1.

0.9
weight_decay float

Weight decay coefficient.

0.0
r float

Exponent of the polynomial step weighting in parameter averaging.

0.0
weight_lr_power float

Exponent of the maximum learning rate seen so far in parameter averaging. 0 disables rate weighting.

2.0
warmup_steps int

Number of linear learning rate warmup steps.

0
eps float

Term added to denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/schedulefree.py
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
class ScheduleFreeSGD(BaseOptimizer):
    """Schedule free SGD with weighted parameter averaging.

    Call `train()` before training and `eval()` before evaluating averaged weights.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Interpolation coefficient for the training iterate, strictly between 0 and 1.
        weight_decay: Weight decay coefficient.
        r: Exponent of the polynomial step weighting in parameter averaging.
        weight_lr_power: Exponent of the maximum learning rate seen so far in parameter averaging. `0` disables
            rate weighting.
        warmup_steps: Number of linear learning rate warmup steps.
        eps: Term added to denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1.0,
        momentum: float = 0.9,
        weight_decay: float = 0.0,
        r: float = 0.0,
        weight_lr_power: float = 2.0,
        warmup_steps: int = 0,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0, range_type='()')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'momentum': momentum,
            'weight_decay': weight_decay,
            'r': r,
            'weight_lr_power': weight_lr_power,
            'warmup_steps': warmup_steps,
            'eps': eps,
            'train_mode': True,
            'weight_sum': 0.0,
            'lr_max': -1.0,
        }

        super().__init__(params, defaults)

        self.base_lrs: list[float] = [group['lr'] for group in self.param_groups]

    def __str__(self) -> str:
        return 'ScheduleFreeSGD'

    def eval(self):
        """Switch to averaged weights for evaluation."""
        for group in self.param_groups:
            momentum = group['momentum']
            if group['train_mode']:
                for p in group['params']:
                    state = self.state[p]
                    if 'z' in state:
                        p.data.lerp_(end=state['z'], weight=1.0 - 1.0 / momentum)
                group['train_mode'] = False

    def train(self):
        """Switch to training weights before forward and backward passes."""
        for group in self.param_groups:
            momentum = group['momentum']
            if not group['train_mode']:
                for p in group['params']:
                    state = self.state[p]
                    if 'z' in state:
                        p.data.lerp_(end=state['z'], weight=1.0 - momentum)
                group['train_mode'] = True

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['z'] = p.clone()

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            warmup_steps: int = group['warmup_steps']
            schedule: float = group['step'] / warmup_steps if group['step'] < warmup_steps else 1.0

            momentum = group['momentum']

            lr: float = group['lr'] * schedule
            lr_max = group['lr_max'] = max(lr, group['lr_max'])

            weight: float = (group['step'] ** group['r']) * (lr_max ** group['weight_lr_power'])
            weight_sum = group['weight_sum'] = group['weight_sum'] + weight

            checkpoint: float = weight / weight_sum if weight_sum != 0.0 else 0.0

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                z = state['z']

                p, grad, z = self.view_as_real(p, grad, z)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=lr,
                    weight_decay=group['weight_decay'],
                    weight_decouple=False,
                    fixed_decay=False,
                )

                p.lerp_(z, weight=checkpoint)
                p.add_(grad, alpha=lr * (momentum * (1.0 - checkpoint) - 1))

                z.sub_(grad, alpha=lr)

        return loss

eval()

Switch to averaged weights for evaluation.

Source code in pytorch_optimizer/optimizer/schedulefree.py
80
81
82
83
84
85
86
87
88
89
def eval(self):
    """Switch to averaged weights for evaluation."""
    for group in self.param_groups:
        momentum = group['momentum']
        if group['train_mode']:
            for p in group['params']:
                state = self.state[p]
                if 'z' in state:
                    p.data.lerp_(end=state['z'], weight=1.0 - 1.0 / momentum)
            group['train_mode'] = False

train()

Switch to training weights before forward and backward passes.

Source code in pytorch_optimizer/optimizer/schedulefree.py
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
def train(self):
    """Switch to training weights before forward and backward passes."""
    for group in self.param_groups:
        momentum = group['momentum']
        if not group['train_mode']:
            for p in group['params']:
                state = self.state[p]
                if 'z' in state:
                    p.data.lerp_(end=state['z'], weight=1.0 - momentum)
            group['train_mode'] = True

ScheduleFreeWrapper

Bases: BaseOptimizer

Wrap an optimizer with schedule free parameter averaging.

Call train() before training and eval() before evaluation or saving evaluation weights. The wrapper supplies momentum, so you can disable the base optimizer's momentum. Base optimizer weight decay acts on the fast iterate z. weight_decay_at_y applies additional decay at the training iterate y using the group's current learning rate.

Parameters:

Name Type Description Default
optimizer OptimizerInstanceOrClass

Base optimizer instance or class to wrap.

required
momentum float

Momentum factor.

0.9
weight_decay float

Weight decay coefficient.

0.0
r float

Exponent of the polynomial step weighting in parameter averaging.

0.0
weight_lr_power float

Exponent of the maximum learning rate seen so far in parameter averaging. 0 disables rate weighting.

2.0
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/schedulefree.py
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
class ScheduleFreeWrapper(BaseOptimizer):
    """Wrap an optimizer with schedule free parameter averaging.

    Call `train()` before training and `eval()` before evaluation or saving evaluation
    weights. The wrapper supplies momentum, so you can disable the base optimizer's momentum.
    Base optimizer weight decay acts on the fast iterate `z`. `weight_decay_at_y` applies
    additional decay at the training iterate `y` using the group's current learning rate.

    Args:
        optimizer: Base optimizer instance or class to wrap.
        momentum: Momentum factor.
        weight_decay: Weight decay coefficient.
        r: Exponent of the polynomial step weighting in parameter averaging.
        weight_lr_power: Exponent of the maximum learning rate seen so far in parameter averaging. `0` disables
            rate weighting.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        optimizer: OptimizerInstanceOrClass,
        momentum: float = 0.9,
        weight_decay: float = 0.0,
        r: float = 0.0,
        weight_lr_power: float = 2.0,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_range(momentum, 'momentum', 0.0, 1.0, '()')
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.momentum = momentum
        self.weight_decay = weight_decay
        self.r = r
        self.weight_lr_power = weight_lr_power
        self.train_mode: bool = False
        self.maximize = maximize

        self.optimizer: Optimizer = self.load_optimizer(optimizer, **kwargs)

        self._optimizer_step_pre_hooks: dict[int, Callable] = OrderedDict()
        self._optimizer_step_post_hooks: dict[int, Callable] = OrderedDict()
        self._patch_step_function()

        self.state: State = defaultdict(dict)
        self.defaults: Defaults = self.optimizer.defaults

    def __str__(self) -> str:
        return 'ScheduleFree'

    @property
    def param_groups(self):
        return self.optimizer.param_groups

    def __getstate__(self):
        return {'state': self.state, 'optimizer': self.optimizer}

    def add_param_group(self, param_group):
        return self.optimizer.add_param_group(param_group)

    def state_dict(self) -> State:
        schedulefree_state: State = {
            (group_index, parameter_index): dict(self.state[p])
            for group_index, group in enumerate(self.param_groups)
            for parameter_index, p in enumerate(group['params'])
            if p in self.state
        }
        return {
            'schedulefree_state': schedulefree_state,
            'base_optimizer': self.optimizer.state_dict(),
            'train_mode': self.train_mode,
        }

    def load_state_dict(self, state: State) -> None:
        """Restore the base optimizer state from a checkpoint."""
        saved_state = state['schedulefree_state']
        restored_state: State = {}
        for group_index, group in enumerate(self.param_groups):
            for parameter_index, p in enumerate(group['params']):
                key = (group_index, parameter_index)
                if key in saved_state:
                    restored_state[p] = dict(saved_state[key])
                elif p in saved_state:
                    restored_state[p] = dict(saved_state[p])

        if len(restored_state) != len(saved_state):
            raise ValueError('schedule-free state does not match the current parameters')

        self.optimizer.load_state_dict(state['base_optimizer'])
        for p, parameter_state in restored_state.items():
            for key, value in parameter_state.items():
                if isinstance(value, torch.Tensor):
                    parameter_state[key] = value.to(device=p.device, dtype=p.dtype)
        self.state = defaultdict(dict, restored_state)
        self.train_mode = state.get('train_mode', self.train_mode)

    def zero_grad(self, set_to_none: bool = True) -> None:
        self.optimizer.zero_grad(set_to_none)

    @torch.no_grad()
    def eval(self):
        """Switch to averaged weights for evaluation."""
        if not self.train_mode:
            return

        for group in self.param_groups:
            for p in group['params']:
                state = self.state[p]
                if 'z' in state:
                    p.lerp_(end=state['z'], weight=1.0 - 1.0 / self.momentum)

        self.train_mode = False

    @torch.no_grad()
    def train(self):
        """Switch to training weights before forward and backward passes."""
        if self.train_mode:
            return

        for group in self.param_groups:
            for p in group['params']:
                state = self.state[p]
                if 'z' in state:
                    p.lerp_(end=state['z'], weight=1.0 - self.momentum)

        self.train_mode = True

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if 'z' not in state:
                state['z'] = p.clone()

    @staticmethod
    def swap(x: torch.Tensor, y: torch.Tensor) -> None:
        x.view(torch.uint8).bitwise_xor_(y.view(torch.uint8))
        y.view(torch.uint8).bitwise_xor_(x.view(torch.uint8))
        x.view(torch.uint8).bitwise_xor_(y.view(torch.uint8))

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        if not self.train_mode:
            raise ValueError('optimizer was not in train mode when step is called. call .train() before training')

        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                z = state['z']

                self.apply_weight_decay(
                    z,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=self.weight_decay,
                    weight_decouple=True,
                    fixed_decay=False,
                )

                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=self.weight_decay,
                    weight_decouple=True,
                    fixed_decay=False,
                    ratio=1.0 - self.momentum,
                )

                p.lerp_(end=z, weight=1.0 - 1.0 / self.momentum)

                self.swap(z, p)

        self.optimizer.step()

        for group in self.param_groups:
            lr: float = group['lr'] * group.get('d', 1.0)
            lr_max = group['lr_max'] = max(lr, group.get('lr_max', 0))

            weight: float = (group['step'] ** self.r) * (lr_max ** self.weight_lr_power)
            weight_sum = group['weight_sum'] = group.get('weight_sum', 0.0) + weight

            checkpoint: float = weight / weight_sum if weight_sum != 0.0 else 0.0

            for p in group['params']:
                if p.grad is None:
                    continue

                state = self.state[p]

                z = state['z']

                self.swap(z, p)

                p.lerp_(end=z, weight=checkpoint)

                p.lerp_(end=state['z'], weight=1.0 - self.momentum)

        return loss

eval()

Switch to averaged weights for evaluation.

Source code in pytorch_optimizer/optimizer/schedulefree.py
628
629
630
631
632
633
634
635
636
637
638
639
640
@torch.no_grad()
def eval(self):
    """Switch to averaged weights for evaluation."""
    if not self.train_mode:
        return

    for group in self.param_groups:
        for p in group['params']:
            state = self.state[p]
            if 'z' in state:
                p.lerp_(end=state['z'], weight=1.0 - 1.0 / self.momentum)

    self.train_mode = False

load_state_dict(state)

Restore the base optimizer state from a checkpoint.

Source code in pytorch_optimizer/optimizer/schedulefree.py
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
def load_state_dict(self, state: State) -> None:
    """Restore the base optimizer state from a checkpoint."""
    saved_state = state['schedulefree_state']
    restored_state: State = {}
    for group_index, group in enumerate(self.param_groups):
        for parameter_index, p in enumerate(group['params']):
            key = (group_index, parameter_index)
            if key in saved_state:
                restored_state[p] = dict(saved_state[key])
            elif p in saved_state:
                restored_state[p] = dict(saved_state[p])

    if len(restored_state) != len(saved_state):
        raise ValueError('schedule-free state does not match the current parameters')

    self.optimizer.load_state_dict(state['base_optimizer'])
    for p, parameter_state in restored_state.items():
        for key, value in parameter_state.items():
            if isinstance(value, torch.Tensor):
                parameter_state[key] = value.to(device=p.device, dtype=p.dtype)
    self.state = defaultdict(dict, restored_state)
    self.train_mode = state.get('train_mode', self.train_mode)

train()

Switch to training weights before forward and backward passes.

Source code in pytorch_optimizer/optimizer/schedulefree.py
642
643
644
645
646
647
648
649
650
651
652
653
654
@torch.no_grad()
def train(self):
    """Switch to training weights before forward and backward passes."""
    if self.train_mode:
        return

    for group in self.param_groups:
        for p in group['params']:
            state = self.state[p]
            if 'z' in state:
                p.lerp_(end=state['z'], weight=1.0 - self.momentum)

    self.train_mode = True

SCION

Bases: BaseOptimizer

Norm constrained updates using linear minimization oracles.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
momentum float

Weight of the new gradient in the momentum average, equal to 1 - usual_momentum.

0.1
constraint bool

Use conditional gradient updates within the selected norm radius.

False
norm_type int

Linear minimization oracle type, as an LMONorm value or matching integer.

AUTO
norm_kwargs dict | None

Options for the normalization class.

None
scale float

Radius or update scale for the selected normalization.

1.0
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
maximize bool

Maximize the objective instead of minimizing it.

False

Examples:

from pytorch_optimizer import SCION
from pytorch_optimizer.optimizer.scion import LMONorm

radius = 50.0
parameter_groups = [{
    'params': model.transformer.h.parameters(),
    'norm_type': LMONorm.SPECTRAL,
    'scale': radius,
}, {
    'params': model.lm_head.parameters(),
    'norm_type': LMONorm.SIGN,
    'scale': radius * 60.0,
}]
optimizer = SCION(parameter_groups)
Source code in pytorch_optimizer/optimizer/scion.py
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
class SCION(BaseOptimizer):
    """Norm constrained updates using linear minimization oracles.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Weight of the new gradient in the momentum average, equal to `1 - usual_momentum`.
        constraint: Use conditional gradient updates within the selected norm radius.
        norm_type: Linear minimization oracle type, as an `LMONorm` value or matching integer.
        norm_kwargs: Options for the normalization class.
        scale: Radius or update scale for the selected normalization.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        maximize: Maximize the objective instead of minimizing it.

    Examples:
        ```python
        from pytorch_optimizer import SCION
        from pytorch_optimizer.optimizer.scion import LMONorm

        radius = 50.0
        parameter_groups = [{
            'params': model.transformer.h.parameters(),
            'norm_type': LMONorm.SPECTRAL,
            'scale': radius,
        }, {
            'params': model.lm_head.parameters(),
            'norm_type': LMONorm.SIGN,
            'scale': radius * 60.0,
        }]
        optimizer = SCION(parameter_groups)
        ```

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        momentum: float = 0.1,
        constraint: bool = False,
        norm_type: int = LMONorm.AUTO,
        norm_kwargs: dict | None = None,
        scale: float = 1.0,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        foreach: bool | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0, '(]')
        self.validate_positive(scale, 'scale')

        self.foreach = foreach
        self.maximize = maximize

        if norm_kwargs is None:
            norm_kwargs = {}

        defaults: Defaults = {
            'lr': lr,
            'momentum': momentum,
            'constraint': constraint,
            'norm_type': norm_type,
            'norm_kwargs': norm_kwargs,
            'scale': scale,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'foreach': foreach,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SCION'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if 'd' not in state:
                state['d'] = torch.zeros_like(grad)

    @torch.no_grad()
    def init(self):
        for group in self.param_groups:
            norm = build_lmo_norm(group['norm_type'], **group['norm_kwargs'])
            for p in group['params']:
                p.copy_(norm.init(p))
                p.mul_(group['scale'])

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        if group.get('foreach') is False:
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        norm: Norm,
        ds: list[torch.Tensor],
    ) -> None:
        if self.maximize:
            torch._foreach_neg_(grads)

        if not group['constraint'] and group['weight_decay'] > 0.0:
            self.apply_weight_decay_foreach(
                params,
                grads=grads,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=False,
            )

        torch._foreach_lerp_(ds, grads, group['momentum'])

        updates = [norm.lmo(d) for d in ds]
        torch._foreach_mul_(updates, group['scale'])

        if group['constraint']:
            torch._foreach_mul_(params, 1.0 - group['lr'])

        torch._foreach_add_(params, updates, alpha=-group['lr'])

    def _step_per_param(self, group: ParamGroup, norm: Norm) -> None:
        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            d = state['d']

            if not group['constraint'] and group['weight_decay'] > 0.0:
                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=False,
                )

            d.lerp_(grad, weight=group['momentum'])

            update = norm.lmo(d).mul_(group['scale'])

            if group['constraint']:
                p.mul_(1.0 - group['lr'])

            p.add_(update, alpha=-group['lr'])

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            norm = build_lmo_norm(group['norm_type'], **group['norm_kwargs'])

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(group, self.state, state_keys=['d'])
                if params:
                    self._step_foreach(group, params, grads, norm, state_dict['d'])
            else:
                self._step_per_param(group, norm)

        return loss

SCIONLight

Bases: BaseOptimizer

Variant of SCION that stores momentum in gradient buffers.

Reuse gradient buffers between backward passes to retain momentum. Each step() scales those buffers by 1 - momentum after updating parameters.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
momentum float

Complement of the factor retained in gradient buffers after each update.

0.1
constraint bool

Use conditional gradient updates within the selected norm radius.

False
norm_type int

Linear minimization oracle type, as an LMONorm value or matching integer.

AUTO
norm_kwargs dict | None

Options for the normalization class.

None
scale float

Radius or update scale for the selected normalization.

1.0
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
maximize bool

Maximize the objective instead of minimizing it.

False

Examples:

from pytorch_optimizer import SCIONLight
from pytorch_optimizer.optimizer.scion import LMONorm

radius = 50.0
parameter_groups = [{
    'params': model.transformer.h.parameters(),
    'norm_type': LMONorm.SPECTRAL,
    'scale': radius,
}, {
    'params': model.lm_head.parameters(),
    'norm_type': LMONorm.SIGN,
    'scale': radius * 60.0,
}]
optimizer = SCIONLight(parameter_groups)
Source code in pytorch_optimizer/optimizer/scion.py
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
class SCIONLight(BaseOptimizer):
    """Variant of SCION that stores momentum in gradient buffers.

    Reuse gradient buffers between backward passes to retain momentum. Each `step()`
    scales those buffers by `1 - momentum` after updating parameters.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Complement of the factor retained in gradient buffers after each update.
        constraint: Use conditional gradient updates within the selected norm radius.
        norm_type: Linear minimization oracle type, as an `LMONorm` value or matching integer.
        norm_kwargs: Options for the normalization class.
        scale: Radius or update scale for the selected normalization.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        maximize: Maximize the objective instead of minimizing it.

    Examples:
        ```python
        from pytorch_optimizer import SCIONLight
        from pytorch_optimizer.optimizer.scion import LMONorm

        radius = 50.0
        parameter_groups = [{
            'params': model.transformer.h.parameters(),
            'norm_type': LMONorm.SPECTRAL,
            'scale': radius,
        }, {
            'params': model.lm_head.parameters(),
            'norm_type': LMONorm.SIGN,
            'scale': radius * 60.0,
        }]
        optimizer = SCIONLight(parameter_groups)
        ```

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        momentum: float = 0.1,
        constraint: bool = False,
        norm_type: int = LMONorm.AUTO,
        norm_kwargs: dict | None = None,
        scale: float = 1.0,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        foreach: bool | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0, '(]')
        self.validate_positive(scale, 'scale')

        self.foreach = foreach
        self.maximize = maximize

        if norm_kwargs is None:
            norm_kwargs = {}

        defaults: Defaults = {
            'lr': lr,
            'momentum': momentum,
            'constraint': constraint,
            'norm_type': norm_type,
            'norm_kwargs': norm_kwargs,
            'scale': scale,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'foreach': foreach,
        }
        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SCIONLight'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

    @torch.no_grad()
    def init(self):
        for group in self.param_groups:
            norm = build_lmo_norm(group['norm_type'], **group['norm_kwargs'])
            for p in group['params']:
                p.copy_(norm.init(p))
                p.mul_(group['scale'])

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        if group.get('foreach') is False:
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        norm: Norm,
    ) -> None:
        momentum = group['momentum']

        if self.maximize:
            torch._foreach_neg_(grads)

        updates = [norm.lmo(grad) for grad in grads]
        torch._foreach_mul_(updates, group['scale'])

        if group['constraint']:
            torch._foreach_mul_(params, 1.0 - group['lr'])

        if not group['constraint'] and group['weight_decay'] > 0.0:
            self.apply_weight_decay_foreach(
                params,
                grads=grads,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=False,
            )

        torch._foreach_add_(params, updates, alpha=-group['lr'])

        if momentum != 1.0:
            torch._foreach_mul_(grads, 1.0 - momentum)

    def _step_per_param(self, group: ParamGroup, norm: Norm) -> None:
        momentum = group['momentum']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            self.maximize_gradient(grad, maximize=self.maximize)

            update = norm.lmo(grad).mul_(group['scale'])

            if group['constraint']:
                p.mul_(1.0 - group['lr'])

            if not group['constraint'] and group['weight_decay'] > 0.0:
                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=False,
                )

            p.add_(update, alpha=-group['lr'])

            if momentum != 1.0:
                grad.mul_(1.0 - momentum)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            norm = build_lmo_norm(group['norm_type'], **group['norm_kwargs'])

            if self._can_use_foreach(group):
                params, grads, _ = self.collect_trainable_params(group, self.state)
                if params:
                    self._step_foreach(group, params, grads, norm)
            else:
                self._step_per_param(group, norm)

        return loss

SGDP

Bases: BaseOptimizer

SGD with projected updates for scale invariant weights.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
momentum float

Momentum factor.

0.0
dampening float

Dampening factor for momentum.

0.0
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
delta float

Threshold that determines whether a set of parameters is scale invariant or not.

0.1
wd_ratio float

Relative weight decay applied on scale invariant parameters compared to that applied on scale-variant parameters.

0.1
nesterov bool

Use Nesterov momentum.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adamp.py
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
class SGDP(BaseOptimizer):
    """SGD with projected updates for scale invariant weights.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Momentum factor.
        dampening: Dampening factor for momentum.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        delta: Threshold that determines whether a set of parameters is scale invariant or not.
        wd_ratio: Relative weight decay applied on scale invariant parameters compared to that applied on
            scale-variant parameters.
        nesterov: Use Nesterov momentum.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        momentum: float = 0.0,
        dampening: float = 0.0,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        delta: float = 0.1,
        wd_ratio: float = 0.1,
        nesterov: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(wd_ratio, 'wd_ratio', 0.0, 1.0)
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'momentum': momentum,
            'dampening': dampening,
            'delta': delta,
            'wd_ratio': wd_ratio,
            'nesterov': nesterov,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SGDP'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]
            if len(state) == 0:
                state['momentum'] = torch.zeros_like(grad)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            momentum = group['momentum']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                buf = state['momentum']

                p, grad, buf = self.view_as_real(p, grad, buf)

                buf.mul_(momentum).add_(grad, alpha=1.0 - group['dampening'])

                d_p = buf.clone()
                if group['nesterov']:
                    d_p = d_p.mul_(momentum).add_(grad)

                wd_ratio: float = 1.0
                if len(p.shape) > 1:
                    d_p, wd_ratio = projection(
                        p,
                        grad,
                        d_p,
                        group['delta'],
                        group['wd_ratio'],
                        group['eps'],
                    )

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                    ratio=wd_ratio / (1.0 - momentum),
                )

                p.add_(d_p, alpha=-group['lr'])

        return loss

SGDSaI

Bases: BaseOptimizer

SGD with learning rate scaling from the initial gradient signal-to-noise ratio.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.01
momentum float

Momentum factor.

0.9
weight_decay float

Weight decay coefficient.

0.01
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
eps float

Term added to denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/sgd.py
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
class SGDSaI(BaseOptimizer):
    """SGD with learning rate scaling from the initial gradient signal-to-noise ratio.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Momentum factor.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        eps: Term added to denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-2,
        momentum: float = 0.9,
        weight_decay: float = 1e-2,
        weight_decouple: bool = True,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'momentum': momentum,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SGDSaI'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if group['momentum'] > 0.0 and 'momentum_buffer' not in state:
                state['momentum_buffer'] = torch.zeros_like(p)

            if 'gsnr' not in state:
                sigma = grad.std().nan_to_num_() if grad.ndim > 1 and grad.size(0) != 1 else 0
                grad_norm = grad.norm()
                state['gsnr'] = grad_norm / (sigma + group['eps']) if sigma != 0.0 else grad_norm

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            momentum: float = group['momentum']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                if momentum > 0.0:
                    buf = state['momentum_buffer']
                    buf.lerp_(grad, weight=1.0 - momentum)
                else:
                    buf = grad

                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=False,
                )

                p.add_(buf, alpha=-group['lr'] * state['gsnr'])

        return loss

SGDW

Bases: BaseOptimizer

SGD with optional decoupled weight decay.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float | Tensor

Learning rate.

0.0001
momentum float

Momentum factor.

0.0
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
dampening float

Dampening factor for momentum.

0.0
nesterov bool

Use Nesterov momentum.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/sgd.py
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
class SGDW(BaseOptimizer):
    """SGD with optional decoupled weight decay.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Momentum factor.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        dampening: Dampening factor for momentum.
        nesterov: Use Nesterov momentum.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        maximize: Maximize the objective instead of minimizing it.

    """

    _supports_compiled_foreach = True

    def __init__(
        self,
        params: ParamsT,
        lr: float | torch.Tensor = 1e-4,
        momentum: float = 0.0,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        dampening: float = 0.0,
        nesterov: bool = False,
        foreach: bool | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0)
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'momentum': momentum,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'dampening': dampening,
            'nesterov': nesterov,
            'foreach': foreach,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SGDW'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        if group.get('foreach') is False:
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _init_momentum_buffers(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor] | tuple[torch.Tensor, ...],
    ) -> tuple[list[torch.Tensor], list[bool]]:
        buffers, first_steps = [], []
        if group['momentum'] > 0.0:
            for p, grad in zip(params, grads):
                state = self.state[p]
                first_step = 'momentum_buffer' not in state
                if first_step:
                    state['momentum_buffer'] = torch.empty_like(grad)
                buffers.append(state['momentum_buffer'])
                first_steps.append(first_step)

        return buffers, first_steps

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor] | tuple[torch.Tensor, ...],
        buffers: list[torch.Tensor],
        first_steps: list[bool],
    ) -> None:
        lr, momentum, dampening = group['lr'], group['momentum'], group['dampening']

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=lr,
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=False,
        )

        if momentum > 0.0:
            new_buffers = [buf for buf, first in zip(buffers, first_steps) if first]
            new_grads = [grad for grad, first in zip(grads, first_steps) if first]
            if new_buffers:
                torch._foreach_copy_(new_buffers, new_grads)

            existing_buffers = [buf for buf, first in zip(buffers, first_steps) if not first]
            existing_grads = [grad for grad, first in zip(grads, first_steps) if not first]
            if existing_buffers:
                torch._foreach_mul_(existing_buffers, momentum)
                torch._foreach_add_(existing_buffers, existing_grads, alpha=1.0 - dampening)

            grads = torch._foreach_add(grads, buffers, alpha=momentum) if group['nesterov'] else buffers

        foreach_add_(params, grads, alpha=-lr)

    def _step_per_param(self, group: ParamGroup) -> None:
        momentum = group['momentum']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            self.apply_weight_decay(
                p,
                grad=grad,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=False,
            )

            if momentum > 0.0:
                state = self.state[p]
                buf = state.get('momentum_buffer')
                if buf is None:
                    state['momentum_buffer'] = buf = grad.clone()
                else:
                    buf.mul_(momentum).add_(grad, alpha=1.0 - group['dampening'])

                grad = grad.add_(buf, alpha=momentum) if group['nesterov'] else buf

            p.add_(grad, alpha=-group['lr'])

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if self._can_use_foreach(group):
                params, grads, _ = self.collect_trainable_params(group, self.state)
                for tensors in group_tensors_by_device_and_dtype(params, grads):
                    buffers, first_steps = self._init_momentum_buffers(group, tensors['params'], tensors['grads'])
                    self._step_foreach(group, tensors['params'], tensors['grads'], buffers, first_steps)
            else:
                self._step_per_param(group)

        return loss

Shampoo

Bases: BaseOptimizer

Stochastic tensor optimization with matrix preconditioning.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
momentum float

Momentum factor.

0.0
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
preconditioning_compute_steps int

How often to compute the preconditioner, tuning memory and compute requirements.

1
matrix_eps float

Term added to denominator to improve numerical stability.

1e-06
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/shampoo.py
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
class Shampoo(BaseOptimizer):
    """Stochastic tensor optimization with matrix preconditioning.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Momentum factor.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        preconditioning_compute_steps: How often to compute the preconditioner, tuning memory and compute
            requirements.
        matrix_eps: Term added to denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        momentum: float = 0.0,
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        preconditioning_compute_steps: int = 1,
        matrix_eps: float = 1e-6,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_step(preconditioning_compute_steps, 'preconditioning_compute_steps')
        self.validate_non_negative(matrix_eps, 'matrix_eps')

        self.preconditioning_compute_steps = preconditioning_compute_steps
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'momentum': momentum,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'matrix_eps': matrix_eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Shampoo'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                if group['momentum'] > 0.0:
                    state['momentum_buffer'] = grad.clone()

                for dim_id, dim in enumerate(grad.size()):
                    state[f'pre_cond_{dim_id}'] = (
                        torch.eye(dim, device=grad.device).to(grad.dtype).mul_(group['matrix_eps'])
                    )
                    state[f'inv_pre_cond_{dim_id}'] = grad.new(dim, dim).zero_()

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            momentum = group['momentum']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                if momentum > 0.0:
                    grad.lerp_(state['momentum_buffer'], weight=momentum)

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                order: int = grad.ndimension()
                original_size: int = grad.size()
                statistics_grad = grad
                for dim_id, dim in enumerate(grad.size()):
                    pre_cond, inv_pre_cond = state[f'pre_cond_{dim_id}'], state[f'inv_pre_cond_{dim_id}']

                    unfolded_grad = statistics_grad.movedim(dim_id, 0).reshape(dim, -1)
                    pre_cond.add_(unfolded_grad @ unfolded_grad.t())

                    grad = grad.transpose_(0, dim_id).contiguous()
                    transposed_size = grad.size()

                    grad = grad.view(dim, -1)
                    grad_t = grad.t()

                    if group['step'] % self.preconditioning_compute_steps == 0:
                        inv_pre_cond.copy_(compute_power_svd(pre_cond, 2 * order))

                    if dim_id == order - 1:
                        grad = grad_t @ inv_pre_cond
                        grad = grad.view(original_size)
                    else:
                        grad = inv_pre_cond @ grad
                        grad = grad.view(transposed_size)

                state['momentum_buffer'] = grad

                p.add_(grad, alpha=-group['lr'])

        return loss

SignSGD

Bases: BaseOptimizer

Sign based SGD with optional momentum.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float | Tensor

Learning rate.

0.001
momentum float

Momentum factor. 0 gives SignSGD. Positive values give Signum.

0.9
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/sgd.py
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
class SignSGD(BaseOptimizer):
    """Sign based SGD with optional momentum.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Momentum factor. `0` gives SignSGD. Positive values give Signum.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        maximize: Maximize the objective instead of minimizing it.

    """

    _supports_compiled_foreach = True

    def __init__(
        self,
        params: ParamsT,
        lr: float | torch.Tensor = 1e-3,
        momentum: float = 0.9,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        foreach: bool | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'beta', 0.0, 1.0)
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'momentum': momentum,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'foreach': foreach,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SignSGD'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if group['momentum'] > 0.0 and 'momentum_buffer' not in state:
                state['momentum_buffer'] = torch.zeros_like(p)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        if group.get('foreach') is False:
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        momentum_buffers: list[torch.Tensor],
    ) -> None:
        lr = group['lr']

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=lr,
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=False,
        )

        if group['momentum'] > 0.0:
            torch._foreach_lerp_(momentum_buffers, grads, weight=1.0 - group['momentum'])
            grads = momentum_buffers

        updates = torch._foreach_sign(grads)
        foreach_add_(params, updates, alpha=-lr)

    def _step_per_param(self, group: ParamGroup) -> None:
        momentum = group['momentum']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            self.apply_weight_decay(
                p,
                grad=grad,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=False,
            )

            state = self.state[p]

            if momentum > 0.0:
                buf = state['momentum_buffer']
                buf.lerp_(grad, weight=1.0 - momentum)
            else:
                buf = grad

            p.add_(torch.sign(buf) if not torch.is_complex(buf) else torch.sgn(buf), alpha=-group['lr'])

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if self._can_use_foreach(group):
                state_keys = ['momentum_buffer'] if group['momentum'] > 0.0 else []
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=state_keys
                )
                for tensors in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(group, tensors['params'], tensors['grads'], tensors.get('momentum_buffer', []))
            else:
                self._step_per_param(group)

        return loss

SimplifiedAdEMAMix

Bases: BaseOptimizer

Adaptive updates that mix the current gradient with gradient momentum.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.0001
betas Betas

Decay rates for gradient momentum and squared gradients.

(0.99, 0.95)
alpha float

Coefficient for mixing the current gradient and EMA.

0.0
beta1_warmup int | None

Number of warmup steps used to increase beta1.

None
min_beta1 float

Minimum value of beta1 to start from.

0.9
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
Source code in pytorch_optimizer/optimizer/ademamix.py
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
class SimplifiedAdEMAMix(BaseOptimizer):
    """Adaptive updates that mix the current gradient with gradient momentum.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for gradient momentum and squared gradients.
        alpha: Coefficient for mixing the current gradient and EMA.
        beta1_warmup: Number of warmup steps used to increase beta1.
        min_beta1: Minimum value of beta1 to start from.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-4,
        betas: Betas = (0.99, 0.95),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        alpha: float = 0.0,
        beta1_warmup: int | None = None,
        min_beta1: float = 0.9,
        eps: float = 1e-8,
        maximize: bool = False,
        foreach: bool | None = None,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(alpha, 'alpha')
        self.validate_non_negative(min_beta1, 'min_beta1')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'alpha': alpha,
            'beta1_warmup': beta1_warmup,
            'min_beta1': min_beta1,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
            'foreach': foreach,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SimplifiedAdEMAMix'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['num_sum'] = 0.0
                state['den_sum'] = 0.0

    @staticmethod
    def linear_hl_warmup_scheduler(step: int, beta_end: float, beta_start: float = 0.0, warmup: int = 1) -> float:
        def f(beta: float, eps: float = 1e-8) -> float:
            return math.log(0.5) / math.log(beta + eps) - 1.0

        def f_inv(t: float) -> float:
            return math.pow(0.5, 1.0 / (t + 1))

        if step < warmup:
            a: float = step / float(warmup)
            return f_inv((1.0 - a) * f(beta_start) + a * f(beta_end))

        return beta_end

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        beta1: float,
    ) -> None:
        beta2 = group['betas'][1]

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params, grads, group['lr'], group['weight_decay'], group['weight_decouple'], group['fixed_decay']
        )

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - beta1)
        torch._foreach_mul_(exp_avg_sqs, beta2)
        torch._foreach_addcmul_(exp_avg_sqs, grads, grads, value=1.0 - beta2)

        den_sums: list[float] = []
        for p in params:
            state = self.state[p]
            state['num_sum'] = beta1 * state['num_sum'] + 1.0
            state['den_sum'] = beta2 * state['den_sum'] + (1.0 - beta2)
            den_sums.append(math.sqrt(state['den_sum']))

        de_noms = torch._foreach_sqrt(exp_avg_sqs)
        torch._foreach_add_(de_noms, [den_sum * group['eps'] for den_sum in den_sums])

        updates = torch._foreach_mul(grads, group['alpha'])
        torch._foreach_add_(updates, exp_avgs)
        torch._foreach_div_(updates, de_noms)
        torch._foreach_div_(updates, den_sums)

        torch._foreach_add_(params, updates, alpha=-group['lr'])

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            if group['beta1_warmup']:
                beta1 = self.linear_hl_warmup_scheduler(
                    group['step'], beta_end=beta1, beta_start=group['min_beta1'], warmup=group['beta1_warmup']
                )

            if self.can_use_foreach(group, group.get('foreach')):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_avg_sq']
                )
                for batch in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(
                        group, batch['params'], batch['grads'], batch['exp_avg'], batch['exp_avg_sq'], beta1
                    )
                continue

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                state['num_sum'] = beta1 * state['num_sum'] + 1.0
                state['den_sum'] = beta2 * state['den_sum'] + (1.0 - beta2)

                de_nom = exp_avg_sq.sqrt().add_(math.sqrt(state['den_sum']) * group['eps'])

                update = (group['alpha'] * grad + exp_avg).div_(de_nom).div_(math.sqrt(state['den_sum']))

                p.add_(update, alpha=-group['lr'])

        return loss

SM3

Bases: BaseOptimizer

Adaptive updates with memory efficient per dimension accumulators.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.1
momentum float

Momentum factor.

0.0
beta float

Coefficient used for exponential moving averages.

0.0
eps float

Term added to the denominator to improve numerical stability.

1e-30
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/sm3.py
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
class SM3(BaseOptimizer):
    """Adaptive updates with memory efficient per dimension accumulators.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Momentum factor.
        beta: Coefficient used for exponential moving averages.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-1,
        momentum: float = 0.0,
        beta: float = 0.0,
        eps: float = 1e-30,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0)
        self.validate_range(beta, 'beta', 0.0, 1.0, range_type='[]')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {'lr': lr, 'momentum': momentum, 'beta': beta, 'eps': eps}

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SM3'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            grad = p.grad

            shape = grad.shape
            rank: int = len(shape)

            state = self.state[p]

            if len(state) == 0:
                state['momentum_buffer'] = torch.zeros_like(grad)

                if grad.is_sparse:
                    state['accumulator_0'] = torch.zeros(shape[0], dtype=grad.dtype, device=grad.device)
                elif rank == 0:
                    state['accumulator_0'] = torch.zeros_like(grad)
                else:
                    for i in range(rank):
                        state[f'accumulator_{i}'] = torch.zeros(
                            [1] * i + [shape[i]] + [1] * (rank - 1 - i), dtype=grad.dtype, device=grad.device
                        )

    @staticmethod
    def make_sparse(grad: torch.Tensor, values: torch.Tensor) -> torch.Tensor:
        if grad._indices().dim() == 0 or values.dim() == 0:
            return grad.new().resize_as_(grad)
        return grad.new(grad._indices(), values, grad.size())

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            momentum, beta = group['momentum'], group['beta']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                shape = grad.shape
                rank: int = len(shape)

                state = self.state[p]

                if grad.is_sparse:
                    grad = grad.coalesce()

                    acc = state['accumulator_0']
                    update_values = torch.gather(acc, 0, grad._indices()[0])
                    update_values = update_values.reshape([-1] + [1] * (grad._values().ndim - 1))
                    if update_values.shape != grad._values().shape:
                        update_values = update_values.expand_as(grad._values()).clone()
                    if beta > 0.0:
                        update_values.mul_(beta)
                    update_values.addcmul_(grad._values(), grad._values(), value=1.0 - beta)

                    nu_max = reduce_max_except_dim(self.make_sparse(grad, update_values).to_dense(), 0).squeeze_()

                    if beta > 0.0:
                        torch.max(acc, nu_max, out=acc)
                    else:
                        acc.copy_(nu_max)

                    update_values.add_(group['eps']).rsqrt_().mul_(grad._values())

                    update = self.make_sparse(grad, update_values)
                else:
                    update = state['accumulator_0'].clone()
                    for i in range(1, rank):
                        update = torch.min(update, state[f'accumulator_{i}'])

                    if beta > 0.0:
                        update.mul_(beta)
                    update.addcmul_(grad, grad, value=1.0 - beta)

                    if rank == 0:
                        state['accumulator_0'].copy_(update)

                    for i in range(rank):
                        acc = state[f'accumulator_{i}']
                        nu_max = reduce_max_except_dim(update, i)
                        if beta > 0.0:
                            torch.max(acc, nu_max, out=acc)
                        else:
                            acc.copy_(nu_max)

                    update.add_(group['eps']).rsqrt_().mul_(grad)

                    if momentum > 0.0:
                        m = state['momentum_buffer']
                        m.lerp_(update, weight=1.0 - momentum)
                        update = m

                p.add_(update, alpha=-group['lr'])

        return loss

SOAP

Bases: BaseOptimizer

Adam updates in Shampoo preconditioner eigenbases.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.003
betas Betas

Decay rates for the first and second moments.

(0.95, 0.95)
shampoo_beta float | None

Decay rate for preconditioner statistics. None uses beta2.

None
weight_decay float

Weight decay coefficient.

0.01
precondition_frequency int

Number of steps between eigenbasis updates.

10
max_precondition_dim int

Largest dimension to precondition. Larger dimensions use an identity transform.

10000
merge_dims bool

Whether to merge dimensions of the preconditioner.

False
precondition_1d bool

Whether to precondition 1D gradients.

False
correct_bias bool

Whether to correct bias in Adam.

True
normalize_gradient bool

Whether to normalize the gradients.

False
data_format DataFormat

Tensor layout for dimension merging: 'channels_first' (NCHW) or 'channels_last' (NHWC).

'channels_first'
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/soap.py
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
class SOAP(BaseOptimizer):
    """Adam updates in Shampoo preconditioner eigenbases.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        shampoo_beta: Decay rate for preconditioner statistics. `None` uses `beta2`.
        weight_decay: Weight decay coefficient.
        precondition_frequency: Number of steps between eigenbasis updates.
        max_precondition_dim: Largest dimension to precondition. Larger dimensions use an identity transform.
        merge_dims: Whether to merge dimensions of the preconditioner.
        precondition_1d: Whether to precondition 1D gradients.
        correct_bias: Whether to correct bias in Adam.
        normalize_gradient: Whether to normalize the gradients.
        data_format: Tensor layout for dimension merging: `'channels_first'` (NCHW) or `'channels_last'` (NHWC).
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 3e-3,
        betas: Betas = (0.95, 0.95),
        shampoo_beta: float | None = None,
        weight_decay: float = 1e-2,
        precondition_frequency: int = 10,
        max_precondition_dim: int = 10000,
        merge_dims: bool = False,
        precondition_1d: bool = False,
        correct_bias: bool = True,
        normalize_gradient: bool = False,
        data_format: DataFormat = 'channels_first',
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(shampoo_beta, 'shampoo_beta')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_positive(precondition_frequency, 'precondition_frequency')
        self.validate_positive(max_precondition_dim, 'max_precondition_dim')
        self.validate_non_negative(eps, 'eps')

        self.data_format = data_format
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'shampoo_beta': shampoo_beta,
            'weight_decay': weight_decay,
            'precondition_frequency': precondition_frequency,
            'max_precondition_dim': max_precondition_dim,
            'merge_dims': merge_dims,
            'precondition_1d': precondition_1d,
            'correct_bias': correct_bias,
            'normalize_gradient': normalize_gradient,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SOAP'

    def merge_dims(self, grad: torch.Tensor, max_precondition_dim: int) -> torch.Tensor:
        """Merge dimensions after converting channels-last tensors to channels-first."""
        if self.data_format == 'channels_last' and grad.dim() == 4:
            grad = grad.permute(0, 3, 1, 2)

        return grad.reshape(merge_small_dims(grad.size(), max_precondition_dim))

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        _, beta2 = group['betas']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(grad)
                state['exp_avg_sq'] = torch.zeros_like(grad)

                self.init_pre_conditioner(
                    grad,
                    state,
                    precondition_frequency=group['precondition_frequency'],
                    shampoo_beta=group['shampoo_beta'] if group['shampoo_beta'] is not None else beta2,
                    max_precondition_dim=group['max_precondition_dim'],
                    precondition_1d=group['precondition_1d'],
                    merge_dims=group['merge_dims'],
                )

                self.update_pre_conditioner(
                    grad,
                    state,
                    step=group['step'],
                    max_precondition_dim=group['max_precondition_dim'],
                    precondition_1d=group['precondition_1d'],
                    merge_dims=group['merge_dims'],
                )

    def project(
        self,
        grad: torch.Tensor,
        state,
        merge_dims: bool = False,
        max_precondition_dim: int = 10000,
        project_type: str = 'forward',
    ) -> torch.Tensor:
        original_shape = grad.shape
        permuted_shape = original_shape

        do_permute: bool = self.data_format == 'channels_last' and len(original_shape) == 4

        if merge_dims:
            if do_permute:
                permuted_shape = grad.permute(0, 3, 1, 2).shape

            grad = self.merge_dims(grad, max_precondition_dim)

        for mat in state['Q']:
            if len(mat) > 0:
                grad = torch.tensordot(grad, mat, dims=[[0], [0 if project_type == 'forward' else 1]])
            else:
                grad = grad.permute([*list(range(1, len(grad.shape))), 0])

        if merge_dims:
            grad = grad.reshape(permuted_shape).permute(0, 2, 3, 1) if do_permute else grad.reshape(original_shape)

        return grad

    @staticmethod
    def get_orthogonal_matrix(mat: torch.Tensor) -> list[torch.Tensor]:
        matrices: list = []
        for m in mat:
            if len(m) == 0:
                matrices.append([])
                continue

            try:
                _, q = torch.linalg.eigh(m + 1e-30 * torch.eye(m.shape[0], device=m.device, dtype=m.dtype))
            except Exception:  # pragma: no cover
                _, q = torch.linalg.eigh(
                    m.to(torch.float64) + 1e-30 * torch.eye(m.shape[0], device=m.device, dtype=torch.float64)
                )
                q = q.to(m.dtype)

            q = torch.flip(q, dims=[1])

            matrices.append(q)

        return matrices

    def get_orthogonal_matrix_qr(self, state, max_precondition_dim: int = 10000, merge_dims: bool = False):
        """Compute the eigenbases of the preconditioner using one round of power iteration."""
        original_shape = state['exp_avg_sq'].shape
        permuted_shape = original_shape
        if self.data_format == 'channels_last' and len(original_shape) == 4:
            permuted_shape = state['exp_avg_sq'].permute(0, 3, 1, 2).shape

        exp_avg_sq = state['exp_avg_sq']
        if merge_dims:
            exp_avg_sq = self.merge_dims(exp_avg_sq, max_precondition_dim)

        matrices = []
        for ind, (m, o) in enumerate(zip(state['GG'], state['Q'])):
            if len(m) == 0:
                matrices.append([])
                continue

            est_eig = torch.diag(o.T @ m @ o)
            sort_idx = torch.argsort(est_eig, descending=True)
            exp_avg_sq = exp_avg_sq.index_select(ind, sort_idx)

            power_iter = m @ o[:, sort_idx]

            # Compute QR decomposition
            # We cast to float32 because:
            #  - torch.linalg.qr does not have support for types like bfloat16 as of PyTorch 2.5.1
            #  - the correctness / numerical stability of the Q orthogonality is important for the stability
            #    of the optimizer
            q, _ = torch.linalg.qr(power_iter.to(torch.float32))
            q = q.to(power_iter.dtype)

            matrices.append(q)

        if merge_dims:
            if self.data_format == 'channels_last' and len(original_shape) == 4:
                exp_avg_sq = exp_avg_sq.reshape(permuted_shape).permute(0, 2, 3, 1)
            else:
                exp_avg_sq = exp_avg_sq.reshape(original_shape)

        state['exp_avg_sq'] = exp_avg_sq

        return matrices

    def init_pre_conditioner(
        self,
        grad,
        state,
        precondition_frequency: int = 10,
        shampoo_beta: float = 0.95,
        max_precondition_dim: int = 10000,
        precondition_1d: bool = False,
        merge_dims: bool = False,
    ) -> None:
        state['GG'] = []
        if grad.dim() == 1:
            if not precondition_1d or grad.shape[0] > max_precondition_dim:
                state['GG'].append([])
            else:
                state['GG'].append(torch.zeros(grad.shape[0], grad.shape[0], device=grad.device, dtype=grad.dtype))
        else:
            if merge_dims:
                grad = self.merge_dims(grad, max_precondition_dim)

            for sh in grad.shape:
                if sh > max_precondition_dim:
                    state['GG'].append([])
                else:
                    state['GG'].append(torch.zeros(sh, sh, device=grad.device, dtype=grad.dtype))

        state['Q'] = None
        state['precondition_frequency'] = precondition_frequency
        state['shampoo_beta'] = shampoo_beta

    def update_pre_conditioner(
        self,
        grad,
        state,
        step: int,
        max_precondition_dim: int = 10000,
        precondition_1d: bool = False,
        merge_dims: bool = False,
    ) -> None:
        if grad.dim() == 1:
            if precondition_1d and grad.shape[0] <= max_precondition_dim:
                state['GG'][0].lerp_(
                    (grad.unsqueeze(1) @ grad.unsqueeze(0)).to(state['GG'][0].dtype),
                    weight=1.0 - state['shampoo_beta'],
                )
        else:
            if merge_dims:
                grad = self.merge_dims(grad, max_precondition_dim)

            for idx, dim in enumerate(grad.shape):
                if dim <= max_precondition_dim:
                    outer_product = torch.tensordot(
                        grad,
                        grad,
                        dims=[[*chain(range(idx), range(idx + 1, len(grad.shape)))]] * 2,
                    )

                    state['GG'][idx].lerp_(
                        outer_product.to(state['GG'][idx].dtype), weight=1.0 - state['shampoo_beta']
                    )

        if state['Q'] is None:
            state['Q'] = self.get_orthogonal_matrix(state['GG'])

        if step > 0 and step % state['precondition_frequency'] == 0:
            state['Q'] = self.get_orthogonal_matrix_qr(state, max_precondition_dim, merge_dims)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if group['step'] == 1:
                continue

            beta1, beta2 = group['betas']

            step: int = group['step'] - 1
            step_size: float = group['lr']
            if group['correct_bias']:
                bias_correction1: float = self.debias(beta1, step)
                bias_correction2_sq: float = math.sqrt(self.debias(beta2, step))

                step_size *= bias_correction2_sq / bias_correction1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                grad_projected = self.project(
                    grad, state, merge_dims=group['merge_dims'], max_precondition_dim=group['max_precondition_dim']
                )

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.lerp_(grad_projected.square(), weight=1.0 - beta2)

                de_nom = exp_avg_sq.sqrt().add_(group['eps'])

                exp_avg_projected = self.project(
                    exp_avg, state, merge_dims=group['merge_dims'], max_precondition_dim=group['max_precondition_dim']
                )

                norm_grad = self.project(
                    exp_avg_projected / de_nom,
                    state,
                    merge_dims=group['merge_dims'],
                    max_precondition_dim=group['max_precondition_dim'],
                    project_type='backward',
                )

                if group['normalize_gradient']:
                    norm_grad.div_(torch.mean(norm_grad.square()).sqrt_().add_(group['eps']))

                p.add_(norm_grad, alpha=-step_size)

                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=True,
                    fixed_decay=False,
                )

                self.update_pre_conditioner(
                    grad,
                    state,
                    step=step,
                    max_precondition_dim=group['max_precondition_dim'],
                    merge_dims=group['merge_dims'],
                    precondition_1d=group['precondition_1d'],
                )

        return loss

get_orthogonal_matrix_qr(state, max_precondition_dim=10000, merge_dims=False)

Compute the eigenbases of the preconditioner using one round of power iteration.

Source code in pytorch_optimizer/optimizer/soap.py
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
def get_orthogonal_matrix_qr(self, state, max_precondition_dim: int = 10000, merge_dims: bool = False):
    """Compute the eigenbases of the preconditioner using one round of power iteration."""
    original_shape = state['exp_avg_sq'].shape
    permuted_shape = original_shape
    if self.data_format == 'channels_last' and len(original_shape) == 4:
        permuted_shape = state['exp_avg_sq'].permute(0, 3, 1, 2).shape

    exp_avg_sq = state['exp_avg_sq']
    if merge_dims:
        exp_avg_sq = self.merge_dims(exp_avg_sq, max_precondition_dim)

    matrices = []
    for ind, (m, o) in enumerate(zip(state['GG'], state['Q'])):
        if len(m) == 0:
            matrices.append([])
            continue

        est_eig = torch.diag(o.T @ m @ o)
        sort_idx = torch.argsort(est_eig, descending=True)
        exp_avg_sq = exp_avg_sq.index_select(ind, sort_idx)

        power_iter = m @ o[:, sort_idx]

        # Compute QR decomposition
        # We cast to float32 because:
        #  - torch.linalg.qr does not have support for types like bfloat16 as of PyTorch 2.5.1
        #  - the correctness / numerical stability of the Q orthogonality is important for the stability
        #    of the optimizer
        q, _ = torch.linalg.qr(power_iter.to(torch.float32))
        q = q.to(power_iter.dtype)

        matrices.append(q)

    if merge_dims:
        if self.data_format == 'channels_last' and len(original_shape) == 4:
            exp_avg_sq = exp_avg_sq.reshape(permuted_shape).permute(0, 2, 3, 1)
        else:
            exp_avg_sq = exp_avg_sq.reshape(original_shape)

    state['exp_avg_sq'] = exp_avg_sq

    return matrices

merge_dims(grad, max_precondition_dim)

Merge dimensions after converting channels-last tensors to channels-first.

Source code in pytorch_optimizer/optimizer/soap.py
81
82
83
84
85
86
def merge_dims(self, grad: torch.Tensor, max_precondition_dim: int) -> torch.Tensor:
    """Merge dimensions after converting channels-last tensors to channels-first."""
    if self.data_format == 'channels_last' and grad.dim() == 4:
        grad = grad.permute(0, 3, 1, 2)

    return grad.reshape(merge_small_dims(grad.size(), max_precondition_dim))

SophiaH

Bases: BaseOptimizer

Clipped second-order updates using Hutchinson Hessian estimates.

Use loss.backward(create_graph=True) for internal Hessian estimation, or supply external estimates through step(hessian=...).

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.06
betas Betas

Decay rates for gradient momentum and Hessian diagonal estimates.

(0.96, 0.99)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
p float

Maximum absolute entry of the preconditioned update.

0.01
update_period int

Number of steps after which to apply Hessian approximation.

10
num_samples int

Number of noise samples for each Hessian diagonal estimate.

1
hessian_distribution HutchinsonG

Type of distribution to initialize Hessian.

'gaussian'
eps float

Term added to the denominator to improve numerical stability.

1e-12
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/sophia.py
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
class SophiaH(BaseOptimizer):
    """Clipped second-order updates using Hutchinson Hessian estimates.

    Use `loss.backward(create_graph=True)` for internal Hessian estimation, or supply
    external estimates through `step(hessian=...)`.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for gradient momentum and Hessian diagonal estimates.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        p: Maximum absolute entry of the preconditioned update.
        update_period: Number of steps after which to apply Hessian approximation.
        num_samples: Number of noise samples for each Hessian diagonal estimate.
        hessian_distribution: Type of distribution to initialize Hessian.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 6e-2,
        betas: Betas = (0.96, 0.99),
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        p: float = 1e-2,
        update_period: int = 10,
        num_samples: int = 1,
        hessian_distribution: HutchinsonG = 'gaussian',
        eps: float = 1e-12,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(p, 'p (gradient clip)')
        self.validate_step(update_period, 'update_period')
        self.validate_positive(num_samples, 'num_samples')
        self.validate_options(hessian_distribution, 'hessian_distribution', ['gaussian', 'rademacher'])
        self.validate_non_negative(eps, 'eps')

        self.update_period = update_period
        self.num_samples = num_samples
        self.distribution = hessian_distribution
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'p': p,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SophiaH'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if 'momentum' not in state:
                state['momentum'] = torch.zeros_like(grad)
                state['hessian_moment'] = torch.zeros_like(grad)

    @torch.no_grad()
    def step(self, closure: Closure = None, hessian: list[torch.Tensor] | None = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)

        step: int = self.param_groups[0]['step'] + 1
        update_hessian = hessian is not None or (step - 1) % self.update_period == 0

        if hessian is not None:
            self.set_hessian(self.param_groups, self.state, hessian)
        elif update_hessian:
            self.zero_hessian(self.param_groups, self.state)
            self.compute_hutchinson_hessian(
                param_groups=self.param_groups,
                state=self.state,
                num_samples=self.num_samples,
                distribution=self.distribution,
            )

        for group in self.param_groups:
            group['step'] += 1

            beta1, beta2 = group['betas']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                momentum, hessian_moment = state['momentum'], state['hessian_moment']
                momentum.lerp_(grad, weight=1.0 - beta1)

                if 'hessian' in state and update_hessian:
                    hessian_moment.lerp_(state['hessian'], weight=1.0 - beta2)

                update = (momentum / torch.clip(hessian_moment, min=group['eps'])).clamp_(-group['p'], group['p'])

                p.add_(update, alpha=-group['lr'])

        return loss

SPAM

Bases: BaseOptimizer

Adam with sparse update masks, gradient spike clipping, and momentum resets.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
density float

Expected fraction of 2D parameter entries to update between mask resets.

1.0
weight_decay float

Weight decay coefficient.

0.0
warmup_epoch int

Number of steps to warm up after each momentum reset.

50
threshold int

Squared gradient to second moment ratio above which to clip spikes.

5000
grad_accu_steps int

Steps after a reset before spike clipping begins.

20
update_proj_gap int

Number of steps between mask updates and momentum resets.

500
eps float

Term added to the denominator to improve numerical stability.

1e-06
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/spam.py
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
class SPAM(BaseOptimizer):
    """Adam with sparse update masks, gradient spike clipping, and momentum resets.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        density: Expected fraction of 2D parameter entries to update between mask resets.
        weight_decay: Weight decay coefficient.
        warmup_epoch: Number of steps to warm up after each momentum reset.
        threshold: Squared gradient to second moment ratio above which to clip spikes.
        grad_accu_steps: Steps after a reset before spike clipping begins.
        update_proj_gap: Number of steps between mask updates and momentum resets.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        density: float = 1.0,
        weight_decay: float = 0.0,
        warmup_epoch: int = 50,
        threshold: int = 5000,
        grad_accu_steps: int = 20,
        update_proj_gap: int = 500,
        eps: float = 1e-6,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(warmup_epoch, 'warmup_epoch')
        self.validate_non_negative(density, 'density')
        self.validate_non_negative(threshold, 'threshold')
        self.validate_non_negative(grad_accu_steps, 'grad_accu_steps')
        self.validate_positive(update_proj_gap, 'update_proj_gap')
        self.validate_non_negative(eps, 'eps')

        self.density = density
        self.warmup_epoch = warmup_epoch
        self.threshold = threshold
        self.grad_accu_steps = grad_accu_steps
        self.update_proj_gap = update_proj_gap
        self.maximize = maximize

        defaults: Defaults = {'lr': lr, 'betas': betas, 'weight_decay': weight_decay, 'eps': eps, **kwargs}

        super().__init__(params, defaults)

        self.warmup = CosineDecay(0.99, self.warmup_epoch)

        self.init_masks()

        self.state['total_step'] = 0
        self.state['current_step'] = self.warmup_epoch + 1

    @staticmethod
    def initialize_random_rank_boolean_tensor(m: int, n: int, density: float, device: torch.device) -> torch.Tensor:
        """Create a boolean matrix with an expected fraction of selected entries.

        Args:
            m: Number of rows.
            n: Number of columns.
            density: Probability of selecting each entry. `1` selects all entries.
            device: Device for the matrix.

        Returns:
            torch.Tensor: Boolean mask with shape `(m, n)`.

        """
        total_elements: int = m * n
        non_zero_count: int = int(density * total_elements)

        tensor = torch.zeros(total_elements, dtype=torch.bool, device=device)

        if non_zero_count > 0:
            tensor[torch.randperm(total_elements, device=device)[:non_zero_count]] = True

        return tensor.view(m, n)

    def update_mask_random(self, p: torch.Tensor, old_mask: torch.Tensor) -> torch.Tensor:
        """Resample a parameter mask and retain moments for entries in both masks.

        Args:
            p: Parameter tensor to mask.
            old_mask: Previous boolean mask.

        Returns:
            torch.Tensor: New boolean mask with the configured expected density.

        """
        new_mask: torch.Tensor = torch.rand_like(p) < self.density

        exp_avg = torch.zeros_like(p[new_mask])
        exp_avg_sq = torch.zeros_like(p[new_mask])

        intersection_mask = new_mask & old_mask
        new_intersection_indices = intersection_mask[new_mask]
        old_intersection_indices = intersection_mask[old_mask]

        state = self.state[p]
        exp_avg[new_intersection_indices] = state['exp_avg'][old_intersection_indices]
        exp_avg_sq[new_intersection_indices] = state['exp_avg_sq'][old_intersection_indices]

        state['exp_avg'] = exp_avg
        state['exp_avg_sq'] = exp_avg_sq

        return new_mask

    def update_masks(self) -> None:
        """Resample matrix masks and retain momentum entries shared with the old masks."""
        for group in self.param_groups:
            for p in group['params']:
                state = self.state[p]
                if p.dim() == 2 and 'mask' in state:
                    state['mask'] = self.update_mask_random(p, state['mask'])
                    p.mask = state['mask']

    def init_masks(self) -> None:
        """Initialize sparse update masks for 2D parameters."""
        for group in self.param_groups:
            for p in group['params']:
                state = self.state[p]
                if p.dim() == 2 and 'mask' not in state:
                    state['mask'] = self.initialize_random_rank_boolean_tensor(
                        m=p.shape[0],
                        n=p.shape[1],
                        density=self.density,
                        device=p.device,
                    )

    def __str__(self) -> str:
        return 'SPAM'

    def state_dict(self) -> dict:
        state = super().state_dict()
        state['warmup'] = self.warmup.state_dict()
        return state

    def load_state_dict(self, state_dict: dict) -> None:
        super().load_state_dict(state_dict)
        if 'warmup' in state_dict:
            self.warmup.load_state_dict(state_dict['warmup'])

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        scale_factor: float = 1.0 - self.warmup.get_death_rate(self.state['current_step'])

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size: float = group['lr'] * bias_correction2_sq / bias_correction1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad
                if grad.is_sparse:
                    raise NoSparseGradientError(str(self))

                if torch.is_complex(p):
                    raise NoComplexParameterError(str(self))

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                if 'mask' in state:
                    grad = grad[state['mask']]

                if ('exp_avg' not in state) or (self.state['total_step'] + 1) % self.update_proj_gap == 0:
                    state['exp_avg'] = torch.zeros_like(grad)
                    state['exp_avg_sq'] = torch.zeros_like(grad)

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

                if self.threshold != 0:
                    current_step: int = self.state['total_step'] + 1
                    if current_step >= self.grad_accu_steps and (
                        self.update_proj_gap == 0 or current_step % self.update_proj_gap >= self.grad_accu_steps
                    ):
                        mask = grad.pow(2) > (self.threshold * exp_avg_sq)
                        grad[mask] = grad[mask].sign() * torch.sqrt(exp_avg_sq[mask] * self.threshold)

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                de_nom = exp_avg_sq.sqrt().add_(group['eps'])

                if 'mask' in state:
                    grad_full = torch.zeros_like(p.grad)
                    grad_full[state['mask']] = exp_avg / de_nom
                    p.add_(grad_full, alpha=-step_size * scale_factor)
                else:
                    p.addcdiv_(exp_avg, de_nom, value=-step_size * scale_factor)

                decay_param = p[state['mask']] if 'mask' in state else p
                self.apply_weight_decay(
                    decay_param,
                    grad=None,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=True,
                    fixed_decay=False,
                )
                if 'mask' in state:
                    p[state['mask']] = decay_param

        self.state['total_step'] += 1
        self.state['current_step'] += 1

        if (self.state['total_step'] != 0) and (self.state['total_step'] + 1) % self.update_proj_gap == 0:
            self.update_masks()
            self.state['current_step'] = 0
            self.warmup = CosineDecay(0.99, self.warmup_epoch)

        return loss

init_masks()

Initialize sparse update masks for 2D parameters.

Source code in pytorch_optimizer/optimizer/spam.py
186
187
188
189
190
191
192
193
194
195
196
197
def init_masks(self) -> None:
    """Initialize sparse update masks for 2D parameters."""
    for group in self.param_groups:
        for p in group['params']:
            state = self.state[p]
            if p.dim() == 2 and 'mask' not in state:
                state['mask'] = self.initialize_random_rank_boolean_tensor(
                    m=p.shape[0],
                    n=p.shape[1],
                    density=self.density,
                    device=p.device,
                )

initialize_random_rank_boolean_tensor(m, n, density, device) staticmethod

Create a boolean matrix with an expected fraction of selected entries.

Parameters:

Name Type Description Default
m int

Number of rows.

required
n int

Number of columns.

required
density float

Probability of selecting each entry. 1 selects all entries.

required
device device

Device for the matrix.

required

Returns:

Type Description
Tensor

torch.Tensor: Boolean mask with shape (m, n).

Source code in pytorch_optimizer/optimizer/spam.py
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
@staticmethod
def initialize_random_rank_boolean_tensor(m: int, n: int, density: float, device: torch.device) -> torch.Tensor:
    """Create a boolean matrix with an expected fraction of selected entries.

    Args:
        m: Number of rows.
        n: Number of columns.
        density: Probability of selecting each entry. `1` selects all entries.
        device: Device for the matrix.

    Returns:
        torch.Tensor: Boolean mask with shape `(m, n)`.

    """
    total_elements: int = m * n
    non_zero_count: int = int(density * total_elements)

    tensor = torch.zeros(total_elements, dtype=torch.bool, device=device)

    if non_zero_count > 0:
        tensor[torch.randperm(total_elements, device=device)[:non_zero_count]] = True

    return tensor.view(m, n)

update_mask_random(p, old_mask)

Resample a parameter mask and retain moments for entries in both masks.

Parameters:

Name Type Description Default
p Tensor

Parameter tensor to mask.

required
old_mask Tensor

Previous boolean mask.

required

Returns:

Type Description
Tensor

torch.Tensor: New boolean mask with the configured expected density.

Source code in pytorch_optimizer/optimizer/spam.py
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
def update_mask_random(self, p: torch.Tensor, old_mask: torch.Tensor) -> torch.Tensor:
    """Resample a parameter mask and retain moments for entries in both masks.

    Args:
        p: Parameter tensor to mask.
        old_mask: Previous boolean mask.

    Returns:
        torch.Tensor: New boolean mask with the configured expected density.

    """
    new_mask: torch.Tensor = torch.rand_like(p) < self.density

    exp_avg = torch.zeros_like(p[new_mask])
    exp_avg_sq = torch.zeros_like(p[new_mask])

    intersection_mask = new_mask & old_mask
    new_intersection_indices = intersection_mask[new_mask]
    old_intersection_indices = intersection_mask[old_mask]

    state = self.state[p]
    exp_avg[new_intersection_indices] = state['exp_avg'][old_intersection_indices]
    exp_avg_sq[new_intersection_indices] = state['exp_avg_sq'][old_intersection_indices]

    state['exp_avg'] = exp_avg
    state['exp_avg_sq'] = exp_avg_sq

    return new_mask

update_masks()

Resample matrix masks and retain momentum entries shared with the old masks.

Source code in pytorch_optimizer/optimizer/spam.py
177
178
179
180
181
182
183
184
def update_masks(self) -> None:
    """Resample matrix masks and retain momentum entries shared with the old masks."""
    for group in self.param_groups:
        for p in group['params']:
            state = self.state[p]
            if p.dim() == 2 and 'mask' in state:
                state['mask'] = self.update_mask_random(p, state['mask'])
                p.mask = state['mask']

SpectralSphere

Bases: BaseOptimizer

Matrix updates with a spectral tangent constraint.

Supports dense, real 2D parameters. Uses the leading singular vectors to constrain the update direction, with a Lagrange multiplier from a bisection solver.

References

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.0003
momentum float

Momentum factor.

0.9
weight_decay float

Weight decay coefficient.

0.01
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
nesterov bool

Use Nesterov momentum.

True
power_iteration_steps int

Number of power iteration steps for spectral norm computation.

10
msign_steps int

Number of Newton-Schulz iterations for msign (uses Polar Express).

5
solver_tolerance_f float

Function value tolerance for solver.

1e-08
solver_max_iterations int

Maximum iterations for solver.

100
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/sso.py
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
class SpectralSphere(BaseOptimizer):
    """Matrix updates with a spectral tangent constraint.

    Supports dense, real 2D parameters. Uses the leading singular vectors to constrain
    the update direction, with a Lagrange multiplier from a bisection solver.

    References:
        - Spectral MuP: Spectral Control of Feature Learning.
        - Modular Duality in Deep Learning: https://arxiv.org/abs/2410.21265.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Momentum factor.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        nesterov: Use Nesterov momentum.
        power_iteration_steps: Number of power iteration steps for spectral norm computation.
        msign_steps: Number of Newton-Schulz iterations for msign (uses Polar Express).
        solver_tolerance_f: Function value tolerance for solver.
        solver_max_iterations: Maximum iterations for solver.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 3e-4,
        momentum: float = 0.9,
        weight_decay: float = 1e-2,
        weight_decouple: bool = True,
        nesterov: bool = True,
        power_iteration_steps: int = 10,
        msign_steps: int = 5,
        solver_tolerance_f: float = 1e-8,
        solver_max_iterations: int = 100,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(momentum, 'momentum', 0.0, 1.0, range_type='[)')
        self.validate_positive(power_iteration_steps, 'power_iteration_steps')
        self.validate_positive(msign_steps, 'msign_steps')

        self.power_iteration_steps = power_iteration_steps
        self.msign_steps = msign_steps
        self.solver_tolerance_f = solver_tolerance_f
        self.solver_max_iterations = solver_max_iterations

        self.maximize = maximize

        defaults = {
            'lr': lr,
            'momentum': momentum,
            'nesterov': nesterov,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SpectralSphere'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            if p.dim() != 2:
                raise ValueError(f'{self} only supports 2D parameters')

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if 'momentum_buffer' not in state:
                state['momentum_buffer'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=False,
                )

                buf = state['momentum_buffer']
                buf.lerp_(grad, weight=1.0 - group['momentum'])

                update = grad.lerp_(buf, weight=group['momentum']) if group['nesterov'] else buf

                update = compute_spectral_ball_update(
                    p,
                    momentum=update,
                    power_iteration_steps=self.power_iteration_steps,
                    msign_steps=self.msign_steps,
                    solver_tolerance_f=self.solver_tolerance_f,
                    solver_max_iterations=self.solver_max_iterations,
                )

                p.add_(update, alpha=-group['lr'])

        return loss

SPlus

Bases: BaseOptimizer

Adaptive updates with matrix whitening and gradient momentum.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.1
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.01
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
ema_rate float

Exponential moving average decay rate.

0.999
inverse_steps int

Number of steps between inverse root preconditioner updates.

100
nonstandard_constant float

Scale factor for the learning rate in case of a nonlinear layer.

0.001
max_dim int

Largest tensor dimension to include in matrix preconditioning.

10000
eps float

Term added to the denominator to improve numerical stability.

1e-30
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/splus.py
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
class SPlus(BaseOptimizer):
    """Adaptive updates with matrix whitening and gradient momentum.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        ema_rate: Exponential moving average decay rate.
        inverse_steps: Number of steps between inverse root preconditioner updates.
        nonstandard_constant: Scale factor for the learning rate in case of a nonlinear layer.
        max_dim: Largest tensor dimension to include in matrix preconditioning.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-1,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 1e-2,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        ema_rate: float = 0.999,
        inverse_steps: int = 100,
        nonstandard_constant: float = 1e-3,
        max_dim: int = 10000,
        eps: float = 1e-30,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_range(ema_rate, 'ema_rate', 0.0, 1.0)
        self.validate_positive(inverse_steps, 'inverse_steps')
        self.validate_positive(max_dim, 'max_dim')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'ema_rate': ema_rate,
            'inverse_steps': inverse_steps,
            'max_dim': max_dim,
            'nonstandard_constant': nonstandard_constant,
            'eps': eps,
            'train_mode': True,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SPlus'

    @torch.no_grad()
    def eval(self):
        for group in self.param_groups:
            if group.get('train_mode'):
                for p in group['params']:
                    state = self.state[p]
                    state['param_buffer'] = p.clone()
                    p.lerp_(state['ema'], weight=1.0).mul_(1.0 / (1.0 - group['ema_rate'] ** group['step']))
                group['train_mode'] = False

    @torch.no_grad()
    def train(self):
        for group in self.param_groups:
            if 'train_mode' in group and not group['train_mode']:
                for p in group['params']:
                    state = self.state[p]
                    if 'param_buffer' in state:
                        p.lerp_(state['param_buffer'], weight=1.0)
                        del state['param_buffer']
                group['train_mode'] = True

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['momentum'] = torch.zeros_like(p)
                state['ema'] = torch.zeros_like(p)
                if len(p.shape) == 2:
                    state['sides'] = [
                        torch.zeros((d, d), device=p.device, dtype=p.dtype) if d < group['max_dim'] else None
                        for d in p.shape
                    ]
                    state['q_sides'] = [
                        torch.eye(d, device=p.device).to(p.dtype) if d < group['max_dim'] else None for d in p.shape
                    ]

    @staticmethod
    def get_scaled_lr(shape: tuple[int, int], lr: float, nonstandard_constant: float, max_dim: int = 10000) -> float:
        scale: float = (
            nonstandard_constant
            if len(shape) != 2 or shape[0] > max_dim or shape[1] > max_dim
            else 2.0 / (shape[0] + shape[1])
        )
        return lr * scale

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                scaled_lr: float = self.get_scaled_lr(
                    p.shape, group['lr'], group['nonstandard_constant'], group['max_dim']
                )

                if not group['weight_decouple']:
                    self.apply_weight_decay(
                        p=p,
                        grad=grad,
                        lr=scaled_lr,
                        weight_decay=group['weight_decay'],
                        weight_decouple=False,
                        fixed_decay=group['fixed_decay'],
                    )

                m, ema = state['momentum'], state['ema']
                m.lerp_(grad, weight=1.0 - beta1)

                if len(p.shape) == 2:
                    sides, q_sides = state['sides'], state['q_sides']

                    m = q_sides[0].T @ m if q_sides[0] is not None else m
                    m = m @ q_sides[1] if q_sides[1] is not None else m

                    if sides[0] is not None:
                        torch.lerp(sides[0], grad @ grad.T, weight=1.0 - beta2, out=sides[0])

                    if sides[1] is not None:
                        torch.lerp(sides[1], grad.T @ grad, weight=1.0 - beta2, out=sides[1])

                    update = torch.sign(m)

                    if q_sides[0] is not None:
                        update = q_sides[0] @ update

                    if q_sides[1] is not None:
                        update = update @ q_sides[1].T

                    if group['step'] == 1 or group['step'] % group['inverse_steps'] == 0:
                        if sides[0] is not None:
                            _, eig_vecs = torch.linalg.eigh(
                                sides[0].float() + torch.eye(sides[0].shape[0], device=p.device).mul_(group['eps'])
                            )
                            state['q_sides'][0] = eig_vecs.to(sides[0].dtype)
                        if sides[1] is not None:
                            _, eig_vecs = torch.linalg.eigh(
                                sides[1].float() + torch.eye(sides[1].shape[0], device=p.device).mul_(group['eps'])
                            )
                            state['q_sides'][1] = eig_vecs.to(sides[1].dtype)
                else:
                    update = torch.sign(m)

                p.add_(update, alpha=-scaled_lr)

                ema.lerp_(p, weight=1.0 - group['ema_rate'])

                if group['weight_decouple']:
                    self.apply_weight_decay(
                        p=p,
                        grad=grad,
                        lr=scaled_lr,
                        weight_decay=group['weight_decay'],
                        weight_decouple=True,
                        fixed_decay=group['fixed_decay'],
                    )

        return loss

SRMM

Bases: BaseOptimizer

Stochastic regularized majorization-minimization with weakly convex and multi convex surrogates.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.01
beta float

Adaptivity weight.

0.5
memory_length int | None

Internal memory length for moving average. None for no refreshing.

100
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/srmm.py
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
class SRMM(BaseOptimizer):
    """Stochastic regularized majorization-minimization with weakly convex and multi convex surrogates.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        beta: Adaptivity weight.
        memory_length: Internal memory length for moving average. None for no refreshing.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 0.01,
        beta: float = 0.5,
        memory_length: int | None = 100,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(beta, 'beta', 0.0, 1.0, range_type='[]')

        self.maximize = maximize

        defaults: Defaults = {'lr': lr, 'beta': beta, 'memory_length': memory_length}

        super().__init__(params, defaults)

        self.base_lrs: list[float] = [group['lr'] for group in self.param_groups]

    def __str__(self) -> str:
        return 'SRMM'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['mov_avg_grad'] = torch.zeros_like(grad)
                state['mov_avg_param'] = torch.zeros_like(grad)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            w_t: float = (
                (group['step'] % (group['memory_length'] if group['memory_length'] is not None else 1)) + 1
            ) ** -group['beta']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                mov_avg_grad, mov_avg_param = state['mov_avg_grad'], state['mov_avg_param']

                mov_avg_grad.lerp_(grad, weight=w_t)
                mov_avg_param.lerp_(p, weight=w_t)

                mov_avg_param.add_(mov_avg_grad, alpha=-group['lr'])

                p.copy_(mov_avg_param)

        return loss

StableAdamW

Bases: BaseOptimizer

AdamW with update clipping and optional low precision Kahan summation.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float | Tensor

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.99)
kahan_sum bool

Enables Kahan summation for more accurate parameter updates when training in low precision (float16 or bfloat16). Float16 parameters use float32 second moments to preserve squared gradients.

True
weight_decay float

Weight decay coefficient.

0.01
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
eps float

Term added to the denominator to improve numerical stability.

1e-08
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/adamw.py
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
class StableAdamW(BaseOptimizer):
    """AdamW with update clipping and optional low precision Kahan summation.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        kahan_sum: Enables Kahan summation for more accurate parameter updates when training in low precision
            (float16 or bfloat16). Float16 parameters use float32 second moments to preserve squared gradients.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        eps: Term added to the denominator to improve numerical stability.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float | torch.Tensor = 1e-3,
        betas: Betas = (0.9, 0.99),
        kahan_sum: bool = True,
        weight_decay: float = 1e-2,
        weight_decouple: bool = True,
        eps: float = 1e-8,
        foreach: bool | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.foreach = foreach
        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'kahan_sum': kahan_sum,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'eps': eps,
            'foreach': foreach,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'StableAdamW'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p, dtype=torch.float32 if p.dtype == torch.float16 else p.dtype)

                state['kahan_comp'] = (
                    torch.zeros_like(p)
                    if (group['kahan_sum'] and p.dtype in {torch.float16, torch.bfloat16})
                    else None
                )

    def load_state_dict(self, state_dict: dict) -> None:
        super().load_state_dict(state_dict)

        for group, saved_group in zip(self.param_groups, state_dict['param_groups']):
            for p, saved_id in zip(group['params'], saved_group['params']):
                saved_state = state_dict['state'].get(saved_id, {})
                if p.dtype == torch.float16 and 'exp_avg_sq' in saved_state:
                    self.state[p]['exp_avg_sq'] = saved_state['exp_avg_sq'].to(device=p.device, dtype=torch.float32)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        if group.get('foreach') is False:
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        kahan_comps: list[torch.Tensor],
    ) -> None:
        beta1, beta2 = group['betas']
        eps = group['eps']
        lr = group['lr']

        beta1_comp: float = 1.0 - self.debias_beta(beta1, group['step'])
        beta2_hat: float = self.debias_beta(beta2, group['step'])

        eps_p2: float = math.pow(eps, 2)

        if self.maximize:
            torch._foreach_neg_(grads)

        if not group['weight_decouple']:
            self.apply_weight_decay_foreach(
                params, grads, lr, group['weight_decay'], weight_decouple=False, fixed_decay=False
            )

        torch._foreach_lerp_(exp_avgs, grads, weight=beta1_comp)

        stats_grads = [grad.float() for grad in grads] if params[0].dtype == torch.float16 else grads
        torch._foreach_mul_(exp_avg_sqs, beta2_hat)
        torch._foreach_addcmul_(exp_avg_sqs, stats_grads, stats_grads, value=1.0 - beta2_hat)

        step_sizes: list[torch.Tensor] = [
            -lr / self.get_stable_adamw_rms(grad, exp_avg_sq, eps=eps_p2)
            for grad, exp_avg_sq in zip(stats_grads, exp_avg_sqs)
        ]

        if group['weight_decay'] != 0.0 and group['weight_decouple']:
            wd_step_sizes = [1.0 + group['weight_decay'] * step_size for step_size in step_sizes]
            torch._foreach_mul_(params, wd_step_sizes)

        de_noms = torch._foreach_sqrt(exp_avg_sqs)
        torch._foreach_add_(de_noms, eps)
        de_noms = [de_nom.to(dtype=step_size.dtype) for de_nom, step_size in zip(de_noms, step_sizes)]
        torch._foreach_div_(de_noms, step_sizes)

        if group['kahan_sum'] and params[0].dtype in (torch.float16, torch.bfloat16):
            torch._foreach_addcdiv_(kahan_comps, exp_avgs, de_noms)

            with torch.no_grad():
                torch._foreach_copy_(grads, params)

            torch._foreach_add_(params, kahan_comps)

            torch._foreach_sub_(grads, params)
            torch._foreach_add_(kahan_comps, grads)
        else:
            torch._foreach_addcdiv_(params, exp_avgs, de_noms)

    def _step_per_param(self, group: ParamGroup) -> None:
        beta1, beta2 = group['betas']

        beta1_comp: float = 1.0 - self.debias_beta(beta1, group['step'])
        beta2_hat: float = self.debias_beta(beta2, group['step'])

        eps_p2: float = math.pow(group['eps'], 2)

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            state = self.state[p]

            exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

            p, grad, exp_avg, exp_avg_sq = self.view_as_real(p, grad, exp_avg, exp_avg_sq)

            self.maximize_gradient(grad, maximize=self.maximize)

            if not group['weight_decouple']:
                self.apply_weight_decay(
                    p, grad, group['lr'], group['weight_decay'], weight_decouple=False, fixed_decay=False
                )

            exp_avg.lerp_(grad, weight=beta1_comp)
            stats_grad = grad.float() if grad.dtype == torch.float16 else grad
            exp_avg_sq.mul_(beta2_hat).addcmul_(stats_grad, stats_grad, value=1.0 - beta2_hat)

            lr = group['lr'] / self.get_stable_adamw_rms(stats_grad, exp_avg_sq, eps=eps_p2)

            if group['weight_decouple']:
                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=lr,
                    weight_decay=group['weight_decay'],
                    weight_decouple=True,
                    fixed_decay=False,
                )

            de_nom = exp_avg_sq.sqrt().add_(group['eps']).to(dtype=lr.dtype).div_(-lr)

            if group['kahan_sum'] and p.dtype in (torch.float16, torch.bfloat16):
                kahan_comp = state['kahan_comp']
                kahan_comp.addcdiv_(exp_avg, de_nom)

                grad.copy_(p.detach())
                p.add_(kahan_comp)

                kahan_comp.add_(grad.sub_(p))
            else:
                p.addcdiv_(exp_avg, de_nom)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_avg_sq', 'kahan_comp']
                )
                for batch in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(
                        group,
                        batch['params'],
                        batch['grads'],
                        batch['exp_avg'],
                        batch['exp_avg_sq'],
                        batch['kahan_comp'],
                    )
            else:
                self._step_per_param(group)

        return loss

StableSPAM

Bases: BaseOptimizer

Adam with adaptive gradient scaling, spike clipping, and momentum resets.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
gamma1 float

Decay rate for the gradient norm average. -1 uses beta1.

0.7
gamma2 float

Decay rate for the squared gradient norm average.

0.9
theta float

Decay rate for the maximum absolute gradient average.

0.999
t_max int | None

Steps for cosine decay of momentum coefficients. None disables decay.

None
eta_min float

Minimum multiplier for the cosine decayed momentum coefficients.

0.5
weight_decay float

Weight decay coefficient.

0.0
update_proj_gap int

Steps between momentum resets.

1000
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/spam.py
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
class StableSPAM(BaseOptimizer):
    """Adam with adaptive gradient scaling, spike clipping, and momentum resets.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        gamma1: Decay rate for the gradient norm average. `-1` uses `beta1`.
        gamma2: Decay rate for the squared gradient norm average.
        theta: Decay rate for the maximum absolute gradient average.
        t_max: Steps for cosine decay of momentum coefficients. `None` disables decay.
        eta_min: Minimum multiplier for the cosine decayed momentum coefficients.
        weight_decay: Weight decay coefficient.
        update_proj_gap: Steps between momentum resets.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        gamma1: float = 0.7,
        gamma2: float = 0.9,
        theta: float = 0.999,
        t_max: int | None = None,
        eta_min: float = 0.5,
        weight_decay: float = 0.0,
        update_proj_gap: int = 1000,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_positive(update_proj_gap, 'update_proj_gap')
        self.validate_non_negative(eps, 'eps')

        self.gamma1: float = betas[0] if gamma1 == -1.0 else gamma1
        self.gamma2: float = gamma2
        self.theta: float = theta
        self.t_max = t_max
        self.update_proj_gap = update_proj_gap
        self.warmup = CosineDecay(1.0, t_max, eta_min=eta_min) if t_max is not None else None
        self.maximize = maximize

        self.total_step: int = 0

        defaults: Defaults = {'lr': lr, 'betas': betas, 'weight_decay': weight_decay, 'eps': eps, **kwargs}

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'StableSPAM'

    def state_dict(self) -> dict:
        state = super().state_dict()
        state['total_step'] = self.total_step
        if self.warmup is not None:
            state['warmup'] = self.warmup.state_dict()
        return state

    def load_state_dict(self, state_dict: dict) -> None:
        super().load_state_dict(state_dict)
        self.total_step = state_dict.get('total_step', 0)
        if self.warmup is not None and 'warmup' in state_dict:
            self.warmup.load_state_dict(state_dict['warmup'])

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if 'exp_avg' not in state:
                state['exp_avg'] = torch.zeros_like(grad)
                state['exp_avg_sq'] = torch.zeros_like(grad)
                state['m_norm_t'] = torch.zeros(1, device=grad.device, dtype=grad.dtype)
                state['v_norm_t'] = torch.zeros(1, device=grad.device, dtype=grad.dtype)
                state['m_max_t'] = torch.zeros(1, device=grad.device, dtype=grad.dtype)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        self.total_step += 1

        scale: float = self.warmup.get_death_rate(self.total_step) if self.warmup is not None else 1.0

        for group in self.param_groups:
            self.init_group(group)
            if self.total_step % self.update_proj_gap == 0:
                group['step'] = 1
                for p in group['params']:
                    if 'exp_avg' in self.state[p]:
                        self.state[p]['exp_avg'].zero_()
                        self.state[p]['exp_avg_sq'].zero_()
            else:
                group['step'] += 1

            beta1, beta2 = group['betas']
            beta1 *= scale

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2: float = self.debias(beta2, group['step'])
            bias_correction2_sq: float = math.sqrt(bias_correction2)

            step_size: float = group['lr'] / bias_correction1

            theta_t: float = 1.0 - self.theta ** self.total_step

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=True,
                    fixed_decay=False,
                )

                max_grad = torch.max(grad.abs())

                exp_avg, exp_avg_sq, m_max_t = state['exp_avg'], state['exp_avg_sq'], state['m_max_t']

                m_max_t.lerp_(max_grad, weight=1.0 - self.theta)

                m_max_hat = m_max_t / theta_t

                mask = grad.abs() > m_max_hat
                if mask.sum() > 0:
                    grad[mask] = grad[mask] / max_grad * m_max_hat

                grad_norm = torch.linalg.norm(grad)
                if grad_norm == 0:
                    continue

                m_norm_t, v_norm_t = state['m_norm_t'], state['v_norm_t']
                m_norm_t.lerp_(grad_norm, weight=1.0 - self.gamma1 * scale)
                v_norm_t.lerp_(grad_norm.pow(2), weight=1.0 - self.gamma2)

                m_norm_hat = m_norm_t / (1.0 - (self.gamma1 * scale) ** self.total_step)
                v_norm_hat = v_norm_t / (1.0 - self.gamma2 ** self.total_step)

                c_norm_t = m_norm_hat.div_(v_norm_hat.sqrt_().add_(group['eps']))

                grad.div_(grad_norm).mul_(c_norm_t)

                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                de_nom = exp_avg_sq.sqrt().div_(bias_correction2_sq).add_(group['eps'])

                p.addcdiv_(exp_avg, de_nom, value=-step_size)

        return loss

SWATS

Bases: BaseOptimizer

Adaptive updates that switch from Adam to SGD.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

False
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
ams_bound bool

Use the running maximum of the second moment to bound adaptive updates.

False
nesterov bool

Use Nesterov momentum.

False
eps float

Term added to the denominator to improve numerical stability.

1e-06
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/swats.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
class SWATS(BaseOptimizer):
    """Adaptive updates that switch from Adam to SGD.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        ams_bound: Use the running maximum of the second moment to bound adaptive updates.
        nesterov: Use Nesterov momentum.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        betas: Betas = (0.9, 0.999),
        weight_decay: float = 0.0,
        weight_decouple: bool = False,
        fixed_decay: bool = False,
        ams_bound: bool = False,
        nesterov: bool = False,
        eps: float = 1e-6,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'ams_bound': ams_bound,
            'nesterov': nesterov,
            'phase': 'adam',
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'SWATS'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            if torch.is_complex(p):
                raise NoComplexParameterError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(p)
                state['exp_avg_sq'] = torch.zeros_like(p)
                state['exp_avg2'] = torch.zeros((1,), dtype=p.dtype, device=p.device)

                if group['ams_bound']:
                    state['max_exp_avg_sq'] = torch.zeros_like(p)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2: float = self.debias(beta2, group['step'])

            step_size: float = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False),
                step_size=group['lr'] * math.sqrt(bias_correction2),
                bias_correction1=bias_correction1,
            )

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p=p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=group['fixed_decay'],
                )

                if group['phase'] == 'sgd':
                    if 'momentum_buffer' not in state:
                        state['momentum_buffer'] = torch.zeros_like(grad)

                    buf = state['momentum_buffer']
                    buf.mul_(beta1).add_(grad)

                    update = buf.clone()
                    update.mul_(1.0 - beta1)

                    if group['nesterov']:
                        update.add_(buf, alpha=beta1)

                    p.add_(update, alpha=-group['lr'])

                    continue

                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']
                exp_avg.lerp_(grad, weight=1.0 - beta1)
                exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1.0 - beta2)

                de_nom = self.apply_ams_bound(
                    ams_bound=group['ams_bound'],
                    exp_avg_sq=exp_avg_sq,
                    max_exp_avg_sq=state.get('max_exp_avg_sq', None),
                    eps=group['eps'],
                )

                perturb = exp_avg.clone()
                perturb.div_(de_nom).mul_(-step_size)

                p.add_(perturb)

                perturb_view = perturb.view(-1)
                pg = perturb_view.dot(grad.view(-1))

                if pg != 0:
                    scaling = perturb_view.dot(perturb_view).div_(-pg)

                    exp_avg2 = state['exp_avg2']
                    exp_avg2.lerp_(scaling, weight=1.0 - beta2)

                    corrected_exp_avg = exp_avg2 / bias_correction2

                    if (
                        group['step'] > 1
                        and corrected_exp_avg > 0.0
                        and corrected_exp_avg.allclose(scaling, rtol=group['eps'])
                    ):
                        group['phase'] = 'sgd'
                        group['lr'] = corrected_exp_avg.item()

        return loss

TAM

Bases: BaseOptimizer

SGD with gradient momentum alignment scaling.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.001
momentum float

Momentum factor.

0.9
decay_rate float

Decay rate for the gradient momentum alignment average.

0.9
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/tam.py
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
class TAM(BaseOptimizer):
    """SGD with gradient momentum alignment scaling.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        momentum: Momentum factor.
        decay_rate: Decay rate for the gradient momentum alignment average.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-3,
        momentum: float = 0.9,
        decay_rate: float = 0.9,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(momentum, 'momentum', 0.0, 1.0)
        self.validate_range(decay_rate, 'decay_rate', 0.0, 1.0)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'momentum': momentum,
            'decay_rate': decay_rate,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'TAM'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['s'] = torch.zeros_like(grad)
                state['momentum_buffer'] = grad.clone()
                self.maximize_gradient(state['momentum_buffer'], maximize=self.maximize)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            momentum: float = group['momentum']
            decay_rate: float = group['decay_rate']

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                s, momentum_buffer = state['s'], state['momentum_buffer']

                corr = normalize(momentum_buffer, p=2.0, dim=0).mul_(normalize(grad, p=2.0, dim=0))
                s.lerp_(corr, weight=1.0 - decay_rate)

                d = ((1.0 + s) / 2.0).add_(group['eps']).mul_(grad)

                momentum_buffer.mul_(momentum).add_(d)

                self.apply_weight_decay(
                    p,
                    grad,
                    group['lr'],
                    group['weight_decay'],
                    group['weight_decouple'],
                    group['fixed_decay'],
                )

                p.add_(momentum_buffer, alpha=-group['lr'])

        return loss

Tiger

Bases: BaseOptimizer

Sign based updates with a single gradient momentum buffer.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float | Tensor

Learning rate.

0.001
beta float

Decay rate for gradient momentum.

0.965
weight_decay float

Weight decay coefficient.

0.01
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/tiger.py
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
class Tiger(BaseOptimizer):
    """Sign based updates with a single gradient momentum buffer.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        beta: Decay rate for gradient momentum.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.
        maximize: Maximize the objective instead of minimizing it.

    """

    _supports_compiled_foreach = True

    def __init__(
        self,
        params: ParamsT,
        lr: float | torch.Tensor = 1e-3,
        beta: float = 0.965,
        weight_decay: float = 0.01,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        foreach: bool | None = None,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_range(beta, 'beta', 0.0, 1.0, range_type='[)')
        self.validate_non_negative(weight_decay, 'weight_decay')

        self.maximize = maximize
        self.foreach = foreach

        defaults: Defaults = {
            'lr': lr,
            'beta': beta,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'foreach': foreach,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Tiger'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.zeros_like(grad)

    def _can_use_foreach(self, group: ParamGroup) -> bool:
        if group.get('foreach') is False:
            return False

        return self.can_use_foreach(group, group.get('foreach'))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
    ) -> None:
        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=group['lr'],
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - group['beta'])

        updates = torch._foreach_sign(exp_avgs)

        foreach_add_(params, updates, alpha=-group['lr'])

    def _step_per_param(self, group: ParamGroup) -> None:
        beta = group['beta']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            self.apply_weight_decay(
                p=p,
                grad=grad,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=group['fixed_decay'],
            )

            exp_avg = state['exp_avg']
            exp_avg.lerp_(grad, weight=1.0 - beta)

            p.add_(torch.sign(exp_avg) if not torch.is_complex(exp_avg) else torch.sgn(exp_avg), alpha=-group['lr'])

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            if self._can_use_foreach(group):
                params, grads, state_dict = self.collect_trainable_params(group, self.state, state_keys=['exp_avg'])
                for tensors in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(group, tensors['params'], tensors['grads'], tensors['exp_avg'])
            else:
                self._step_per_param(group)

        return loss

TRAC

Bases: BaseOptimizer

Optimizer wrapper with parameter free scale adaptation.

Parameters:

Name Type Description Default
optimizer OptimizerInstanceOrClass

Base optimizer.

required
betas Sequence[float]

Decay rates for the online learners whose scale estimates are combined.

(0.9, 0.99, 0.999, 0.9999, 0.99999, 0.999999)
num_coefs int

Number of polynomial coefficients to use in the approximation.

128
s_prev float

Initial scale value.

1e-08
eps float

Term added to the denominator to improve numerical stability.

1e-08

Examples:

model = YourModel()
optimizer = TRAC(AdamW(model.parameters()))

for input, output in data:
    optimizer.zero_grad()
    loss = loss_fn(model(input), output)
    loss.backward()
    optimizer.step()
Source code in pytorch_optimizer/optimizer/trac.py
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
class TRAC(BaseOptimizer):
    """Optimizer wrapper with parameter free scale adaptation.

    Args:
        optimizer: Base optimizer.
        betas: Decay rates for the online learners whose scale estimates are combined.
        num_coefs: Number of polynomial coefficients to use in the approximation.
        s_prev: Initial scale value.
        eps: Term added to the denominator to improve numerical stability.

    Examples:
        ```python
        model = YourModel()
        optimizer = TRAC(AdamW(model.parameters()))

        for input, output in data:
            optimizer.zero_grad()
            loss = loss_fn(model(input), output)
            loss.backward()
            optimizer.step()
        ```

    """

    def __init__(
        self,
        optimizer: OptimizerInstanceOrClass,
        betas: Sequence[float] = (0.9, 0.99, 0.999, 0.9999, 0.99999, 0.999999),
        num_coefs: int = 128,
        s_prev: float = 1e-8,
        eps: float = 1e-8,
        **kwargs,
    ):
        self.validate_positive(num_coefs, 'num_coefs')
        self.validate_non_negative(s_prev, 's_prev')
        self.validate_non_negative(eps, 'eps')

        self._optimizer_step_pre_hooks: dict[int, Callable] = OrderedDict()
        self._optimizer_step_post_hooks: dict[int, Callable] = OrderedDict()
        self._patch_step_function()

        self.optimizer: Optimizer = self.load_optimizer(optimizer, **kwargs)

        self.betas = betas
        self.s_prev = s_prev
        self.eps = eps

        self.erf: nn.Module = ERF1994(num_coefs=num_coefs)
        self.f_term: torch.Tensor = self.s_prev / self.erf_imag(1.0 / torch.sqrt(torch.tensor(2.0)))

        self.defaults: Defaults = self.optimizer.defaults

    def __str__(self) -> str:
        return 'TRAC'

    @property
    def param_groups(self):
        return self.optimizer.param_groups

    @property
    def state(self) -> State:
        return self.optimizer.state

    def state_dict(self) -> State:
        state_dict = self.optimizer.state_dict()
        if 'trac' in state_dict['state']:
            parameter_indices = {
                p: index
                for group, saved_group in zip(self.param_groups, state_dict['param_groups'])
                for p, index in zip(group['params'], saved_group['params'])
            }
            state_dict['state']['trac'] = {
                parameter_indices.get(key, key): value for key, value in state_dict['state']['trac'].items()
            }
        return state_dict

    def load_state_dict(self, state_dict: State) -> None:
        saved_trac = state_dict['state'].get('trac')
        if saved_trac is not None:
            parameters = [p for group in self.param_groups for p in group['params']]
            references = [value for key, value in saved_trac.items() if isinstance(key, torch.Tensor)]
            if not references:
                references = [
                    saved_trac.get(index) for group in state_dict['param_groups'] for index in group['params']
                ]

            metadata = {key: value for key, value in saved_trac.items() if isinstance(key, str)}
            if (
                len(references) != len(parameters)
                or len(saved_trac) != len(metadata) + len(parameters)
                or any(
                    not isinstance(ref, torch.Tensor) or ref.shape != p.shape for p, ref in zip(parameters, references)
                )
            ):
                raise ValueError('TRAC state does not match the current parameters')

            trac_state = {
                key: value.to(device=parameters[0].device) if isinstance(value, torch.Tensor) else value
                for key, value in metadata.items()
            }
            trac_state.update({p: ref.to(p) for p, ref in zip(parameters, references)})
            state_dict = {**state_dict, 'state': {**state_dict['state'], 'trac': trac_state}}

        self.optimizer.load_state_dict(state_dict)

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        updates: dict[torch.Tensor, torch.Tensor] = kwargs.get('updates', {})

        for p in group['params']:
            self.state['trac'][p] = updates[p].clone()

    @torch.no_grad()
    def zero_grad(self, set_to_none: bool = True) -> None:
        self.optimizer.zero_grad(set_to_none=set_to_none)

    @torch.no_grad()
    def erf_imag(self, x: torch.Tensor) -> torch.Tensor:
        if not torch.is_floating_point(x):
            x = x.real.to(torch.float32)

        ix = torch.complex(torch.zeros_like(x), x)

        return self.erf(ix).imag

    @torch.no_grad()
    def backup_params_and_grads(self) -> tuple[dict, dict]:
        updates, grads = {}, {}

        for group in self.param_groups:
            for p in group['params']:
                updates[p] = p.clone()
                grads[p] = p.grad.clone() if p.grad is not None else None

        return updates, grads

    @torch.no_grad()
    def trac_step(self, updates: dict, grads: dict) -> None:
        self.state['trac']['step'] += 1

        deltas = {}

        device = self.param_groups[0]['params'][0].device

        s = self.state['trac']['s']
        h = torch.zeros((1,), device=device)
        for group in self.param_groups:
            for p in group['params']:
                if grads[p] is None:
                    continue

                theta_ref = self.state['trac'][p]
                update = updates[p]

                deltas[p] = (update - theta_ref) / s.add(self.eps)
                update.neg_().add_(p)

                grad, delta = grads[p], deltas[p]

                product = torch.dot(delta.flatten(), grad.flatten())
                h.add_(product)

                delta.add_(update)

                p.copy_(theta_ref)

        betas = self.state['trac']['betas']
        variance = self.state['trac']['variance']
        sigma = self.state['trac']['sigma']

        variance.mul_(betas.pow(2)).add_(h.pow(2))
        sigma.mul_(betas).sub_(h)

        term = self.erf_imag(sigma / (2.0 * variance).sqrt_().add_(self.eps)).mul_(self.f_term)
        s.copy_(torch.sum(term))

        scale = max(s, 0.0)

        for group in self.param_groups:
            for p in group['params']:
                if grads[p] is None:
                    continue

                p.add_(deltas[p] * scale)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        # TODO(kozistr): backup is first to get the delta of param and grad, but it does not work.
        with torch.enable_grad():
            loss = self.optimizer.step(closure)

        updates, grads = self.backup_params_and_grads()

        if 'trac' not in self.state:
            device = self.param_groups[0]['params'][0].device

            self.state['trac'] = {
                'betas': torch.tensor(self.betas, device=device),
                's': torch.zeros(1, device=device),
                'variance': torch.zeros(len(self.betas), device=device),
                'sigma': torch.full((len(self.betas),), 1e-8, device=device),
                'step': 0,
            }

            for group in self.param_groups:
                self.init_group(group, updates=updates)

        self.trac_step(updates, grads)

        return loss

VSGD

Bases: BaseOptimizer

SGD with variational gradient estimates.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.1
ghattg float

Prior variance ratio between ghat and g, Var(ghat_t-g_t)/Var(g_t-g_{t-1}).

30.0
ps float

Prior strength.

1e-08
tau1 float

Remember rate for the gamma parameters of g.

0.81
tau2 float

Remember rate for the gamma parameter of ghat.

0.9
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
eps float

Term added to denominator to improve numerical stability.

1e-08
maximize bool

Maximize the objective instead of minimizing it.

False
Source code in pytorch_optimizer/optimizer/sgd.py
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
class VSGD(BaseOptimizer):
    """SGD with variational gradient estimates.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        ghattg: Prior variance ratio between ghat and g, Var(ghat_t-g_t)/Var(g_t-g_{t-1}).
        ps: Prior strength.
        tau1: Remember rate for the gamma parameters of g.
        tau2: Remember rate for the gamma parameter of ghat.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        eps: Term added to denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.

    """

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-1,
        ghattg: float = 30.0,
        ps: float = 1e-8,
        tau1: float = 0.81,
        tau2: float = 0.9,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        eps: float = 1e-8,
        maximize: bool = False,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_non_negative(ghattg, 'ghattg')
        self.validate_non_negative(ps, 'ps')
        self.validate_non_negative(tau1, 'tau1')
        self.validate_non_negative(tau2, 'tau2')
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize

        defaults: Defaults = {
            'lr': lr,
            'tau1': tau1,
            'tau2': tau2,
            'pa2': 2.0 * ps + 1.0 + 1e-4,
            'pbg2': 2.0 * ps,
            'pbhg2': 2.0 * ghattg * ps,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'eps': eps,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'VSGD'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['mug'] = torch.zeros_like(p)
                state['bg'] = torch.full_like(p, group['pbg2'])
                state['bhg'] = torch.full_like(p, group['pbhg2'])

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            pa2, pbg2, pbhg2 = group['pa2'], group['pbg2'], group['pbhg2']

            rho1: float = math.pow(group['step'], -group['tau1'])
            rho2: float = math.pow(group['step'], -group['tau2'])

            for p in group['params']:
                if p.grad is None:
                    continue

                grad = p.grad

                self.maximize_gradient(grad, maximize=self.maximize)

                state = self.state[p]

                self.apply_weight_decay(
                    p,
                    grad=grad,
                    lr=group['lr'],
                    weight_decay=group['weight_decay'],
                    weight_decouple=group['weight_decouple'],
                    fixed_decay=False,
                )

                bg, bhg = state['bg'], state['bhg']

                if group['step'] == 1:
                    sg = pbg2 / (pa2 - 1.0)
                    shg = pbhg2 / (pa2 - 1.0)
                else:
                    sg = bg / pa2
                    shg = bhg / pa2

                mug = state['mug']
                mug_prev = mug.clone()

                mug.mul_(shg).add_(grad * sg).div_(sg + shg)

                sigg = (sg * shg) / (sg + shg)
                mug_sq = mug.pow(2).add_(sigg)

                bg2 = pbg2 + mug_sq - 2.0 * mug * mug_prev + mug_prev.pow(2)
                bhg2 = pbhg2 + mug_sq - 2.0 * grad * mug + grad.pow(2)

                bg.lerp_(bg2, weight=rho1)
                bhg.lerp_(bhg2, weight=rho2)

                p.add_(group['lr'] / mug_sq.sqrt().add_(group['eps']) * mug, alpha=-1.0)

        return loss

WSAM

Bases: BaseOptimizer

Sharpness-aware minimization with weighted sharpness regularization.

Parameters:

Name Type Description Default
model Module | DistributedDataParallel

Model used for training. Supports DistributedDataParallel.

required
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
base_optimizer OptimizerType

Optimizer class to instantiate for parameter updates.

required
rho float

Size of the neighborhood for computing the max loss.

0.05
gamma float

Sharpness mixing coefficient, used as gamma / (1 - gamma).

0.9
adaptive bool

Elementwise adaptive SAM.

False
decouple bool

Apply the sharpness correction after the base optimizer update.

True
max_norm float | None

Max norm of the gradients.

None
eps float

Term added to the denominator of WSAM to improve numerical stability.

1e-12
**kwargs dict

Parameters for optimizer.

{}
Source code in pytorch_optimizer/optimizer/sam.py
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
class WSAM(BaseOptimizer):
    """Sharpness-aware minimization with weighted sharpness regularization.

    Args:
        model: Model used for training. Supports DistributedDataParallel.
        params: Parameters to optimize or dictionaries defining parameter groups.
        base_optimizer: Optimizer class to instantiate for parameter updates.
        rho: Size of the neighborhood for computing the max loss.
        gamma: Sharpness mixing coefficient, used as `gamma / (1 - gamma)`.
        adaptive: Elementwise adaptive SAM.
        decouple: Apply the sharpness correction after the base optimizer update.
        max_norm: Max norm of the gradients.
        eps: Term added to the denominator of WSAM to improve numerical stability.
        **kwargs (dict): Parameters for optimizer.

    """

    def __init__(
        self,
        model: nn.Module | DistributedDataParallel,
        params: ParamsT,
        base_optimizer: OptimizerType,
        rho: float = 0.05,
        gamma: float = 0.9,
        adaptive: bool = False,
        decouple: bool = True,
        max_norm: float | None = None,
        eps: float = 1e-12,
        **kwargs,
    ):
        self.validate_non_negative(rho, 'rho')

        self.model = model
        self.decouple = decouple
        self.max_norm = max_norm

        alpha: float = gamma / (1.0 - gamma)

        defaults: Defaults = {'rho': rho, 'alpha': alpha, 'adaptive': adaptive, 'sam_eps': eps, **kwargs}

        super().__init__(params, defaults)

        self.base_optimizer = base_optimizer(self.param_groups, **kwargs)
        self.param_groups = self.base_optimizer.param_groups

    def __str__(self) -> str:
        return 'WSAM'

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        pass

    @torch.no_grad()
    def first_step(self, zero_grad: bool = False):
        grad_norm = get_global_gradient_norm(self.param_groups, weight_adaptive=True).sqrt_().squeeze_(0)

        for group in self.param_groups:
            scale = group['rho'] / (grad_norm + group['sam_eps'])

            for p in group['params']:
                if p.grad is None:
                    continue

                e_w = (torch.pow(p, 2) if group['adaptive'] else 1.0) * p.grad * scale.to(p)

                p.add_(e_w)

                self.state[p]['e_w'] = e_w

                if is_initialized():  # pragma: no cover
                    all_reduce(p.grad, op=ReduceOp.AVG)

        for group in self.param_groups:
            for p in group['params']:
                self.state[p].pop('grad', None)
                if p.grad is None:
                    continue

                self.state[p]['grad'] = p.grad.clone()

        if zero_grad:
            self.zero_grad()

    @torch.no_grad()
    def second_step(self, zero_grad: bool = False):
        for group in self.param_groups:
            for p in group['params']:
                if 'e_w' in self.state[p]:
                    p.sub_(self.state[p].pop('e_w'))
                if p.grad is None:
                    continue

                if is_initialized():  # pragma: no cover
                    all_reduce(p.grad, ReduceOp.AVG)

        if self.max_norm is not None:
            clip_grad_norm_(self.model.parameters(), self.max_norm)

        for group in self.param_groups:
            for p in group['params']:
                old_grad = self.state[p].pop('grad', None)
                if p.grad is None:
                    continue

                if old_grad is None:
                    old_grad = torch.zeros_like(p.grad)
                if not self.decouple:
                    p.grad.lerp_(old_grad, weight=1.0 - group['alpha'])
                else:
                    self.state[p]['sharpness'] = p.grad.clone() - old_grad
                    p.grad.copy_(old_grad)

        self.base_optimizer.step()

        if self.decouple:
            for group in self.param_groups:
                for p in group['params']:
                    if p.grad is None:
                        continue

                    p.add_(self.state[p]['sharpness'], alpha=-group['lr'] * group['alpha'])

        if zero_grad:
            self.zero_grad()

    @torch.no_grad()
    def step(self, closure: Closure = None):
        if closure is None:
            raise NoClosureError(str(self))

        closure = torch.enable_grad()(closure)

        enable_running_stats(self.model)
        loss = closure()

        self.first_step(zero_grad=True)

        disable_running_stats(self.model)
        closure()

        self.second_step()

        return loss

    def state_dict(self) -> dict:
        state = super().state_dict()
        state['base_optimizer'] = self.base_optimizer.state_dict()
        return state

    def load_state_dict(self, state_dict: dict):
        super().load_state_dict(state_dict)
        if 'base_optimizer' in state_dict:
            self.base_optimizer.load_state_dict(state_dict['base_optimizer'])
            self.param_groups = self.base_optimizer.param_groups
        else:
            self.base_optimizer.param_groups = self.param_groups

Yogi

Bases: BaseOptimizer

Adaptive updates with sign controlled second moment accumulation.

Parameters:

Name Type Description Default
params ParamsT

Parameters to optimize or dictionaries defining parameter groups.

required
lr float

Learning rate.

0.01
betas Betas

Decay rates for the first and second moments.

(0.9, 0.999)
initial_accumulator float

Initial values for first and second moments.

1e-06
weight_decay float

Weight decay coefficient.

0.0
weight_decouple bool

Apply weight decay to parameters instead of adding it to the gradient.

True
fixed_decay bool

Apply decoupled weight decay without scaling it by the learning rate.

False
eps float

Term added to the denominator to improve numerical stability.

0.001
maximize bool

Maximize the objective instead of minimizing it.

False
foreach bool | None

Use batched tensor operations. None enables them for supported parameter groups.

None
Source code in pytorch_optimizer/optimizer/yogi.py
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
class Yogi(BaseOptimizer):
    """Adaptive updates with sign controlled second moment accumulation.

    Args:
        params: Parameters to optimize or dictionaries defining parameter groups.
        lr: Learning rate.
        betas: Decay rates for the first and second moments.
        initial_accumulator: Initial values for first and second moments.
        weight_decay: Weight decay coefficient.
        weight_decouple: Apply weight decay to parameters instead of adding it to the gradient.
        fixed_decay: Apply decoupled weight decay without scaling it by the learning rate.
        eps: Term added to the denominator to improve numerical stability.
        maximize: Maximize the objective instead of minimizing it.
        foreach: Use batched tensor operations. `None` enables them for supported parameter groups.

    """

    _supports_compiled_foreach = True

    def __init__(
        self,
        params: ParamsT,
        lr: float = 1e-2,
        betas: Betas = (0.9, 0.999),
        initial_accumulator: float = 1e-6,
        weight_decay: float = 0.0,
        weight_decouple: bool = True,
        fixed_decay: bool = False,
        eps: float = 1e-3,
        maximize: bool = False,
        foreach: bool | None = None,
        **kwargs,
    ):
        self.validate_learning_rate(lr)
        self.validate_betas(betas)
        self.validate_non_negative(weight_decay, 'weight_decay')
        self.validate_non_negative(eps, 'eps')

        self.maximize = maximize
        self.foreach = foreach
        self._compiled_foreach = False

        defaults: Defaults = {
            'lr': lr,
            'betas': betas,
            'weight_decay': weight_decay,
            'weight_decouple': weight_decouple,
            'fixed_decay': fixed_decay,
            'initial_accumulator': initial_accumulator,
            'eps': eps,
            'foreach': foreach,
            **kwargs,
        }

        super().__init__(params, defaults)

    def __str__(self) -> str:
        return 'Yogi'

    def _compile_foreach(self, compile_kwargs: dict | None = None) -> None:
        # Keep rounded gradient squares for the second-moment sign comparison.
        self._apply_update_foreach = compile_foreach_step(  # ty: ignore[invalid-assignment]
            self._apply_update_foreach, compile_kwargs
        )
        self._compiled_foreach = True

    def init_group(self, group: ParamGroup, **kwargs) -> None:
        if 'step' not in group:
            group['step'] = 0

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad
            if grad.is_sparse:
                raise NoSparseGradientError(str(self))

            state = self.state[p]

            if len(state) == 0:
                state['exp_avg'] = torch.full_like(grad, fill_value=group['initial_accumulator'])
                state['exp_avg_sq'] = torch.full_like(grad, fill_value=group['initial_accumulator'])

    @staticmethod
    def _update_second_moment_foreach(
        exp_avg_sqs: list[torch.Tensor], grad_p2: list[torch.Tensor], beta2: float
    ) -> None:
        signs = torch._foreach_sub(exp_avg_sqs, grad_p2)
        torch._foreach_sign_(signs)

        torch._foreach_addcmul_(exp_avg_sqs, signs, grad_p2, value=-(1.0 - beta2))

    def _step_foreach(
        self,
        group: ParamGroup,
        params: list[torch.Tensor],
        grads: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        step_size: float | torch.Tensor,
        bias_correction2_sq: float | torch.Tensor,
    ) -> None:
        beta1, beta2 = group['betas']

        if self.maximize:
            torch._foreach_neg_(grads)

        self.apply_weight_decay_foreach(
            params=params,
            grads=grads,
            lr=group['lr'],
            weight_decay=group['weight_decay'],
            weight_decouple=group['weight_decouple'],
            fixed_decay=group['fixed_decay'],
        )

        torch._foreach_lerp_(exp_avgs, grads, weight=1.0 - beta1)

        grad_p2 = None
        if self._compiled_foreach:
            grad_p2 = torch._foreach_mul(grads, grads)
        else:
            self._update_second_moment_foreach(exp_avg_sqs, torch._foreach_mul(grads, grads), beta2)

        self._apply_update_foreach(group, params, exp_avgs, exp_avg_sqs, step_size, bias_correction2_sq, grad_p2)

    @staticmethod
    def _apply_update_foreach(
        group: ParamGroup,
        params: list[torch.Tensor],
        exp_avgs: list[torch.Tensor],
        exp_avg_sqs: list[torch.Tensor],
        step_size: float | torch.Tensor,
        bias_correction2_sq: float | torch.Tensor,
        grad_p2: list[torch.Tensor] | None,
    ) -> None:
        if grad_p2 is not None:
            Yogi._update_second_moment_foreach(exp_avg_sqs, grad_p2, group['betas'][1])

        de_noms = torch._foreach_sqrt(exp_avg_sqs)
        torch._foreach_div_(de_noms, bias_correction2_sq)
        torch._foreach_add_(de_noms, group['eps'])

        foreach_addcdiv_(params, exp_avgs, de_noms, value=-step_size)

    def _step_per_param(self, group: ParamGroup, step_size: float, bias_correction2_sq: float) -> None:
        beta1, beta2 = group['betas']

        for p in group['params']:
            if p.grad is None:
                continue

            grad = p.grad

            self.maximize_gradient(grad, maximize=self.maximize)

            state = self.state[p]

            self.apply_weight_decay(
                p=p,
                grad=grad,
                lr=group['lr'],
                weight_decay=group['weight_decay'],
                weight_decouple=group['weight_decouple'],
                fixed_decay=group['fixed_decay'],
            )

            grad_p2 = grad.mul(grad)

            exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']

            exp_avg.lerp_(grad, weight=1.0 - beta1)

            exp_avg_sq.addcmul_(
                (
                    (exp_avg_sq - grad_p2).sign_()
                    if not torch.is_complex(exp_avg_sq)
                    else (exp_avg_sq - grad_p2).sgn_()
                ),
                grad_p2,
                value=-(1.0 - beta2),
            )

            de_nom = exp_avg_sq.sqrt().div_(bias_correction2_sq).add_(group['eps'])

            p.addcdiv_(exp_avg, de_nom, value=-step_size)

    @torch.no_grad()
    def step(self, closure: Closure = None) -> Loss:
        loss: Loss = None
        if closure is not None:
            with torch.enable_grad():
                loss = closure()

        for group in self.param_groups:
            self.init_group(group)
            group['step'] += 1

            beta1, beta2 = group['betas']

            bias_correction1: float = self.debias(beta1, group['step'])
            bias_correction2_sq: float = math.sqrt(self.debias(beta2, group['step']))

            step_size: float = self.apply_adam_debias(
                adam_debias=group.get('adam_debias', False), step_size=group['lr'], bias_correction1=bias_correction1
            )

            if self.can_use_foreach(group, group.get('foreach')):
                params, grads, state_dict = self.collect_trainable_params(
                    group, self.state, state_keys=['exp_avg', 'exp_avg_sq']
                )

                for tensors in group_tensors_by_device_and_dtype(params, grads, state_dict):
                    self._step_foreach(
                        group,
                        tensors['params'],
                        tensors['grads'],
                        tensors['exp_avg'],
                        tensors['exp_avg_sq'],
                        step_size,
                        bias_correction2_sq,
                    )
            else:
                self._step_per_param(group, step_size, bias_correction2_sq)

        return loss