Skip to content

Commit b3ae5f4

Browse files
committed
test(pr2): memory probe adds --opt adam8bit + per-phase peak decomposition
1 parent fda4671 commit b3ae5f4

1 file changed

Lines changed: 5 additions & 2 deletions

File tree

scripts/pr2_seq4096_memory_probe.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
from cppmega_mlx.models.hybrid_lm import HybridTinyConfig, HybridTinyLM
2222
from cppmega_mlx.training.loop import one_step_train
2323
from cppmega_mlx.training.optimizers import make_adamw
24+
from cppmega_mlx.training.optimizers_quantized import make_adam8bit
2425

2526

2627
def _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

Comments
 (0)