MCPcopy Create free account
hub / github.com/AnswerDotAI/ModernBERT / stableadamw

Function stableadamw

src/optimizer.py:183–247  ·  view source on GitHub ↗

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,
)

Source from the content-addressed store, hash-verified

181
182
183def 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,

Callers 1

stepMethod · 0.85

Calls 1

_foreach_stableadamwFunction · 0.85

Tested by

no test coverage detected