| 13 | return optimizer, scheduler |
| 14 | |
| 15 | def build_optimizer(model, config): |
| 16 | name = config.optimizer |
| 17 | lr = config.canonical_lr |
| 18 | |
| 19 | if name == "adam": |
| 20 | return torch.optim.Adam(model.parameters(), lr=lr, weight_decay=config.adam_decay, eps=config.epsilon) |
| 21 | elif name == "adamw": |
| 22 | if hasattr(config, 'twins_lr_factor'): |
| 23 | factor = config.twins_lr_factor |
| 24 | print("[Decrease lr of pre-trained model by factor {}]".format(factor)) |
| 25 | param_dicts = [ |
| 26 | {"params": [p for n, p in model.named_parameters() if "feat_encoder" not in n and 'context_encoder' not in n and p.requires_grad]}, |
| 27 | { |
| 28 | "params": [p for n, p in model.named_parameters() if ("feat_encoder" in n or 'context_encoder' in n) and p.requires_grad], |
| 29 | "lr": lr*factor, |
| 30 | }, |
| 31 | ] |
| 32 | full = [n for n, _ in model.named_parameters()] |
| 33 | return torch.optim.AdamW(param_dicts, lr=lr, weight_decay=config.adamw_decay, eps=config.epsilon) |
| 34 | else: |
| 35 | return torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=config.adamw_decay, eps=config.epsilon) |
| 36 | else: |
| 37 | raise ValueError(f"TRAINER.OPTIMIZER = {name} is not a valid optimizer!") |
| 38 | |
| 39 | |
| 40 | def build_scheduler(config, optimizer): |