# # zyquo_train.py # Zyquo MLX # # Author: Simon-Pierre Boucher # Mail: contact@spboucher.ai # # Training driver for the Zyquo MLX Swift app. # # Wraps mlx_lm.lora.run() with a custom TrainingCallback and emits a stable # JSON-lines protocol on stdout (one JSON object per line, event-typed). # We NEVER rely on mlx-lm's own stdout format (it changed between 0.31.3 and # main — docs/TRAINING-RESEARCH.md §5). Pinned against mlx-lm==0.31.3. # # Usage: zyquo_train.py --config # # Events: {"event":"start", ...} {"event":"train", ...} {"event":"val", ...} # {"event":"save", ...} {"event":"done"} {"event":"error", ...} import argparse import json import sys import time import types def emit(obj): sys.stdout.write(json.dumps(obj) + "\n") sys.stdout.flush() class ZyquoCallback: """Receives the stable mlx-lm TrainingCallback dict payloads (docs/TRAINING-RESEARCH.md §5.2) and re-emits them as JSON lines.""" def on_train_loss_report(self, info): emit({"event": "train", **info, "ts": time.time()}) def on_val_loss_report(self, info): emit({"event": "val", **info, "ts": time.time()}) def main(): parser = argparse.ArgumentParser() parser.add_argument("--config", required=True, help="YAML run config (mlx-lm schema)") args = parser.parse_args() try: import numpy as np import yaml from mlx_lm import lora from mlx_lm.tuner.datasets import load_dataset from mlx_lm.utils import load with open(args.config) as f: config = yaml.safe_load(f) # Build the args namespace exactly like mlx_lm.lora's CLI does: # defaults first, then config overrides. run_args = dict(lora.CONFIG_DEFAULTS) run_args.update(config) ns = types.SimpleNamespace(**run_args) emit({ "event": "start", "model": ns.model, "fine_tune_type": ns.fine_tune_type, "iters": ns.iters, "batch_size": ns.batch_size, "learning_rate": ns.learning_rate, "adapter_path": ns.adapter_path, }) # NOTE: we deliberately do NOT call lora.run() — in mlx-lm 0.31.3 it # overwrites the training_callback argument with # get_reporting_callbacks(args.report_to) (None here), silently # discarding ours. Replicate run()'s exact flow instead. np.random.seed(ns.seed) emit({"event": "stage", "stage": "loading_model"}) model, tokenizer = load(ns.model, tokenizer_config={"trust_remote_code": True}) emit({"event": "stage", "stage": "loading_datasets"}) train_set, valid_set, _test_set = load_dataset(ns, tokenizer) emit({"event": "stage", "stage": "training"}) lora.train_model(ns, model, train_set, valid_set, ZyquoCallback()) emit({"event": "done"}) except KeyboardInterrupt: emit({"event": "error", "message": "cancelled"}) sys.exit(130) except Exception as exc: # noqa: BLE001 - single funnel to the app emit({"event": "error", "message": str(exc), "type": type(exc).__name__}) sys.exit(1) if __name__ == "__main__": main()