Clip gradients. Inputs: clip_type (int): The way to clip, 0 for 'value', 1 for 'norm'. clip_value (float): Specifies how much to clip. grad (tuple[Tensor]): Gradients. Outputs: tuple[Tensor], clipped gradients.
(clip_type, clip_value, grad)
| 35 | |
| 36 | @clip_grad.register("Number", "Number", "Tensor") |
| 37 | def _clip_grad(clip_type, clip_value, grad): |
| 38 | """ |
| 39 | Clip gradients. |
| 40 | |
| 41 | Inputs: |
| 42 | clip_type (int): The way to clip, 0 for 'value', 1 for 'norm'. |
| 43 | clip_value (float): Specifies how much to clip. |
| 44 | grad (tuple[Tensor]): Gradients. |
| 45 | |
| 46 | Outputs: |
| 47 | tuple[Tensor], clipped gradients. |
| 48 | """ |
| 49 | if clip_type not in [0, 1]: |
| 50 | return grad |
| 51 | dt = F.dtype(grad) |
| 52 | # 0 for clip_by_value and 1 for clip_by_norm |
| 53 | if clip_type == 0: |
| 54 | new_grad = C.clip_by_value( |
| 55 | grad, |
| 56 | F.cast(F.tuple_to_array((-clip_value,)), dt), |
| 57 | F.cast(F.tuple_to_array((clip_value,)), dt), |
| 58 | ) |
| 59 | else: |
| 60 | new_grad = nn.ClipByNorm()( |
| 61 | grad, F.cast(F.tuple_to_array((clip_value,)), dt) |
| 62 | ) |
| 63 | return new_grad |
| 64 | |
| 65 | |
| 66 | grad_scale = C.MultitypeFuncGraph("grad_scale") |