Skip to content

Commit f70cd32

Browse files
authored
Merge pull request #26 from theomgdev/copilot/improve-chaosgrad-optimizer
Improve ChaosGrad diagnostics correctness and decay stability
2 parents b04a695 + bd7dde2 commit f70cd32

2 files changed

Lines changed: 33 additions & 3 deletions

File tree

odyssnet/training/chaos_optimizer.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -484,7 +484,8 @@ def step(self, closure=None):
484484

485485
# ---- Weight decay ----
486486
if per_decay > 0.0 and not is_hebbian:
487-
p.data.mul_(1.0 - genesis_lr * per_decay)
487+
decay_factor = max(self._EPS, 1.0 - genesis_lr * per_decay)
488+
p.data.mul_(decay_factor)
488489

489490
# ---- Parameter update ----
490491
p.data.add_(v_hat, alpha=-(genesis_lr * per_lr / denom))
@@ -569,10 +570,11 @@ def get_diagnostics(self, debug: bool = False) -> dict:
569570
genesis_lr = self.defaults['lr']
570571
effective_lrs = []
571572
init_lrs = []
572-
for state, _ in all_states:
573+
for state, group in all_states:
573574
if 'per_param_lr' in state and 'init_lr' in state:
575+
group_genesis_lr = group.get('lr', self.defaults['lr'])
574576
effective_lrs.append(state['per_param_lr'] / state['init_lr'])
575-
init_lrs.append(state['init_lr'] * genesis_lr)
577+
init_lrs.append(state['init_lr'] * group_genesis_lr)
576578

577579
diag: dict = {
578580
'global_step': self._global_step,
@@ -601,14 +603,17 @@ def get_diagnostics(self, debug: bool = False) -> dict:
601603
g_states = [self.state[p] for p in group['params'] if self.state.get(p)]
602604
if not g_states:
603605
continue
606+
g_genesis_lr = group.get('lr', self.defaults['lr'])
604607
g_lrs = [s['per_param_lr'] / s['init_lr'] for s in g_states if 'per_param_lr' in s]
608+
g_init_lrs = [s['init_lr'] * g_genesis_lr for s in g_states if 'init_lr' in s]
605609
g_betas = [s['per_param_beta'] for s in g_states if 'per_param_beta' in s]
606610
g_alphas = [s['per_param_alpha'] for s in g_states if 'per_param_alpha' in s]
607611
g_decays = [s['per_param_decay'] for s in g_states if 'per_param_decay' in s]
608612
group_stats.append({
609613
'group_name': gname,
610614
'param_count': len(g_states),
611615
'avg_effective_lr': (sum(g_lrs) / len(g_lrs)) if g_lrs else 0.0,
616+
'avg_init_lr': (sum(g_init_lrs) / len(g_init_lrs)) if g_init_lrs else 0.0,
612617
'avg_beta': (sum(g_betas) / len(g_betas)) if g_betas else 0.0,
613618
'avg_alpha': (sum(g_alphas) / len(g_alphas)) if g_alphas else 0.0,
614619
'avg_decay': (sum(g_decays) / len(g_decays)) if g_decays else 0.0,

tests/training/test_chaos_optimizer_extra.py

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -138,6 +138,31 @@ def test_lr_group_override_propagates(self):
138138
assert torch.allclose(m.W.data.diagonal(), torch.zeros(m.num_neurons))
139139
assert math.isfinite(m.W.data.norm().item())
140140

141+
def test_get_diagnostics_respects_group_lr_override(self):
142+
m = _model()
143+
opt = _opt(m, lr=1e-4)
144+
_one_step_raw(opt, m)
145+
146+
base_diag = opt.get_diagnostics()
147+
base_debug_diag = opt.get_diagnostics(debug=True)
148+
for pg in opt.param_groups:
149+
pg['lr'] = 5e-5
150+
new_diag = opt.get_diagnostics()
151+
new_debug_diag = opt.get_diagnostics(debug=True)
152+
153+
assert new_diag['avg_init_lr'] == pytest.approx(base_diag['avg_init_lr'] * 0.5, rel=1e-6)
154+
155+
base_group_init = {
156+
g['group_name']: g['avg_init_lr']
157+
for g in base_debug_diag['param_groups']
158+
}
159+
new_group_init = {
160+
g['group_name']: g['avg_init_lr']
161+
for g in new_debug_diag['param_groups']
162+
}
163+
for gname, base_avg in base_group_init.items():
164+
assert new_group_init[gname] == pytest.approx(base_avg * 0.5, rel=1e-6)
165+
141166

142167
# ---------------------------------------------------------------------------
143168
# reset_param_state

0 commit comments

Comments
 (0)