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

Method step

src/optimizer.py:139–180  ·  view source on GitHub ↗

Performs a single optimization step on the whole model or individual parameter. Args: closure: A closure which reevaluates the model and returns the loss. Incompatible with performing an optimization step on a single `param`. param: An individual para

(self, closure: Callable | None = None, param: Tensor | None = None)

Source from the content-addressed store, hash-verified

137
138 @torch.no_grad()
139 def step(self, closure: Callable | None = None, param: Tensor | None = None):
140 """Performs a single optimization step on the whole model or individual parameter.
141
142 Args:
143 closure: A closure which reevaluates the model and returns the loss. Incompatible with
144 performing an optimization step on a single `param`.
145 param: An individual parameter to perform a fused optimization step during the backward
146 pass. Requires optimizer to be initialized with `gradient_release=True` and model
147 hooks created with `register_gradient_release`. Incompatible with `closure`.
148 """
149 loss = None
150 if closure is not None and param is None:
151 with torch.enable_grad():
152 loss = closure()
153
154 for group in self.param_groups:
155 params, grads, exp_avgs, exp_avg_sqs, eps_sqs, kahan_comps = [], [], [], [], [], []
156 self._init_group(group, params, grads, exp_avgs, exp_avg_sqs, eps_sqs, kahan_comps)
157
158 l1_norm, l2_norm = stableadamw(
159 params=params,
160 grads=grads,
161 exp_avgs=exp_avgs,
162 exp_avg_sqs=exp_avg_sqs,
163 eps_sqs=eps_sqs,
164 kahan_comps=kahan_comps,
165 lr=group["lr"],
166 beta1=group["beta1"],
167 beta2=group["beta2"],
168 weight_decay=group["weight_decay"],
169 eps=group["eps"],
170 step=group["step"],
171 decouple_lr=group["decouple_lr"],
172 max_lr=group["max_lr"],
173 kahan_sum=group["kahan_sum"],
174 return_norms=self.return_norms,
175 )
176
177 self.grad_norms["l1_norm"] = l1_norm
178 self.grad_norms["l2_norm"] = l2_norm
179
180 return loss
181
182
183def stableadamw(

Callers 1

benchmark_trainingFunction · 0.80

Calls 2

_init_groupMethod · 0.95
stableadamwFunction · 0.85

Tested by

no test coverage detected