@@ -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 ,
0 commit comments