AdamW object for dygraph
| 142 | |
| 143 | |
| 144 | class 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 | |
| 162 | class LinearDecay(D.learning_rate_scheduler.LearningRateDecay): |