Given a model and its bert module, create parameter groups with different lr
| 87 | |
| 88 | @registry.register('optimizer', 'bertAdamw') |
| 89 | class BertAdamW(transformers.AdamW): |
| 90 | """ |
| 91 | Given a model and its bert module, create parameter groups with different lr |
| 92 | """ |
| 93 | def __init__(self, non_bert_params, bert_params, lr=1e-3, bert_lr=2e-5, **kwargs): |
| 94 | self.bert_param_group = {"params" : bert_params , "lr": bert_lr, "weight_decay": 0} |
| 95 | self.non_bert_param_group = {"params" : non_bert_params} |
| 96 | |
| 97 | params = [self.non_bert_param_group, self.bert_param_group] |
| 98 | if "name" in kwargs: del kwargs["name"] #TODO: fix this |
| 99 | super(BertAdamW, self).__init__(params, lr=lr, **kwargs) |
| 100 | |
| 101 | @registry.register('lr_scheduler', 'bert_warmup_polynomial_group') |
| 102 | @attr.s |
nothing calls this directly
no outgoing calls
no test coverage detected