AdamWeightDecayOp
(self, gradients, clip_value)
| 154 | self.opt.add_prim_attr("primitive_target", "CPU") |
| 155 | |
| 156 | def construct(self, gradients, clip_value): |
| 157 | """AdamWeightDecayOp""" |
| 158 | lr = self.get_lr() |
| 159 | cond = P.GreaterEqual()(clip_value, self.clip_norm) |
| 160 | global_norm = F.select(cond, clip_value, self.clip_norm) |
| 161 | global_norm = P.Cast()(global_norm, mstype.float16) |
| 162 | if self.is_group: |
| 163 | if self.is_group_lr: |
| 164 | optim_result = self.map_reverse(F.partial(_adam_opt, self.opt, global_norm, |
| 165 | self.beta1, self.beta2, self.eps), |
| 166 | lr, self.weight_decay, self.parameters, self.moments1, self.moments2, |
| 167 | gradients, self.decay_flags, self.optim_filter) |
| 168 | else: |
| 169 | optim_result = self.map_reverse(F.partial(_adam_opt, self.opt, global_norm, |
| 170 | self.beta1, self.beta2, self.eps, lr), |
| 171 | self.weight_decay, self.parameters, self.moments1, self.moments2, |
| 172 | gradients, self.decay_flags, self.optim_filter) |
| 173 | else: |
| 174 | optim_result = self.map_reverse(F.partial(_adam_opt, self.opt, global_norm, |
| 175 | self.beta1, self.beta2, self.eps, lr, |
| 176 | self.weight_decay), self.parameters, self.moments1, self.moments2, |
| 177 | gradients, self.decay_flags, self.optim_filter) |
| 178 | if self.use_parallel: |
| 179 | self.broadcast_params(optim_result) |
| 180 | return optim_result |
| 181 | |
| 182 | def clone_param32(self, prefix, init=None): |
| 183 | """ |