MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / _clip_grad

Function _clip_grad

codegeex/mindspore/src/pangu_alpha_wrapcell.py:37–63  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

35
36@clip_grad.register("Number", "Number", "Tensor")
37def _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
66grad_scale = C.MultitypeFuncGraph("grad_scale")

Callers

nothing calls this directly

Calls 1

dtypeMethod · 0.80

Tested by

no test coverage detected