(self, lr_kwargs, block_eigenvalue={})
| 3269 | clip_grad_norm_(parameters=self.module.parameters(), max_norm=self.gradient_clipping(), mpu=self.mpu) |
| 3270 | |
| 3271 | def _take_model_step(self, lr_kwargs, block_eigenvalue={}): |
| 3272 | if self.gradient_clipping() > 0.0: |
| 3273 | if self.torch_autocast_z0_gradscaler: |
| 3274 | # Unscale for gradient clipping |
| 3275 | self.torch_autocast_z0_gradscaler.unscale_(self.optimizer) |
| 3276 | if not (self.fp16_enabled() or self.bfloat16_enabled() or self.amp_enabled() or self.zero_optimization()): |
| 3277 | self.clip_fp32_gradients() |
| 3278 | elif self.amp_enabled(): |
| 3279 | # AMP's recommended way of doing clipping |
| 3280 | # https://nvidia.github.io/apex/advanced.html#gradient-clipping |
| 3281 | master_params = amp.master_params(self.optimizer) |
| 3282 | clip_grad_norm_(parameters=master_params, max_norm=self.gradient_clipping(), mpu=self.mpu) |
| 3283 | if self.torch_autocast_z0_gradscaler: |
| 3284 | self.torch_autocast_z0_gradscaler.step(self.optimizer) |
| 3285 | self.torch_autocast_z0_gradscaler.update() |
| 3286 | else: |
| 3287 | self.optimizer.step() |
| 3288 | |
| 3289 | if hasattr(self.optimizer, '_global_grad_norm'): |
| 3290 | self._global_grad_norm = self.optimizer._global_grad_norm |
| 3291 | |
| 3292 | # Quantize the updated parameter if there is no overflow |
| 3293 | if self.quantizer: |
| 3294 | tensor_to_quantize = self.optimizer.bit16_groups if self.zero_optimization_stage( |
| 3295 | ) == 2 else self.optimizer.fp16_groups |
| 3296 | if self.compression_scheduler.weight_quantization_enabled: |
| 3297 | self.quantizer.quantize( |
| 3298 | tensor_to_quantize, |
| 3299 | (self.optimizer.overflow if self.fp16_enabled() else False), |
| 3300 | self.eigenvalue_enabled(), |
| 3301 | block_eigenvalue, |
| 3302 | ) |
| 3303 | # zero grad in basic optimizer could be unreliable and may not exhibit |
| 3304 | # the behavior that we want |
| 3305 | if self.bfloat16_enabled(): |
| 3306 | # TODO: Temporary until bf16_optimizer and zero_optimizer are integrated |
| 3307 | if hasattr(self.optimizer, "zero_grad"): |
| 3308 | self.optimizer.zero_grad() |
| 3309 | else: |
| 3310 | self.zero_grad() |
| 3311 | elif self.zero_optimization() or self.fp16_enabled() or self.amp_enabled(): |
| 3312 | self.optimizer.zero_grad() |
| 3313 | else: |
| 3314 | self.zero_grad() |
| 3315 | |
| 3316 | # Check overflow here since in DS fp16 optimizer, the overflow is updated in above step() function. |
| 3317 | overflow = False |
| 3318 | if hasattr(self.optimizer, "overflow"): |
| 3319 | overflow = self.optimizer.overflow |
| 3320 | self._step_applied = not overflow |
| 3321 | |
| 3322 | if overflow: |
| 3323 | self.skipped_steps += 1 |
| 3324 | else: |
| 3325 | self.compression_scheduler.step() |
| 3326 | if self.lr_scheduler is not None: |
| 3327 | try: |
| 3328 | self.lr_scheduler.step(**(lr_kwargs or {})) |
no test coverage detected