Functional API to apply a StableAdamW optimization step. See `optimi.StableAdamW` for more details. Args: params: Parameters to update grads: Parameter gradients exp_avgs: Gradient moving averages exp_avg_sqs: Squared gradient moving averages eps_sqs
(
params: list[Tensor],
grads: list[Tensor],
exp_avgs: list[Tensor],
exp_avg_sqs: list[Tensor],
eps_sqs: list[Tensor],
kahan_comps: list[Tensor | None] | None = None,
*,
lr: float,
beta1: float,
beta2: float,
weight_decay: float,
eps: float,
step: Tensor,
decouple_lr: bool = False,
max_lr: float | None = None,
kahan_sum: bool = False,
return_norms: bool = True,
)
| 181 | |
| 182 | |
| 183 | def stableadamw( |
| 184 | params: list[Tensor], |
| 185 | grads: list[Tensor], |
| 186 | exp_avgs: list[Tensor], |
| 187 | exp_avg_sqs: list[Tensor], |
| 188 | eps_sqs: list[Tensor], |
| 189 | kahan_comps: list[Tensor | None] | None = None, |
| 190 | *, |
| 191 | lr: float, |
| 192 | beta1: float, |
| 193 | beta2: float, |
| 194 | weight_decay: float, |
| 195 | eps: float, |
| 196 | step: Tensor, |
| 197 | decouple_lr: bool = False, |
| 198 | max_lr: float | None = None, |
| 199 | kahan_sum: bool = False, |
| 200 | return_norms: bool = True, |
| 201 | ): |
| 202 | """Functional API to apply a StableAdamW optimization step. |
| 203 | |
| 204 | See `optimi.StableAdamW` for more details. |
| 205 | |
| 206 | Args: |
| 207 | params: Parameters to update |
| 208 | grads: Parameter gradients |
| 209 | exp_avgs: Gradient moving averages |
| 210 | exp_avg_sqs: Squared gradient moving averages |
| 211 | eps_sqs: Squared epsilon term tensors |
| 212 | kahan_comps: Kahan summation compensations |
| 213 | lr: Learning rate |
| 214 | beta1: Gradient moving average coefficient |
| 215 | beta2: Squared gradient moving average coefficient |
| 216 | weight_decay: Weight decay coefficient |
| 217 | eps: Added to denominator to improve numerical stability |
| 218 | step: Step counter used for bias correction |
| 219 | decouple_lr: Apply fully decoupled weight decay |
| 220 | max_lr: Maximum scheduled learning rate for `decouple_lr` |
| 221 | kahan_sum: Enables Kahan summation for low precision parameters |
| 222 | """ |
| 223 | # calculate debiased beta hat & complement terms |
| 224 | step.add_(1) |
| 225 | beta1_comp = 1 - debias_beta(beta1, step.item()) |
| 226 | beta2_hat = debias_beta(beta2, step.item()) |
| 227 | |
| 228 | if kahan_comps is None: |
| 229 | kahan_comps = [None] * len(params) |
| 230 | |
| 231 | return _foreach_stableadamw( |
| 232 | params, |
| 233 | grads, |
| 234 | exp_avgs, |
| 235 | exp_avg_sqs, |
| 236 | eps_sqs, |
| 237 | kahan_comps, |
| 238 | lr=lr, |
| 239 | beta1_comp=beta1_comp, |
| 240 | beta2_hat=beta2_hat, |
no test coverage detected