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 |
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 |
False
|
compile_kwargs
|
dict | None
|
Options for |
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 | |
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 |
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 | |
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 | |
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 | |
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
|
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 | |
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
|
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 | |
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 | |
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 |
(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 |
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
|
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 | |
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 | |
get_options(shape)
staticmethod
¶
Return whether the gradient supports factored second moments.
Source code in pytorch_optimizer/optimizer/adafactor.py
166 167 168 169 | |
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 | |
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 |
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
|
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 | |
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 | |
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 | |
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 | |
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. |
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 | |
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
|
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 | |
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 | |
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 |
0.2
|
q
|
float
|
The q value in the numerator function |
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 | |
s(p)
¶
Compute the numerator scaling function p * x ** q.
Source code in pytorch_optimizer/optimizer/adamg.py
86 87 88 | |
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 | |
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
|
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 | |
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 | |
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 | |
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 |
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
|
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 | |
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
|
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 | |
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
|
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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 |
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
|
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 | |
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 |
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
|
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 | |
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 |
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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 |
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
|
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 | |
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 | |
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 | |
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'
|
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 | |
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 | |
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 | |
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 | |
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 | |
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
|
mode
|
Mode
|
Search direction and estimator: |
'c'
|
simple_cond
|
bool
|
Use the simplified conditional estimator in mode |
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 | |
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 | |
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 | |
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 | |
get_options(shape)
staticmethod
¶
Return whether the gradient supports factored second moments.
Source code in pytorch_optimizer/optimizer/came.py
110 111 112 113 | |
get_rms(x)
staticmethod
¶
Compute the root mean square of a tensor.
Source code in pytorch_optimizer/optimizer/came.py
115 116 117 118 | |
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 | |
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 | |
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
|
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 | |
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
|
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 | |
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
|
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 | |
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
|
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 | |
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
|
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 | |
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
|
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'
|
matrix_scaling
|
Literal['power', 'frobenius']
|
Newton-DB scaling: |
'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 | |
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 | |
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 | |
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 | |
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 | |
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
|
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 | |
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 | |
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
|
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 | |
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 |
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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 |
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 |
True
|
master_weight_bits
|
int | None
|
Effective master weight precision for bf16/fp16 parameters. Supports |
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 | |
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 | |
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 | |
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 | |
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 | |
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
|
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 | |
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 | |
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 |
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 | |
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 | |
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 | |
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 |
True
|
eps
|
float
|
Term added to the denominator to improve numerical stability. |
1e-08
|
foreach
|
bool | None
|
Use batched tensor operations. |
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 | |
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 | |
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 | |
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 | |
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
|
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
|
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 | |
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 |
True
|
max_grad_norm
|
float
|
Reference norm for gradient scaling when |
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
|
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 | |
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 | |
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
|
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 | |
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
|
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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'
|
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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.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 | |
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 | |
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'
|
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
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 | |
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'
|
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 | |
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 | |
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 | |
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 |
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
|
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 | |
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 | |
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 |
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'
|
use_adjusted_lr
|
bool
|
Apply the Moonlight shape adjustment in |
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
|
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 | |
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 |
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 | |
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 | |
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
|
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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
|
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 | |
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 | |
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 | |
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 | |
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
|
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 | |
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 | |
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 | |
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
|
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 | |
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 |
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 | |
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 | |
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 |
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 | |
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 | |
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 | |
get_lr()
¶
Return the base optimizer learning rate.
Source code in pytorch_optimizer/optimizer/fp16.py
280 281 282 | |
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 | |
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 | |
set_lr(lr)
¶
Set the base optimizer learning rate.
Source code in pytorch_optimizer/optimizer/fp16.py
284 285 286 287 | |
state_dict()
¶
Return the optimizer state dict.
Source code in pytorch_optimizer/optimizer/fp16.py
159 160 161 162 163 164 165 | |
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 | |
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 | |
zero_grad()
¶
Clear the gradients of all optimized parameters.
Source code in pytorch_optimizer/optimizer/fp16.py
272 273 274 275 276 277 278 | |
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 | |
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 | |
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 |
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 |
-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 | |
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 | |
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 | |
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.
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 |
SGD
|
pre_conditioner_type
|
int
|
Dimensions to precondition, from |
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 | |
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. |
2.0
|
warmup_steps
|
int
|
Number of linear learning rate warmup steps. |
0
|
decoupling_c
|
int
|
Coefficient scaling the parameter averaging weight. |
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 | |
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 | |
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 | |
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. |
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 | |
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 | |
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 | |
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. |
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 | |
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 | |
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 | |
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. |
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 | |
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 | |
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 | |
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 | |
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 |
0.1
|
constraint
|
bool
|
Use conditional gradient updates within the selected norm radius. |
False
|
norm_type
|
int
|
Linear minimization oracle type, as an |
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
|
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 | |
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 |
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
|
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 | |
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 | |
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 | |
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
|
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 | |
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 | |
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.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
|
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 | |
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
|
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 | |
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 | |
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
|
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'
|
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 | |
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 | |
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 | |
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 | |
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 | |
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 | |
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. |
required |
device
|
device
|
Device for the matrix. |
required |
Returns:
| Type | Description |
|---|---|
Tensor
|
torch.Tensor: Boolean mask with shape |
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 | |
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 | |
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 | |
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
- Spectral MuP: Spectral Control of Feature Learning.
- Modular Duality in Deep Learning: https://arxiv.org/abs/2410.21265.
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 | |
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 | |
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 | |
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
|
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 | |
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. |
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
|
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 | |
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 | |
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 | |
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
|
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 | |
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 | |
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 | |
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 |
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 | |
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
|
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 | |