mlxtuner 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- mlxtuner/__init__.py +3 -0
- mlxtuner/cli.py +384 -0
- mlxtuner/config.py +166 -0
- mlxtuner/data.py +323 -0
- mlxtuner/eval.py +183 -0
- mlxtuner/hardware.py +163 -0
- mlxtuner/inference.py +215 -0
- mlxtuner/train.py +232 -0
- mlxtuner-0.1.0.dist-info/METADATA +255 -0
- mlxtuner-0.1.0.dist-info/RECORD +13 -0
- mlxtuner-0.1.0.dist-info/WHEEL +4 -0
- mlxtuner-0.1.0.dist-info/entry_points.txt +2 -0
- mlxtuner-0.1.0.dist-info/licenses/LICENSE +202 -0
mlxtuner/__init__.py
ADDED
mlxtuner/cli.py
ADDED
|
@@ -0,0 +1,384 @@
|
|
|
1
|
+
"""mlxtuner command line interface."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import json
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
from typing import Annotated
|
|
8
|
+
|
|
9
|
+
import typer
|
|
10
|
+
from rich.console import Console
|
|
11
|
+
from rich.table import Table
|
|
12
|
+
|
|
13
|
+
from . import __version__
|
|
14
|
+
from .config import EXAMPLE_CONFIG, RunConfig
|
|
15
|
+
|
|
16
|
+
app = typer.Typer(
|
|
17
|
+
name="mlxtuner",
|
|
18
|
+
help="Fine-tune LLMs on Apple Silicon with one command.",
|
|
19
|
+
no_args_is_help=True,
|
|
20
|
+
rich_markup_mode="rich",
|
|
21
|
+
pretty_exceptions_show_locals=False,
|
|
22
|
+
)
|
|
23
|
+
console = Console()
|
|
24
|
+
|
|
25
|
+
SetOpt = Annotated[
|
|
26
|
+
list[str] | None,
|
|
27
|
+
typer.Option(
|
|
28
|
+
"--set", "-s", help="Override a config value, e.g. --set train.lr=1e-4", show_default=False
|
|
29
|
+
),
|
|
30
|
+
]
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def _version(value: bool) -> None:
|
|
34
|
+
if value:
|
|
35
|
+
console.print(f"mlxtuner {__version__}")
|
|
36
|
+
raise typer.Exit()
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
@app.callback()
|
|
40
|
+
def _main(
|
|
41
|
+
version: Annotated[
|
|
42
|
+
bool, typer.Option("--version", "-V", callback=_version, is_eager=True, help="Show version")
|
|
43
|
+
] = False,
|
|
44
|
+
) -> None:
|
|
45
|
+
pass
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def _load_config(
|
|
49
|
+
config: Path | None,
|
|
50
|
+
model: str | None,
|
|
51
|
+
data: str | None,
|
|
52
|
+
output: str | None,
|
|
53
|
+
overrides: list[str] | None,
|
|
54
|
+
) -> RunConfig:
|
|
55
|
+
overrides = list(overrides or [])
|
|
56
|
+
if model:
|
|
57
|
+
overrides.append(f"model={model}")
|
|
58
|
+
if data:
|
|
59
|
+
overrides.append(f"data.path={data}")
|
|
60
|
+
if output:
|
|
61
|
+
overrides.append(f"train.output={output}")
|
|
62
|
+
if config:
|
|
63
|
+
return RunConfig.from_yaml(config, overrides)
|
|
64
|
+
if not (model and data):
|
|
65
|
+
raise typer.BadParameter("Pass a config file, or both --model and --data.")
|
|
66
|
+
return RunConfig.from_dict({"model": model, "data": {"path": data}}, overrides)
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
# ---------------------------------------------------------------------------
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
@app.command()
|
|
73
|
+
def check() -> None:
|
|
74
|
+
"""Show this Mac's chip and RAM, and the training defaults mlxtuner will use for it."""
|
|
75
|
+
from .hardware import TIER_DEFAULTS, detect
|
|
76
|
+
|
|
77
|
+
m = detect()
|
|
78
|
+
console.print(f"mlxtuner {__version__}")
|
|
79
|
+
console.print(f"chip: {m.chip}")
|
|
80
|
+
console.print(f"ram: {m.ram_gb:g} GB -> tier {m.tier}")
|
|
81
|
+
if not m.apple_silicon:
|
|
82
|
+
console.print("[red]This is not an Apple Silicon Mac; MLX will not run here.[/]")
|
|
83
|
+
raise typer.Exit(1)
|
|
84
|
+
d = TIER_DEFAULTS[m.tier]
|
|
85
|
+
console.print(
|
|
86
|
+
f"defaults: batch_size={d.batch_size} max_seq_length={d.max_seq_length} "
|
|
87
|
+
f"num_layers={d.num_layers} grad_checkpoint={'on' if d.grad_checkpoint else 'off'}"
|
|
88
|
+
)
|
|
89
|
+
console.print("see what fits: mlxtuner models")
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
@app.command()
|
|
93
|
+
def models(
|
|
94
|
+
all_: Annotated[
|
|
95
|
+
bool, typer.Option("--all", help="Include models that don't fit this Mac")
|
|
96
|
+
] = False,
|
|
97
|
+
) -> None:
|
|
98
|
+
"""List recommended 4-bit models and whether each fits this Mac for LoRA training."""
|
|
99
|
+
from .hardware import MODELS, TIER_DEFAULTS, detect, estimate_train_gb, fits
|
|
100
|
+
|
|
101
|
+
m = detect()
|
|
102
|
+
d = TIER_DEFAULTS[m.tier]
|
|
103
|
+
table = Table(
|
|
104
|
+
title=f"LoRA training on {m.chip}, {m.ram_gb:g} GB\n"
|
|
105
|
+
f"(batch {d.batch_size}, seq {d.max_seq_length}, {d.num_layers} layers, grad checkpoint {'on' if d.grad_checkpoint else 'off'})",
|
|
106
|
+
caption="all repos are under mlx-community/",
|
|
107
|
+
)
|
|
108
|
+
table.add_column("model", no_wrap=True)
|
|
109
|
+
table.add_column("params", justify="right", no_wrap=True)
|
|
110
|
+
table.add_column("weights", justify="right", no_wrap=True)
|
|
111
|
+
table.add_column("est. peak", justify="right", no_wrap=True)
|
|
112
|
+
table.add_column("fits", no_wrap=True, min_width=5)
|
|
113
|
+
colour = {"yes": "green", "tight": "yellow", "no": "red"}
|
|
114
|
+
for rec in MODELS:
|
|
115
|
+
est = estimate_train_gb(
|
|
116
|
+
rec.weights_gb,
|
|
117
|
+
d.batch_size,
|
|
118
|
+
d.max_seq_length,
|
|
119
|
+
d.num_layers,
|
|
120
|
+
d.grad_checkpoint,
|
|
121
|
+
rec.vocab_k,
|
|
122
|
+
)
|
|
123
|
+
verdict = fits(m, est)
|
|
124
|
+
if verdict == "no" and not all_:
|
|
125
|
+
continue
|
|
126
|
+
name = rec.repo.removeprefix("mlx-community/")
|
|
127
|
+
if rec.note:
|
|
128
|
+
name += f"\n[dim]{rec.note}[/]"
|
|
129
|
+
table.add_row(
|
|
130
|
+
name, f"{rec.params_b:g}B", f"{rec.weights_gb:.1f} GB",
|
|
131
|
+
f"~{est:g} GB", f"[{colour[verdict]}]{verdict}[/]",
|
|
132
|
+
) # fmt: skip
|
|
133
|
+
console.print(table)
|
|
134
|
+
if not all_:
|
|
135
|
+
console.print(
|
|
136
|
+
"[dim]--all shows models that don't fit. Estimates are ±30 %; peak memory is reported after each run.[/]"
|
|
137
|
+
)
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
@app.command()
|
|
141
|
+
def init(
|
|
142
|
+
path: Annotated[Path, typer.Argument(help="Where to write the starter config")] = Path(
|
|
143
|
+
"config.yaml"
|
|
144
|
+
),
|
|
145
|
+
force: Annotated[bool, typer.Option("--force", "-f")] = False,
|
|
146
|
+
) -> None:
|
|
147
|
+
"""Write a commented starter config."""
|
|
148
|
+
if path.exists() and not force:
|
|
149
|
+
console.print(f"[red]{path} exists[/] (use --force to overwrite)")
|
|
150
|
+
raise typer.Exit(1)
|
|
151
|
+
path.write_text(EXAMPLE_CONFIG)
|
|
152
|
+
console.print(f"[green]wrote {path}[/] next: edit it, then mlxtuner train {path}")
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
@app.command()
|
|
156
|
+
def validate(
|
|
157
|
+
config: Annotated[
|
|
158
|
+
Path | None, typer.Argument(help="Run config (optional if --data is given)")
|
|
159
|
+
] = None,
|
|
160
|
+
model: Annotated[
|
|
161
|
+
str | None, typer.Option("--model", "-m", help="Model id (for tokenizer stats)")
|
|
162
|
+
] = None,
|
|
163
|
+
data: Annotated[str | None, typer.Option("--data", "-d", help="Dataset path or Hub id")] = None,
|
|
164
|
+
overrides: SetOpt = None,
|
|
165
|
+
show: Annotated[int, typer.Option(help="Print this many rendered examples")] = 1,
|
|
166
|
+
) -> None:
|
|
167
|
+
"""Convert a dataset without training: format detection, drop reasons, token-length stats."""
|
|
168
|
+
import tempfile
|
|
169
|
+
|
|
170
|
+
from .data import DataError, prepare, render, token_stats
|
|
171
|
+
from .hardware import detect
|
|
172
|
+
|
|
173
|
+
model = model or "mlx-community/Qwen2.5-0.5B-Instruct-4bit"
|
|
174
|
+
cfg = _load_config(config, model, data, None, overrides).resolved(detect())
|
|
175
|
+
with tempfile.TemporaryDirectory() as tmp:
|
|
176
|
+
try:
|
|
177
|
+
p = prepare(cfg.data, Path(tmp) / "data")
|
|
178
|
+
except DataError as e:
|
|
179
|
+
console.print(f"[red]data error:[/] {e}")
|
|
180
|
+
raise typer.Exit(1) from None
|
|
181
|
+
console.print(f"format: [bold]{p.source_format}[/] -> {p.kind}")
|
|
182
|
+
console.print(f"rows: {p.n_train:,} train / {p.n_valid:,} valid / {p.dropped} dropped")
|
|
183
|
+
for reason, n in sorted(p.drop_reasons.items(), key=lambda kv: -kv[1])[:5]:
|
|
184
|
+
console.print(f" [yellow]{n:>6}[/] {reason}")
|
|
185
|
+
|
|
186
|
+
from .inference import load_tokenizer_only
|
|
187
|
+
|
|
188
|
+
tok = load_tokenizer_only(cfg.model)
|
|
189
|
+
if p.kind == "messages" and getattr(tok, "chat_template", None) is None:
|
|
190
|
+
console.print(f"[red]{cfg.model} has no chat template; use an -Instruct model[/]")
|
|
191
|
+
raise typer.Exit(1)
|
|
192
|
+
with open(p.dir / "train.jsonl") as f:
|
|
193
|
+
rows = [json.loads(line) for line in f]
|
|
194
|
+
st = token_stats(tok, rows, int(cfg.train.max_seq_length))
|
|
195
|
+
table = Table(title=f"token lengths (sampled {st['sampled']:,})")
|
|
196
|
+
for k in ("min", "p50", "p90", "max", "mean"):
|
|
197
|
+
table.add_column(k, justify="right")
|
|
198
|
+
table.add_row(*(str(st[k]) for k in ("min", "p50", "p90", "max", "mean")))
|
|
199
|
+
console.print(table)
|
|
200
|
+
if st["over_max"]:
|
|
201
|
+
console.print(
|
|
202
|
+
f"[yellow]{st['over_max']} of {st['sampled']} rows exceed max_seq_length={cfg.train.max_seq_length} "
|
|
203
|
+
"and will be truncated.[/]"
|
|
204
|
+
)
|
|
205
|
+
for i in range(min(show, len(rows))):
|
|
206
|
+
console.rule(f"[dim]example {i}")
|
|
207
|
+
console.print(render(tok, rows[i]), markup=False, highlight=False)
|
|
208
|
+
console.print("[green]dataset OK[/]")
|
|
209
|
+
|
|
210
|
+
|
|
211
|
+
@app.command()
|
|
212
|
+
def train(
|
|
213
|
+
config: Annotated[
|
|
214
|
+
Path | None, typer.Argument(help="YAML run config (see `mlxtuner init`)")
|
|
215
|
+
] = None,
|
|
216
|
+
model: Annotated[
|
|
217
|
+
str | None, typer.Option("--model", "-m", help="Model id or local path")
|
|
218
|
+
] = None,
|
|
219
|
+
data: Annotated[str | None, typer.Option("--data", "-d", help="Dataset path or Hub id")] = None,
|
|
220
|
+
output: Annotated[
|
|
221
|
+
str | None, typer.Option("--output", "-o", help="Adapter output directory")
|
|
222
|
+
] = None,
|
|
223
|
+
overrides: SetOpt = None,
|
|
224
|
+
dry_run: Annotated[
|
|
225
|
+
bool, typer.Option("--dry-run", help="Convert data and print the plan, train nothing")
|
|
226
|
+
] = False,
|
|
227
|
+
) -> None:
|
|
228
|
+
"""Fine-tune a model with LoRA. Use a config file, or --model + --data for auto defaults."""
|
|
229
|
+
from .train import run
|
|
230
|
+
|
|
231
|
+
run(_load_config(config, model, data, output, overrides), dry_run=dry_run)
|
|
232
|
+
|
|
233
|
+
|
|
234
|
+
@app.command()
|
|
235
|
+
def chat(
|
|
236
|
+
path: Annotated[
|
|
237
|
+
str, typer.Argument(help="Adapter dir from `mlxtuner train`, fused model dir, or Hub id")
|
|
238
|
+
],
|
|
239
|
+
system: Annotated[str | None, typer.Option("--system", help="System prompt")] = None,
|
|
240
|
+
max_tokens: Annotated[int, typer.Option(help="Max tokens per reply")] = 512,
|
|
241
|
+
temperature: Annotated[float, typer.Option(help="0 = greedy")] = 0.7,
|
|
242
|
+
) -> None:
|
|
243
|
+
"""Chat interactively with a fine-tuned model."""
|
|
244
|
+
from .inference import chat_loop
|
|
245
|
+
|
|
246
|
+
chat_loop(path, system=system, max_tokens=max_tokens, temperature=temperature)
|
|
247
|
+
|
|
248
|
+
|
|
249
|
+
@app.command()
|
|
250
|
+
def fuse(
|
|
251
|
+
adapter: Annotated[str, typer.Argument(help="Adapter dir from `mlxtuner train`")],
|
|
252
|
+
output: Annotated[
|
|
253
|
+
str, typer.Option("--output", "-o", help="Where to save the fused model")
|
|
254
|
+
] = "fused-model",
|
|
255
|
+
dequantize: Annotated[
|
|
256
|
+
bool, typer.Option("--dequantize", help="Save in full precision (needed for GGUF)")
|
|
257
|
+
] = False,
|
|
258
|
+
gguf: Annotated[
|
|
259
|
+
str | None,
|
|
260
|
+
typer.Option(
|
|
261
|
+
"--gguf", help="Also export a GGUF file to this path (llama/mistral/mixtral only)"
|
|
262
|
+
),
|
|
263
|
+
] = None,
|
|
264
|
+
ollama: Annotated[
|
|
265
|
+
bool, typer.Option("--ollama", help="Write an Ollama Modelfile next to the GGUF")
|
|
266
|
+
] = False,
|
|
267
|
+
system: Annotated[str | None, typer.Option(help="System prompt for the Modelfile")] = None,
|
|
268
|
+
) -> None:
|
|
269
|
+
"""Merge the adapter into the base model -> standalone MLX model (and optionally GGUF / Ollama)."""
|
|
270
|
+
from .inference import fuse as _fuse
|
|
271
|
+
from .inference import write_ollama_modelfile
|
|
272
|
+
|
|
273
|
+
if gguf and not dequantize:
|
|
274
|
+
console.print("[dim]--gguf implies --dequantize[/]")
|
|
275
|
+
dequantize = True
|
|
276
|
+
_fuse(adapter, output, dequantize=dequantize, gguf=gguf)
|
|
277
|
+
if ollama:
|
|
278
|
+
if not gguf:
|
|
279
|
+
console.print("[red]--ollama needs --gguf <path>[/]")
|
|
280
|
+
raise typer.Exit(1)
|
|
281
|
+
write_ollama_modelfile(gguf, system=system)
|
|
282
|
+
|
|
283
|
+
|
|
284
|
+
@app.command()
|
|
285
|
+
def eval( # noqa: A001 - typer command name
|
|
286
|
+
path: Annotated[
|
|
287
|
+
str, typer.Argument(help="Adapter dir from `mlxtuner train`, fused dir, or Hub id")
|
|
288
|
+
],
|
|
289
|
+
data: Annotated[
|
|
290
|
+
str | None,
|
|
291
|
+
typer.Option(
|
|
292
|
+
"--data", "-d", help="Dataset path or Hub id (default: the run's mlxtuner.yaml data)"
|
|
293
|
+
),
|
|
294
|
+
] = None,
|
|
295
|
+
config: Annotated[
|
|
296
|
+
Path | None, typer.Option("--config", "-c", help="Run config to take data settings from")
|
|
297
|
+
] = None,
|
|
298
|
+
overrides: SetOpt = None,
|
|
299
|
+
compare_base: Annotated[
|
|
300
|
+
bool,
|
|
301
|
+
typer.Option("--compare-base/--no-compare-base", help="Also score the untuned base model"),
|
|
302
|
+
] = True,
|
|
303
|
+
samples: Annotated[int, typer.Option(help="Sample generations to show")] = 3,
|
|
304
|
+
max_examples: Annotated[int, typer.Option(help="Cap on examples scored")] = 100,
|
|
305
|
+
max_tokens: Annotated[int, typer.Option(help="Tokens per sample generation")] = 128,
|
|
306
|
+
output: Annotated[
|
|
307
|
+
Path | None,
|
|
308
|
+
typer.Option("--output", "-o", help="Where to write eval.json (default: <path>/eval.json)"),
|
|
309
|
+
] = None,
|
|
310
|
+
) -> None:
|
|
311
|
+
"""Held-out loss / perplexity and sample generations, tuned vs base."""
|
|
312
|
+
from .data import DataError
|
|
313
|
+
from .eval import print_result, run_eval, save_result
|
|
314
|
+
from .hardware import detect
|
|
315
|
+
|
|
316
|
+
run_yaml = Path(path) / "mlxtuner.yaml"
|
|
317
|
+
if config is None and data is None and run_yaml.exists():
|
|
318
|
+
config = run_yaml
|
|
319
|
+
if config is None and data is None:
|
|
320
|
+
raise typer.BadParameter(
|
|
321
|
+
"Pass --data, or --config, or a run directory containing mlxtuner.yaml."
|
|
322
|
+
)
|
|
323
|
+
cfg = _load_config(
|
|
324
|
+
config, "unused" if config is None else None, data, None, overrides
|
|
325
|
+
).resolved(detect())
|
|
326
|
+
try:
|
|
327
|
+
result = run_eval(
|
|
328
|
+
path,
|
|
329
|
+
cfg.data,
|
|
330
|
+
max_seq_length=int(cfg.train.max_seq_length),
|
|
331
|
+
mask_prompt=cfg.train.mask_prompt,
|
|
332
|
+
compare_base=compare_base,
|
|
333
|
+
n_samples=samples,
|
|
334
|
+
max_examples=max_examples,
|
|
335
|
+
max_tokens=max_tokens,
|
|
336
|
+
)
|
|
337
|
+
except DataError as e:
|
|
338
|
+
console.print(f"[red]data error:[/] {e}")
|
|
339
|
+
raise typer.Exit(1) from None
|
|
340
|
+
print_result(result)
|
|
341
|
+
out = output or (Path(path) / "eval.json" if Path(path).is_dir() else Path("eval.json"))
|
|
342
|
+
save_result(result, out)
|
|
343
|
+
console.print(f"[dim]written to {out}[/]")
|
|
344
|
+
|
|
345
|
+
|
|
346
|
+
@app.command()
|
|
347
|
+
def export(
|
|
348
|
+
model_dir: Annotated[
|
|
349
|
+
str, typer.Argument(help="Fused model dir from `mlxtuner fuse --dequantize`")
|
|
350
|
+
],
|
|
351
|
+
output: Annotated[str | None, typer.Option("--output", "-o", help="Output .gguf path")] = None,
|
|
352
|
+
quant: Annotated[str, typer.Option(help="f16 | bf16 | q8_0 | f32")] = "q8_0",
|
|
353
|
+
llama_cpp: Annotated[str | None, typer.Option(help="Path to a llama.cpp checkout")] = None,
|
|
354
|
+
ollama: Annotated[
|
|
355
|
+
bool, typer.Option("--ollama", help="Also write an Ollama Modelfile")
|
|
356
|
+
] = False,
|
|
357
|
+
system: Annotated[str | None, typer.Option(help="System prompt for the Modelfile")] = None,
|
|
358
|
+
) -> None:
|
|
359
|
+
"""Export a dequantized fused model to GGUF via llama.cpp (any architecture llama.cpp supports)."""
|
|
360
|
+
from .inference import to_gguf, write_ollama_modelfile
|
|
361
|
+
|
|
362
|
+
try:
|
|
363
|
+
gguf = to_gguf(model_dir, output=output, quant=quant, llama_cpp=llama_cpp)
|
|
364
|
+
except (FileNotFoundError, RuntimeError) as e:
|
|
365
|
+
console.print(f"[red]{e}[/]")
|
|
366
|
+
raise typer.Exit(1) from None
|
|
367
|
+
if ollama:
|
|
368
|
+
write_ollama_modelfile(str(gguf), system=system)
|
|
369
|
+
|
|
370
|
+
|
|
371
|
+
@app.command()
|
|
372
|
+
def info(path: Annotated[str, typer.Argument(help="A run output directory")]) -> None:
|
|
373
|
+
"""Summarise a finished run (mlxtuner.json)."""
|
|
374
|
+
meta_path = Path(path) / "mlxtuner.json"
|
|
375
|
+
if not meta_path.exists():
|
|
376
|
+
console.print(f"[red]{meta_path} not found[/]")
|
|
377
|
+
raise typer.Exit(1)
|
|
378
|
+
meta = json.loads(meta_path.read_text())
|
|
379
|
+
meta.pop("history", None)
|
|
380
|
+
console.print_json(json.dumps(meta))
|
|
381
|
+
|
|
382
|
+
|
|
383
|
+
if __name__ == "__main__":
|
|
384
|
+
app()
|
mlxtuner/config.py
ADDED
|
@@ -0,0 +1,166 @@
|
|
|
1
|
+
"""Run configuration. 'auto' values are resolved from the detected Mac at train time."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import copy
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
from typing import Any, Literal
|
|
8
|
+
|
|
9
|
+
import yaml
|
|
10
|
+
from pydantic import BaseModel, Field, field_validator
|
|
11
|
+
|
|
12
|
+
from .hardware import TIER_DEFAULTS, Machine
|
|
13
|
+
|
|
14
|
+
Auto = Literal["auto"]
|
|
15
|
+
DataFormat = Literal["auto", "messages", "alpaca", "sharegpt", "prompt_completion", "text"]
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class DataConfig(BaseModel):
|
|
19
|
+
path: str = Field(
|
|
20
|
+
description="A .jsonl/.json/.csv file, a directory with train/valid.jsonl, or a Hub dataset id"
|
|
21
|
+
)
|
|
22
|
+
format: DataFormat = "auto"
|
|
23
|
+
eval_fraction: float = Field(0.05, ge=0.0, lt=1.0, description="Held out as the validation set")
|
|
24
|
+
max_samples: int | None = None
|
|
25
|
+
system_prompt: str | None = Field(
|
|
26
|
+
None, description="Added to conversations that lack a system turn"
|
|
27
|
+
)
|
|
28
|
+
shuffle_seed: int | None = 42
|
|
29
|
+
text_field: str = "text"
|
|
30
|
+
prompt_field: str = "prompt"
|
|
31
|
+
completion_field: str = "completion"
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
class LoraConfig(BaseModel):
|
|
35
|
+
type: Literal["lora", "dora", "full"] = "lora"
|
|
36
|
+
rank: int = Field(8, ge=1)
|
|
37
|
+
scale: float = Field(
|
|
38
|
+
20.0, gt=0, description="MLX uses a direct scale (not alpha/r). 20 is mlx-lm's default."
|
|
39
|
+
)
|
|
40
|
+
dropout: float = Field(0.0, ge=0.0, le=1.0)
|
|
41
|
+
num_layers: int | Auto = Field(
|
|
42
|
+
"auto", description="How many transformer layers (from the top) get adapters. -1 = all"
|
|
43
|
+
)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class TrainConfig(BaseModel):
|
|
47
|
+
output: str = "adapters/run"
|
|
48
|
+
epochs: float = Field(1.0, gt=0, description="Used to compute iters when iters is null")
|
|
49
|
+
iters: int | None = Field(None, description="Total optimizer steps; overrides epochs")
|
|
50
|
+
batch_size: int | Auto = "auto"
|
|
51
|
+
max_seq_length: int | Auto = "auto"
|
|
52
|
+
grad_checkpoint: bool | Auto = "auto"
|
|
53
|
+
grad_accumulation_steps: int = Field(1, ge=1)
|
|
54
|
+
lr: float = Field(
|
|
55
|
+
1e-5, gt=0, description="mlx-lm's LoRA default; try 1e-4 for small datasets / bigger shifts"
|
|
56
|
+
)
|
|
57
|
+
optimizer: Literal["adam", "adamw", "sgd", "adafactor", "muon"] = "adam"
|
|
58
|
+
warmup_steps: int = Field(0, ge=0, description="Linear warmup then cosine decay when > 0")
|
|
59
|
+
mask_prompt: bool = Field(
|
|
60
|
+
True, description="Loss only on assistant/completion tokens (ignored for text data)"
|
|
61
|
+
)
|
|
62
|
+
steps_per_report: int = 10
|
|
63
|
+
steps_per_eval: int | None = Field(None, description="null = 4 evals per run")
|
|
64
|
+
val_batches: int = Field(25, description="Validation batches per eval; -1 = whole set")
|
|
65
|
+
save_every: int | None = Field(None, description="null = every eval")
|
|
66
|
+
seed: int = 0
|
|
67
|
+
resume: str | None = Field(None, description="Path to adapters.safetensors to continue from")
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
class RunConfig(BaseModel):
|
|
71
|
+
model: str = Field(
|
|
72
|
+
description="mlx-community/* repo id, any HF model id (converted on load), or local path"
|
|
73
|
+
)
|
|
74
|
+
data: DataConfig
|
|
75
|
+
lora: LoraConfig = Field(default_factory=LoraConfig)
|
|
76
|
+
train: TrainConfig = Field(default_factory=TrainConfig)
|
|
77
|
+
|
|
78
|
+
@field_validator("data", mode="before")
|
|
79
|
+
@classmethod
|
|
80
|
+
def _data_str(cls, v: Any) -> Any:
|
|
81
|
+
return {"path": v} if isinstance(v, str) else v
|
|
82
|
+
|
|
83
|
+
@classmethod
|
|
84
|
+
def from_yaml(cls, path: str | Path, overrides: list[str] | None = None) -> RunConfig:
|
|
85
|
+
with open(path) as f:
|
|
86
|
+
raw = yaml.safe_load(f) or {}
|
|
87
|
+
return cls.from_dict(raw, overrides)
|
|
88
|
+
|
|
89
|
+
@classmethod
|
|
90
|
+
def from_dict(cls, raw: dict[str, Any], overrides: list[str] | None = None) -> RunConfig:
|
|
91
|
+
if overrides:
|
|
92
|
+
raw = apply_overrides(raw, overrides)
|
|
93
|
+
return cls.model_validate(raw)
|
|
94
|
+
|
|
95
|
+
def to_yaml(self, path: str | Path) -> None:
|
|
96
|
+
Path(path).parent.mkdir(parents=True, exist_ok=True)
|
|
97
|
+
with open(path, "w") as f:
|
|
98
|
+
yaml.safe_dump(self.model_dump(mode="json"), f, sort_keys=False)
|
|
99
|
+
|
|
100
|
+
# -- resolve 'auto' -----------------------------------------------------
|
|
101
|
+
def resolved(self, machine: Machine) -> RunConfig:
|
|
102
|
+
"""Return a copy with every 'auto' replaced by the value for this Mac's RAM tier."""
|
|
103
|
+
d = TIER_DEFAULTS[machine.tier]
|
|
104
|
+
out = self.model_copy(deep=True)
|
|
105
|
+
if out.train.batch_size == "auto":
|
|
106
|
+
out.train.batch_size = d.batch_size
|
|
107
|
+
if out.train.max_seq_length == "auto":
|
|
108
|
+
out.train.max_seq_length = d.max_seq_length
|
|
109
|
+
if out.train.grad_checkpoint == "auto":
|
|
110
|
+
out.train.grad_checkpoint = d.grad_checkpoint
|
|
111
|
+
if out.lora.num_layers == "auto":
|
|
112
|
+
out.lora.num_layers = d.num_layers
|
|
113
|
+
return out
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
def _parse_scalar(value: str) -> Any:
|
|
117
|
+
try:
|
|
118
|
+
return yaml.safe_load(value)
|
|
119
|
+
except yaml.YAMLError:
|
|
120
|
+
return value
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
def apply_overrides(raw: dict[str, Any], overrides: list[str]) -> dict[str, Any]:
|
|
124
|
+
"""Apply ``["train.lr=1e-4", "lora.rank=16"]`` onto a nested dict (returns a copy)."""
|
|
125
|
+
out = copy.deepcopy(raw)
|
|
126
|
+
for item in overrides:
|
|
127
|
+
if "=" not in item:
|
|
128
|
+
raise ValueError(f"Override must look like section.key=value, got: {item!r}")
|
|
129
|
+
key, _, value = item.partition("=")
|
|
130
|
+
parts = key.strip().split(".")
|
|
131
|
+
node = out
|
|
132
|
+
for p in parts[:-1]:
|
|
133
|
+
node = node.setdefault(p, {})
|
|
134
|
+
if not isinstance(node, dict):
|
|
135
|
+
raise ValueError(f"Cannot set {key!r}: {p!r} is not a mapping")
|
|
136
|
+
node[parts[-1]] = _parse_scalar(value.strip())
|
|
137
|
+
return out
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
EXAMPLE_CONFIG = """\
|
|
141
|
+
# mlxtuner run config. 'auto' values are chosen from your Mac's RAM when training starts.
|
|
142
|
+
# Override anything from the CLI: mlxtuner train config.yaml --set train.lr=1e-4
|
|
143
|
+
|
|
144
|
+
model: mlx-community/Qwen2.5-1.5B-Instruct-4bit # see `mlxtuner models` for what fits your Mac
|
|
145
|
+
|
|
146
|
+
data:
|
|
147
|
+
path: data/train.jsonl # .jsonl/.json/.csv, a dir with train/valid.jsonl, or a Hub dataset id
|
|
148
|
+
format: auto # auto | messages | alpaca | sharegpt | prompt_completion | text
|
|
149
|
+
eval_fraction: 0.05
|
|
150
|
+
# system_prompt: "You are a helpful assistant."
|
|
151
|
+
|
|
152
|
+
lora:
|
|
153
|
+
type: lora # lora | dora | full
|
|
154
|
+
rank: 8
|
|
155
|
+
scale: 20.0
|
|
156
|
+
num_layers: auto # layers that get adapters (8 on 8 GB Macs, 16 above)
|
|
157
|
+
|
|
158
|
+
train:
|
|
159
|
+
output: adapters/my-run
|
|
160
|
+
epochs: 2
|
|
161
|
+
batch_size: auto
|
|
162
|
+
max_seq_length: auto
|
|
163
|
+
grad_checkpoint: auto
|
|
164
|
+
lr: 1.0e-5
|
|
165
|
+
mask_prompt: true # learn only from assistant turns
|
|
166
|
+
"""
|