(self, lr_kwargs, block_eigenvalue={})
| 3166 | clip_grad_norm_(parameters=self.module.parameters(), max_norm=self.gradient_clipping(), mpu=self.mpu) |
| 3167 | |
| 3168 | def _take_model_step(self, lr_kwargs, block_eigenvalue={}): |
| 3169 | if self.gradient_clipping() > 0.0: |
| 3170 | if self.torch_autocast_z0_gradscaler: |
| 3171 | # Unscale for gradient clipping |
| 3172 | self.torch_autocast_z0_gradscaler.unscale_(self.optimizer) |
| 3173 | if not (self.fp16_enabled() or self.bfloat16_enabled() or self.amp_enabled() or self.zero_optimization()): |
| 3174 | self.clip_fp32_gradients() |
| 3175 | elif self.amp_enabled(): |
| 3176 | # AMP's recommended way of doing clipping |
| 3177 | # https://nvidia.github.io/apex/advanced.html#gradient-clipping |
| 3178 | master_params = amp.master_params(self.optimizer) |
| 3179 | clip_grad_norm_(parameters=master_params, max_norm=self.gradient_clipping(), mpu=self.mpu) |
| 3180 | if self.torch_autocast_z0_gradscaler: |
| 3181 | self.torch_autocast_z0_gradscaler.step(self.optimizer) |
| 3182 | self.torch_autocast_z0_gradscaler.update() |
| 3183 | else: |
| 3184 | self.optimizer.step() |
| 3185 | |
| 3186 | if hasattr(self.optimizer, '_global_grad_norm'): |
| 3187 | self._global_grad_norm = self.optimizer._global_grad_norm |
| 3188 | |
| 3189 | # Quantize the updated parameter if there is no overflow |
| 3190 | if self.quantizer: |
| 3191 | tensor_to_quantize = self.optimizer.bit16_groups if self.zero_optimization_stage( |
| 3192 | ) == 2 else self.optimizer.fp16_groups |
| 3193 | if self.compression_scheduler.weight_quantization_enabled: |
| 3194 | self.quantizer.quantize( |
| 3195 | tensor_to_quantize, |
| 3196 | (self.optimizer.overflow if self.fp16_enabled() else False), |
| 3197 | self.eigenvalue_enabled(), |
| 3198 | block_eigenvalue, |
| 3199 | ) |
| 3200 | # zero grad in basic optimizer could be unreliable and may not exhibit |
| 3201 | # the behavior that we want |
| 3202 | if self.bfloat16_enabled(): |
| 3203 | # TODO: Temporary until bf16_optimizer and zero_optimizer are integrated |
| 3204 | if hasattr(self.optimizer, "zero_grad"): |
| 3205 | self.optimizer.zero_grad() |
| 3206 | else: |
| 3207 | self.zero_grad() |
| 3208 | elif self.zero_optimization() or self.fp16_enabled() or self.amp_enabled(): |
| 3209 | self.optimizer.zero_grad() |
| 3210 | else: |
| 3211 | self.zero_grad() |
| 3212 | |
| 3213 | # Check overflow here since in DS fp16 optimizer, the overflow is updated in above step() function. |
| 3214 | overflow = False |
| 3215 | if hasattr(self.optimizer, "overflow"): |
| 3216 | overflow = self.optimizer.overflow |
| 3217 | self._step_applied = not overflow |
| 3218 | |
| 3219 | if overflow: |
| 3220 | self.skipped_steps += 1 |
| 3221 | else: |
| 3222 | self.compression_scheduler.step() |
| 3223 | if self.lr_scheduler is not None: |
| 3224 | try: |
| 3225 | self.lr_scheduler.step(**(lr_kwargs or {})) |
no test coverage detected