From c7e3989d7218f25b355936871bc1bbecaa31f9d3 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 7 May 2026 03:39:55 +0000 Subject: [PATCH] ci(mlx): real MLX training + inference smoke test on Mac M1 Add tests/studio/run_real_mlx_smoke.py and wire it into the macos-14 job as the final step. The script trains unsloth/gemma-3-270m-it for 7 deterministic LoRA steps on an in-memory dataset of the SAME row repeated: "<> My name is Unsloth!" then prompts the trained model with "<> My name is " and asserts the completion contains "Unsloth". Captures and asserts: - per-step training loss (via MLXTrainer.add_step_callback); - pre- and post-training loss + gradient norm (computed manually via mx.nn.value_and_grad over the training row, since MLXTrainer does not currently expose per-step grad norms); - losses are finite, do not diverge, and post-train loss < pre-train; - grad norms are finite and positive; - the inference output contains "Unsloth". Determinism: seeds python random, numpy, and mlx.core.random; passes random_state=SEED to FastMLXModel.from_pretrained and get_peft_model (both invoke _seed_mlx_random_state internally) and seed=SEED to MLXTrainingConfig (drives batch shuffling). Uses fp16 + no quant (gemma-3-270m is small enough to skip 4-bit) and LoRA r=8 on the four attention projections. This is the only place in CI that exercises a real MLX backward pass + optimizer step + mlx_lm.generate call. --- .github/workflows/mlx-ci.yml | 23 +++ tests/studio/run_real_mlx_smoke.py | 266 +++++++++++++++++++++++++++++ 2 files changed, 289 insertions(+) create mode 100644 tests/studio/run_real_mlx_smoke.py diff --git a/.github/workflows/mlx-ci.yml b/.github/workflows/mlx-ci.yml index 6fa88df01b..7041b9efc9 100644 --- a/.github/workflows/mlx-ci.yml +++ b/.github/workflows/mlx-ci.yml @@ -28,6 +28,13 @@ # environment (the test fixture installs a MetaPathFinder that # blocks `import mlx.core` for "no-mlx" profiles, faithfully # simulating a Mac without mlx even when mlx IS installed). +# 4. End-to-end MLX training + inference smoke test: +# run_real_mlx_smoke.py trains unsloth/gemma-3-270m-it for 7 +# deterministic LoRA steps on a single repeated text row, then +# verifies the trained model can complete the prompt and that +# losses + grad norms are finite and well-behaved. This is the +# only place in CI that exercises a real MLX backward pass + +# optimizer step + inference call. # # Three dispatch test files documented in tests/studio/README.md: # - test_hardware_dispatch_matrix.py parametrized 7-profile matrix @@ -60,6 +67,7 @@ on: - 'tests/studio/test_hardware_dispatch_matrix.py' - 'tests/studio/test_is_mlx_dispatch_gate.py' - 'tests/studio/test_mlx_training_worker_behaviors.py' + - 'tests/studio/run_real_mlx_smoke.py' - 'tests/conftest.py' - '.github/workflows/mlx-ci.yml' push: @@ -194,3 +202,18 @@ jobs: tests/studio/test_hardware_dispatch_matrix.py \ tests/studio/test_is_mlx_dispatch_gate.py \ tests/studio/test_mlx_training_worker_behaviors.py + + # Real MLX training + inference smoke test. Trains + # unsloth/gemma-3-270m-it for 7 deterministic LoRA steps on a + # single repeated row ("<> My name is Unsloth!"), + # captures per-step losses and pre/post-training grad norms, + # then completes "<> My name is " and asserts the + # generation contains "Unsloth". This is the only place in CI + # that exercises the real MLX backward pass + optimizer step + + # inference path end to end. + - name: Real MLX training + inference smoke test + env: + HF_TOKEN: ${{ secrets.HF_TOKEN }} + UNSLOTH_COMPILE_DISABLE: '1' + run: | + python tests/studio/run_real_mlx_smoke.py diff --git a/tests/studio/run_real_mlx_smoke.py b/tests/studio/run_real_mlx_smoke.py new file mode 100644 index 0000000000..e2e32bbec1 --- /dev/null +++ b/tests/studio/run_real_mlx_smoke.py @@ -0,0 +1,266 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. + +""" +End-to-end MLX smoke test on real Apple Silicon. + +Trains `unsloth/gemma-3-270m-it` for 7 deterministic LoRA steps on an +in-memory dataset of the SAME row repeated: + + "<> My name is Unsloth!" + +then asks the trained model to complete the prompt + + "<> My name is " + +and asserts the completion contains "Unsloth". + +Captures and asserts: + + - Per-step training loss (from MLXTrainer's add_step_callback). + - Loss is finite and does not diverge across the 7 steps. + - Pre- and post-training gradient norms (computed manually via + mx.nn.value_and_grad over a single batch of the training text; + the trainer does not currently expose per-step grad norms). + - Inference output contains "Unsloth". + +This script is only runnable on a real Apple Silicon host (the import +chain pulls real `mlx`, `mlx-lm`, and `unsloth_zoo.mlx_*`). It is +invoked from .github/workflows/mlx-ci.yml on the macos-14 runner. + +Determinism: seeds Python `random`, `numpy`, and `mlx.core.random` +before any MLX import, and forwards `random_state=SEED` to both +`FastMLXModel.from_pretrained` and `FastMLXModel.get_peft_model` +(both call `_seed_mlx_random_state` internally), and `seed=SEED` to +`MLXTrainingConfig` (drives batch shuffling). Metal still has minor +nondeterminism from reduction-order in atomics, so loss assertions +are bounds rather than exact-match. +""" + +from __future__ import annotations + +import math +import os +import random as _random +import sys + +import numpy as np + + +SEED = 3407 + + +def _seed_everything() -> None: + _random.seed(SEED) + np.random.seed(SEED) + # mlx.core.random must be seeded after the import; we can't avoid + # the import here. This must run BEFORE FastMLXModel.from_pretrained. + import mlx.core as mx + + mx.random.seed(SEED) + + +def _compute_loss_and_grad_norm(model, tokenizer, text: str) -> tuple[float, float]: + """Run one forward+backward over a single training example and + return (loss, ||grad||_2) so we can compare pre- vs post-training. + + Uses the same next-token cross-entropy loss the trainer uses (no + masking — the tiny synthetic dataset has no instruction/response + split). + """ + import mlx.core as mx + import mlx.nn as nn + from mlx.utils import tree_flatten + + ids = list(tokenizer.encode(text)) + eos_id = getattr(tokenizer, "eos_token_id", None) + if eos_id is not None: + ids.append(int(eos_id)) + if len(ids) < 2: + raise RuntimeError( + f"tokenized text too short to compute loss: {len(ids)} tokens" + ) + + inputs = mx.array([ids[:-1]], dtype=mx.int32) + targets = mx.array([ids[1:]], dtype=mx.int32) + + def loss_fn(m): + logits = m(inputs) + return nn.losses.cross_entropy(logits, targets, reduction="mean") + + loss_and_grad = nn.value_and_grad(model, loss_fn) + loss_val, grad = loss_and_grad(model) + + norm_sq = mx.array(0.0, dtype=mx.float32) + for _name, value in tree_flatten(grad): + norm_sq = norm_sq + mx.sum(value.astype(mx.float32) * value.astype(mx.float32)) + grad_norm = mx.sqrt(norm_sq) + + return float(loss_val.item()), float(grad_norm.item()) + + +def main() -> int: + _seed_everything() + + import mlx.core as mx + from unsloth_zoo.mlx_loader import FastMLXModel + from unsloth_zoo.mlx_trainer import MLXTrainer, MLXTrainingConfig + + text_row = "<> My name is Unsloth!" + model_name = "unsloth/gemma-3-270m-it" + hf_token = os.environ.get("HF_TOKEN") or None + + print(f"Loading {model_name} (fp16, no quant)...", flush=True) + model, tokenizer = FastMLXModel.from_pretrained( + model_name, + load_in_4bit=False, + dtype="float16", + text_only=True, + max_seq_length=128, + random_state=SEED, + token=hf_token, + trust_remote_code=False, + ) + + # Re-seed RNG between load and LoRA injection so the LoRA init is + # reproducible regardless of how many random draws the loader did. + mx.random.seed(SEED) + + print("Applying LoRA r=8 on attention modules...", flush=True) + model = FastMLXModel.get_peft_model( + model, + r=8, + lora_alpha=16, + lora_dropout=0.0, + target_modules=["q_proj", "k_proj", "v_proj", "o_proj"], + use_gradient_checkpointing=False, + random_state=SEED, + finetune_language_layers=True, + finetune_attention_modules=True, + finetune_mlp_modules=False, + ) + + # Tiny synthetic in-memory dataset: same row repeated. The trainer + # consumes any iterable of dicts with the dataset_text_field key. + dataset = [{"text": text_row}] * 32 + + print("Pre-training loss + grad norm (single-batch probe)...", flush=True) + pre_loss, pre_grad_norm = _compute_loss_and_grad_norm(model, tokenizer, text_row) + print(f" pre loss={pre_loss:.4f} grad_norm={pre_grad_norm:.4f}", flush=True) + assert math.isfinite(pre_loss), f"pre-train loss is non-finite: {pre_loss}" + assert math.isfinite(pre_grad_norm), ( + f"pre-train grad_norm is non-finite: {pre_grad_norm}" + ) + assert pre_grad_norm > 0, f"pre-train grad_norm is zero: {pre_grad_norm}" + + print("Constructing MLXTrainer (max_steps=7, lr=1e-3, bs=2)...", flush=True) + config = MLXTrainingConfig( + per_device_train_batch_size=2, + gradient_accumulation_steps=1, + max_steps=7, + learning_rate=1e-3, + warmup_steps=0, + lr_scheduler_type="constant", + optim="adamw", + weight_decay=0.0, + max_grad_norm=1.0, + logging_steps=1, + max_seq_length=64, + seed=SEED, + use_cce=False, + compile=False, + gradient_checkpointing=False, + output_dir="/tmp/unsloth_mlx_smoke", + save_steps=0, + eval_steps=0, + dataset_text_field="text", + ) + + trainer = MLXTrainer( + model=model, + tokenizer=tokenizer, + train_dataset=dataset, + args=config, + ) + + losses: list[tuple[int, float]] = [] + lrs: list[tuple[int, float]] = [] + + def _on_step(step, total, loss, lr, tok_s, peak_gb, elapsed, num_tokens): + losses.append((int(step), float(loss))) + lrs.append((int(step), float(lr))) + print( + f" step {step}/{total} loss={loss:.4f} lr={lr:.2e} " + f"tok/s={tok_s:.0f} peak={peak_gb:.2f}GB", + flush=True, + ) + + trainer.add_step_callback(_on_step) + + print("Running 7 training steps...", flush=True) + train_result = trainer.train() + print(f"Trainer summary: {train_result}", flush=True) + + print("Post-training loss + grad norm (single-batch probe)...", flush=True) + post_loss, post_grad_norm = _compute_loss_and_grad_norm(model, tokenizer, text_row) + print(f" post loss={post_loss:.4f} grad_norm={post_grad_norm:.4f}", flush=True) + + # Loss + grad norm assertions + assert len(losses) == 7, f"expected 7 step callbacks, got {len(losses)}: {losses}" + for step, loss in losses: + assert math.isfinite(loss), f"step {step} loss not finite: {loss}" + assert 0 < loss < 50, f"step {step} loss out of bounds: {loss}" + + first_loss = losses[0][1] + last_loss = losses[-1][1] + print(f"loss[0]={first_loss:.4f} loss[6]={last_loss:.4f}", flush=True) + # On a single repeated row the model should bend towards the data. + # Allow some headroom for Metal nondeterminism but require we are + # not wildly diverging. + assert last_loss < first_loss * 1.1, ( + f"loss diverged across 7 steps: first={first_loss:.4f} " + f"last={last_loss:.4f}" + ) + assert math.isfinite(post_loss), f"post-train loss not finite: {post_loss}" + assert math.isfinite(post_grad_norm), ( + f"post-train grad_norm not finite: {post_grad_norm}" + ) + assert post_loss < pre_loss, ( + f"post-train loss {post_loss:.4f} >= pre-train loss {pre_loss:.4f} — " + f"7 steps of LoRA on a single repeated row should reduce loss" + ) + + # Inference: prompt -> "Unsloth" continuation + print("Inference: completing '<> My name is '...", flush=True) + from mlx_lm import generate + + model.eval() + prompt = "<> My name is " + output = generate( + model, + tokenizer, + prompt=prompt, + max_tokens=24, + verbose=False, + ) + print(f" prompt: {prompt!r}", flush=True) + print(f" output: {output!r}", flush=True) + assert "Unsloth" in output, ( + f"expected 'Unsloth' in completion of {prompt!r}; got {output!r}. " + f"Loss went {first_loss:.4f}->{last_loss:.4f}, post={post_loss:.4f}, " + f"pre_grad_norm={pre_grad_norm:.4f} post_grad_norm={post_grad_norm:.4f}." + ) + + print( + f"\nOK: real-MLX training+inference smoke passed.\n" + f" losses: {[round(l, 4) for _, l in losses]}\n" + f" pre loss={pre_loss:.4f} grad_norm={pre_grad_norm:.4f}\n" + f" post loss={post_loss:.4f} grad_norm={post_grad_norm:.4f}\n" + f" generation: {output!r}", + flush=True, + ) + return 0 + + +if __name__ == "__main__": + sys.exit(main())