(enable_grad_fp16, clip_norm, global_norm, grad)
| 88 | |
| 89 | @apply_global_norm.register("Bool", "Tensor", "Tensor", "Tensor") |
| 90 | def _apply_global_norm(enable_grad_fp16, clip_norm, global_norm, grad): |
| 91 | if enable_grad_fp16: |
| 92 | grad = P.Cast()(grad * clip_norm / global_norm, mstype.float16) |
| 93 | else: |
| 94 | grad = grad * clip_norm / global_norm |
| 95 | return grad |
| 96 | |
| 97 | |
| 98 | def _get_model_parallel_group(mp): |
nothing calls this directly
no outgoing calls
no test coverage detected