MCPcopy Create free account
hub / github.com/baidu/DDParser / AdamW

Class AdamW

ddparser/ernie/optimization.py:144–159  ·  view source on GitHub ↗

AdamW object for dygraph

Source from the content-addressed store, hash-verified

142
143
144class AdamW(F.optimizer.AdamOptimizer):
145 """AdamW object for dygraph"""
146 def __init__(self, *args, **kwargs):
147 weight_decay = kwargs.pop('weight_decay', None)
148 var_name_to_exclude = kwargs.pop('var_name_to_exclude', '.*layer_norm_scale|.*layer_norm_bias|.*b_0')
149 super(AdamW, self).__init__(*args, **kwargs)
150 self.wd = weight_decay
151 self.pat = re.compile(var_name_to_exclude)
152
153 def apply_optimize(self, loss, startup_program, params_grads):
154 super(AdamW, self).apply_optimize(loss, startup_program, params_grads)
155 for p, g in params_grads:
156 #log.debug(L.reduce_mean(p))
157 if not self.pat.match(p.name):
158 L.assign(p * (1. - self.wd * self.current_step_lr()), p)
159 #log.debug(L.reduce_mean(p))
160
161
162class LinearDecay(D.learning_rate_scheduler.LearningRateDecay):

Callers 1

trainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected