Clip grads by global norm
| 204 | |
| 205 | |
| 206 | class 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 | |
| 232 | class LearningRate(LearningRateSchedule): |