Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
+24
-7
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user