-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathoptimizer.py
More file actions
75 lines (61 loc) · 1.92 KB
/
Copy pathoptimizer.py
File metadata and controls
75 lines (61 loc) · 1.92 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
# optimizer.py
import torch
from dataclasses import dataclass
from muon.muon import Muon
@dataclass
class OptimizerConfig:
muon_lr: float = 0.03
adamw_lr: float = 0.008
muon_weight_decay: float = 0.0
adamw_weight_decay: float = 0.0
cautious: bool = True
beta1: float = 0.9
beta2: float = 0.95
muon_momentum: float = 0.95
def configure_optimizers(model, config: OptimizerConfig):
muon_params = []
adamw_params = []
for name, p in model.named_parameters():
if not p.requires_grad:
continue
if p.ndim >= 2:
muon_params.append(p)
else:
adamw_params.append(p)
optimizers = []
if muon_params:
muon = Muon(
muon_params,
lr=config.muon_lr,
weight_decay=config.muon_weight_decay,
momentum=config.muon_momentum,
cautious=config.cautious,
)
optimizers.append(muon)
use_cuda = torch.cuda.is_available()
if adamw_params:
adamw = torch.optim.AdamW(
adamw_params,
lr=config.adamw_lr,
weight_decay=config.adamw_weight_decay,
betas=(config.beta1, config.beta2),
fused=use_cuda,
capturable=use_cuda,
)
optimizers.append(adamw)
print(f"Muon optimizer: {len(muon_params)} parameters")
print(f"AdamW optimizer: {len(adamw_params)} parameters")
return optimizers
def get_optimizers(config, model):
optimizer_config = OptimizerConfig(
muon_lr=config["muon_lr"],
adamw_lr=config["adamw_lr"],
muon_weight_decay=config["muon_weight_decay"],
adamw_weight_decay=config["adamw_weight_decay"],
cautious=config["cautious"],
beta1=config["beta1"],
beta2=config["beta2"],
muon_momentum=config["muon_momentum"],
)
optimizers = configure_optimizers(model, optimizer_config)
return optimizers