@@ -47,14 +47,10 @@ def forward(self, x):
4747 "eps" : 1e-8 ,
4848 "precondition_1d" : True , # Enable preconditioning for bias vectors
4949 "precondition_frequency" : 1 , # Update preconditioner every step for testing
50- "trace_normalization" : True ,
51- "shampoo_beta" : 0.9 , # Slightly more aggressive moving average
5250 "fp32_matmul_prec" : "high" ,
5351 "qr_fp32_matmul_prec" : "high" ,
5452 "use_adaptive_criteria" : False ,
5553 "power_iter_steps" : 1 ,
56- "use_nesterov" : True ,
57- "skip_preconditioning_steps" : 0 ,
5854}
5955
6056
@@ -106,15 +102,12 @@ def main() -> None:
106102 # Initialize optimizers
107103 optimizer_soap = SOAP (
108104 model_soap .parameters (),
109- lr = 9.0 * config ["lr" ],
105+ lr = 2.1 * config ["lr" ],
110106 weight_decay = config ["weight_decay" ],
111107 betas = (config ["adam_beta1" ], config ["adam_beta2" ]),
112108 eps = config ["eps" ],
113109 precondition_frequency = config ["precondition_frequency" ],
114- trace_normalization = config ["trace_normalization" ],
115- shampoo_beta = config ["shampoo_beta" ],
116110 precondition_1d = config ["precondition_1d" ],
117- use_nesterov = config ["use_nesterov" ],
118111 fp32_matmul_prec = config ["fp32_matmul_prec" ],
119112 qr_fp32_matmul_prec = config ["qr_fp32_matmul_prec" ],
120113 use_adaptive_criteria = config ["use_adaptive_criteria" ],
0 commit comments