@@ -28,11 +28,12 @@ def _gib(n: int | None) -> float:
2828 return (n or 0 ) / (1024 ** 3 )
2929
3030
31- def build_model (hidden : int , depth : int , vocab : int , seq : int , grad_ckpt : bool ) -> HybridTinyLM :
31+ def build_model (hidden : int , depth : int , vocab : int , seq : int , grad_ckpt : bool ,
32+ pattern : str = "AEMR" ) -> HybridTinyLM :
3233 cfg = HybridTinyConfig (
3334 vocab_size = vocab ,
3435 hidden_size = hidden ,
35- pattern = "AEMR" ,
36+ pattern = pattern ,
3637 depth = depth ,
3738 num_attention_heads = max (1 , hidden // 64 ),
3839 max_seq_length = seq ,
@@ -57,14 +58,15 @@ def main() -> None:
5758 ap .add_argument ("--grad-ckpt" , action = "store_true" )
5859 ap .add_argument ("--clear-cache" , action = "store_true" )
5960 ap .add_argument ("--opt" , choices = ["adamw" , "adam8bit" ], default = "adamw" )
61+ ap .add_argument ("--pattern" , type = str , default = "AEMR" )
6062 ap .add_argument ("--steps" , type = int , default = 2 )
6163 args = ap .parse_args ()
6264
6365 mx .random .seed (0 )
6466 if hasattr (mx , "reset_peak_memory" ):
6567 mx .reset_peak_memory ()
6668
67- model = build_model (args .hidden , args .depth , args .vocab , args .seq , args .grad_ckpt )
69+ model = build_model (args .hidden , args .depth , args .vocab , args .seq , args .grad_ckpt , args . pattern )
6870 nparams = sum (v .size for _ , v in __import__ ("mlx.utils" , fromlist = ["tree_flatten" ]).tree_flatten (model .parameters ()))
6971 print (f"after-model-build peak={ _gib (mx .get_peak_memory () if hasattr (mx ,'get_peak_memory' ) else None ):.2f} GiB" )
7072 opt = make_adam8bit (learning_rate = 1e-4 ) if args .opt == "adam8bit" else make_adamw (learning_rate = 1e-4 )
0 commit comments