Skip to content

Commit 03e536b

Browse files
committed
tooling(probe): add --split-eval (free bwd graph before optimizer) and --adam8bit
1 parent cecdd67 commit 03e536b

1 file changed

Lines changed: 26 additions & 3 deletions

File tree

scripts/probe_real_step_mem_20260601.py

Lines changed: 26 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)