@@ -52,6 +52,13 @@ def main() -> None:
5252 ap .add_argument ("--seq" , type = int , default = 4096 )
5353 ap .add_argument ("--vocab" , type = int , default = 65_536 )
5454 ap .add_argument ("--optimizer" , action = "store_true" , help = "also run AdamW update" )
55+ ap .add_argument ("--adam8bit" , action = "store_true" , help = "use int8 AdamW state" )
56+ ap .add_argument (
57+ "--split-eval" ,
58+ dest = "split_eval" ,
59+ action = "store_true" ,
60+ help = "eval grads + free bwd graph before the optimizer update" ,
61+ )
5562 ap .add_argument (
5663 "--grad-checkpoint" ,
5764 dest = "grad_checkpoint" ,
@@ -91,9 +98,19 @@ def main() -> None:
9198
9299 optimizer = None
93100 if args .optimizer :
94- from cppmega_mlx .training .optimizers import make_adamw
95-
96- optimizer = make_adamw (learning_rate = 1e-4 , weight_decay = 0.0 )
101+ if args .adam8bit :
102+ from cppmega_mlx .training .optimizers import make_adam8bit
103+
104+ optimizer = make_adam8bit (
105+ learning_rate = 1e-4 ,
106+ weight_decay = 0.0 ,
107+ quant_scheme = "dynamic_int8_v1" ,
108+ min_8bit_size = 4096 ,
109+ )
110+ else :
111+ from cppmega_mlx .training .optimizers import make_adamw
112+
113+ optimizer = make_adamw (learning_rate = 1e-4 , weight_decay = 0.0 )
97114 optimizer .init (model .trainable_parameters ())
98115 mx .eval (model .parameters (), optimizer .state )
99116
@@ -108,6 +125,12 @@ def loss_fn(m, b):
108125 t1 = time .time ()
109126 (loss , ntok ), grads = nn .value_and_grad (model , loss_fn )(model , batch )
110127 if optimizer is not None :
128+ if args .split_eval :
129+ # Materialise loss+grads and free the backward graph BEFORE the
130+ # optimizer update, so the bwd activations and the optimizer-update
131+ # working set never coexist in a single eval. Peak becomes
132+ # max(bwd, update) instead of bwd+update.
133+ mx .eval (loss , ntok , grads )
111134 optimizer .update (model , grads )
112135 mx .eval (model .parameters (), optimizer .state , loss , ntok )
113136 else :
0 commit comments