2121from cppmega_mlx .models .hybrid_lm import HybridTinyConfig , HybridTinyLM
2222from cppmega_mlx .training .loop import one_step_train
2323from cppmega_mlx .training .optimizers import make_adamw
24+ from cppmega_mlx .training .optimizers_quantized import make_adam8bit
2425
2526
2627def _gib (n : int | None ) -> float :
@@ -55,6 +56,7 @@ def main() -> None:
5556 ap .add_argument ("--vocab" , type = int , default = 32000 )
5657 ap .add_argument ("--grad-ckpt" , action = "store_true" )
5758 ap .add_argument ("--clear-cache" , action = "store_true" )
59+ ap .add_argument ("--opt" , choices = ["adamw" , "adam8bit" ], default = "adamw" )
5860 ap .add_argument ("--steps" , type = int , default = 2 )
5961 args = ap .parse_args ()
6062
@@ -64,12 +66,13 @@ def main() -> None:
6466
6567 model = build_model (args .hidden , args .depth , args .vocab , args .seq , args .grad_ckpt )
6668 nparams = sum (v .size for _ , v in __import__ ("mlx.utils" , fromlist = ["tree_flatten" ]).tree_flatten (model .parameters ()))
67- opt = make_adamw (learning_rate = 1e-4 )
69+ print (f"after-model-build peak={ _gib (mx .get_peak_memory () if hasattr (mx ,'get_peak_memory' ) else None ):.2f} GiB" )
70+ opt = make_adam8bit (learning_rate = 1e-4 ) if args .opt == "adam8bit" else make_adamw (learning_rate = 1e-4 )
6871 batch = synthetic_token_batch (batch_size = args .batch , seq_length = args .seq , vocab_size = args .vocab )
6972
7073 print (f"config: seq={ args .seq } batch={ args .batch } grad_accum={ args .grad_accum } "
7174 f"hidden={ args .hidden } depth={ args .depth } grad_ckpt={ args .grad_ckpt } "
72- f"clear_cache={ args .clear_cache } params={ nparams / 1e6 :.1f} M" )
75+ f"opt= { args . opt } clear_cache={ args .clear_cache } params={ nparams / 1e6 :.1f} M" )
7376
7477 for step in range (args .steps ):
7578 t0 = time .perf_counter ()
0 commit comments