MCPcopy Create free account
hub / github.com/AlayaLab/Hive / step

Method step

models/flowsep/diffusers/training_utils.py:161–196  ·  view source on GitHub ↗
(self, parameters: Iterable[torch.nn.Parameter])

Source from the content-addressed store, hash-verified

159
160 @torch.no_grad()
161 def step(self, parameters: Iterable[torch.nn.Parameter]):
162 if isinstance(parameters, torch.nn.Module):
163 deprecation_message = (
164 "Passing a `torch.nn.Module` to `ExponentialMovingAverage.step` is deprecated. "
165 "Please pass the parameters of the module instead."
166 )
167 deprecate(
168 "passing a `torch.nn.Module` to `ExponentialMovingAverage.step`",
169 "1.0.0",
170 deprecation_message,
171 standard_warn=False,
172 )
173 parameters = parameters.parameters()
174
175 parameters = list(parameters)
176
177 self.optimization_step += 1
178
179 # Compute the decay factor for the exponential moving average.
180 decay = self.get_decay(self.optimization_step)
181 self.cur_decay_value = decay
182 one_minus_decay = 1 - decay
183
184 context_manager = contextlib.nullcontext
185 if is_transformers_available() and transformers.deepspeed.is_deepspeed_zero3_enabled():
186 import deepspeed
187
188 for s_param, param in zip(self.shadow_params, parameters):
189 if is_transformers_available() and transformers.deepspeed.is_deepspeed_zero3_enabled():
190 context_manager = deepspeed.zero.GatheredParameters(param, modifier_rank=None)
191
192 with context_manager():
193 if param.requires_grad:
194 s_param.sub_(one_minus_decay * (s_param - param))
195 else:
196 s_param.copy_(param)
197
198 def copy_to(self, parameters: Iterable[torch.nn.Parameter]) -> None:
199 """

Callers 2

train_one_epochFunction · 0.45
train_one_epochFunction · 0.45

Calls 3

get_decayMethod · 0.95
deprecateFunction · 0.85

Tested by

no test coverage detected