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

Class ClipByGlobalNorm

codegeex/mindspore/src/utils.py:206–229  ·  view source on GitHub ↗

Clip grads by global norm

Source from the content-addressed store, hash-verified

204
205
206class ClipByGlobalNorm(nn.Cell):
207 """
208
209 Clip grads by global norm
210
211 """
212
213 def __init__(self, params, config, clip_norm=1.0):
214 super(ClipByGlobalNorm, self).__init__()
215 self.global_norm = GlobalNorm(params, config)
216 self.clip_norm = Tensor([clip_norm], mstype.float32)
217 self.hyper_map = C.HyperMap()
218 if config.param_init_type == mstype.float16 and config.enable_offload:
219 self.enable_grad_fp16 = True
220 else:
221 self.enable_grad_fp16 = False
222
223 def construct(self, grads):
224 """Clip grads by global norm construct"""
225 grads, global_norm_value = self.global_norm(grads)
226 cond = P.GreaterEqual()(global_norm_value, self.clip_norm)
227 global_norm = F.select(cond, global_norm_value, self.clip_norm)
228 grads = self.hyper_map(F.partial(apply_global_norm, self.enable_grad_fp16, self.clip_norm, global_norm), grads)
229 return grads, global_norm_value
230
231
232class LearningRate(LearningRateSchedule):

Callers 4

__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected