Skip to content

Commit 82de04a

Browse files
committed
fix(probe): grad-checkpoint only with backward (forward-only mx.checkpoint makes malformed CUDA graph)
1 parent 0b48e05 commit 82de04a

1 file changed

Lines changed: 18 additions & 1 deletion

File tree

scripts/probe_moe_forward_mem.py

Lines changed: 18 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)