(self, parameters: Iterable[torch.nn.Parameter])
| 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 | """ |
no test coverage detected