Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-30 23:50:55 -07:00
parent 5687c90756
commit ac71c67e9a
7 changed files with 580 additions and 13 deletions
+24 -7
View File
@@ -1,5 +1,5 @@
#!/usr/bin/env python3
"""Fine-tune purpose-lite natively on Apple Silicon with MLX."""
"""Fine-tune purpose-lite with MLX, using Metal by default."""
from __future__ import annotations
@@ -36,17 +36,28 @@ DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1"
DEFAULT_OUTPUT_DIR = SCRIPT_DIR / "outputs" / "purpose-lite-v1-mlx"
def _load_mlx() -> tuple[Any, Any, Any]:
def _configure_mlx_device(mx: Any, device: str) -> None:
if device == "metal":
if not mx.metal.is_available():
raise DataError("MLX Metal training requires Apple Silicon")
mx.set_default_device(mx.gpu)
return
if device == "cpu":
mx.set_default_device(mx.cpu)
return
raise DataError(f"unsupported MLX device {device!r}")
def _load_mlx(device: str) -> tuple[Any, Any, Any]:
try:
import mlx.core as mx
import mlx.nn as nn
import mlx.optimizers as optim
except ImportError as exc:
raise DataError(
"MLX training requires Apple Silicon and requirements-mlx.txt"
"MLX training requires requirements-mlx.txt"
) from exc
if not mx.metal.is_available():
raise DataError("MLX training requires the Apple Silicon Metal backend")
_configure_mlx_device(mx, device)
return mx, nn, optim
@@ -310,7 +321,7 @@ def _linear_schedule(
def train(args: argparse.Namespace) -> dict[str, Any]:
mx, nn, optim = _load_mlx()
mx, nn, optim = _load_mlx(args.device)
try:
from transformers import AutoTokenizer
@@ -718,7 +729,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
"baseModel": str(model_dir),
"baseModelRevision": "local-checkpoint",
"trainingBackend": "mlx",
"device": "metal",
"device": args.device,
"fixedInputShape": [1, MAX_LENGTH],
"truncation": {
"strategy": "head-tail-pair",
@@ -774,6 +785,12 @@ def build_parser() -> argparse.ArgumentParser:
parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR)
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
parser.add_argument("--model", type=Path, required=True)
parser.add_argument(
"--device",
choices=("metal", "cpu"),
default="metal",
help="MLX execution device (default: metal; cpu is a diagnostic fallback)",
)
parser.add_argument("--seed", type=int, default=20260730)
parser.add_argument("--epochs", type=int, default=3)
parser.add_argument("--batch-size", type=int, default=32)