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

Method construct

codegeex/mindspore/src/adam.py:156–180  ·  view source on GitHub ↗

AdamWeightDecayOp

(self, gradients, clip_value)

Source from the content-addressed store, hash-verified

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 """

Callers

nothing calls this directly

Calls 1

get_lrMethod · 0.80

Tested by

no test coverage detected