@@ -48,7 +48,23 @@ def main() -> None:
4848 ap .add_argument ("--seq" , type = int , default = 4096 )
4949 ap .add_argument ("--backward" , action = "store_true" , help = "also run one backward" )
5050 ap .add_argument ("--vocab" , type = int , default = 65_536 )
51+ ap .add_argument (
52+ "--grad-checkpoint" ,
53+ dest = "grad_checkpoint" ,
54+ action = "store_true" ,
55+ default = None ,
56+ help = "force grad-checkpoint on (default: on iff --backward)" ,
57+ )
58+ ap .add_argument (
59+ "--no-grad-checkpoint" ,
60+ dest = "grad_checkpoint" ,
61+ action = "store_false" ,
62+ help = "force grad-checkpoint off" ,
63+ )
5164 args = ap .parse_args ()
65+ # mx.checkpoint without a paired backward produces a malformed CUDA graph;
66+ # only enable grad-checkpoint when a backward is actually run (or forced).
67+ grad_ckpt = args .backward if args .grad_checkpoint is None else args .grad_checkpoint
5268
5369 efficient = os .environ .get ("CPPMEGA_MOE_EFFICIENT" , "" ).strip ().lower () in {
5470 "1" ,
@@ -57,7 +73,7 @@ def main() -> None:
5773 "yes" ,
5874 }
5975 t0 = time .time ()
60- model = local_gb10_quarter (dtype = mx .bfloat16 , grad_checkpoint = True )
76+ model = local_gb10_quarter (dtype = mx .bfloat16 , grad_checkpoint = grad_ckpt )
6177 mx .eval (model .parameters ())
6278 mx .synchronize ()
6379 build_s = time .time () - t0
@@ -97,6 +113,7 @@ def loss_fn(m, x):
97113 "batch" : args .batch ,
98114 "seq" : args .seq ,
99115 "backward" : bool (args .backward ),
116+ "grad_checkpoint" : bool (grad_ckpt ),
100117 "build_s" : round (build_s , 2 ),
101118 "run_s" : round (run_s , 2 ),
102119 "after_params_peak_gb" : round (after_params_gb , 3 ),
0 commit comments