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 ADDED
@@ -0,0 +1,3 @@
1
+ """mlxtuner – fine-tune LLMs on Apple Silicon with one command."""
2
+
3
+ __version__ = "0.1.0"
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
+ """