MCPcopy Create free account
hub / github.com/drinkingcoder/FlowFormer-Official / build_optimizer

Function build_optimizer

core/optimizer/__init__.py:15–37  ·  view source on GitHub ↗
(model, config)

Source from the content-addressed store, hash-verified

13 return optimizer, scheduler
14
15def 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
40def build_scheduler(config, optimizer):

Callers 1

fetch_optimizerFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected