tunekit 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.
- tunekit/__init__.py +3 -0
- tunekit/cli.py +363 -0
- tunekit/config.py +250 -0
- tunekit/data.py +371 -0
- tunekit/eval.py +196 -0
- tunekit/export.py +128 -0
- tunekit/inference.py +113 -0
- tunekit/model.py +228 -0
- tunekit/train.py +235 -0
- tunekit-0.1.0.dist-info/METADATA +293 -0
- tunekit-0.1.0.dist-info/RECORD +14 -0
- tunekit-0.1.0.dist-info/WHEEL +4 -0
- tunekit-0.1.0.dist-info/entry_points.txt +2 -0
- tunekit-0.1.0.dist-info/licenses/LICENSE +202 -0
tunekit/__init__.py
ADDED
tunekit/cli.py
ADDED
|
@@ -0,0 +1,363 @@
|
|
|
1
|
+
"""tunekit 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="tunekit",
|
|
18
|
+
help="One-command LoRA/QLoRA fine-tuning for open LLMs.",
|
|
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"tunekit {__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.name={model}")
|
|
58
|
+
if data:
|
|
59
|
+
overrides.append(f"data.path={data}")
|
|
60
|
+
if output:
|
|
61
|
+
overrides.append(f"train.output_dir={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": {"name": model}, "data": {"path": data}}, overrides)
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
# ---------------------------------------------------------------------------
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
@app.command()
|
|
73
|
+
def init(
|
|
74
|
+
path: Annotated[Path, typer.Argument(help="Where to write the starter config")] = Path(
|
|
75
|
+
"config.yaml"
|
|
76
|
+
),
|
|
77
|
+
force: Annotated[bool, typer.Option("--force", "-f", help="Overwrite if it exists")] = False,
|
|
78
|
+
) -> None:
|
|
79
|
+
"""Write a commented starter config you can edit."""
|
|
80
|
+
if path.exists() and not force:
|
|
81
|
+
console.print(f"[red]{path} exists[/] (use --force to overwrite)")
|
|
82
|
+
raise typer.Exit(1)
|
|
83
|
+
path.write_text(EXAMPLE_CONFIG)
|
|
84
|
+
console.print(f"[green]wrote {path}[/] next: edit it, then tunekit train {path}")
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
@app.command()
|
|
88
|
+
def validate(
|
|
89
|
+
config: Annotated[
|
|
90
|
+
Path | None, typer.Argument(help="Run config (optional if --data is given)")
|
|
91
|
+
] = None,
|
|
92
|
+
model: Annotated[
|
|
93
|
+
str | None, typer.Option("--model", "-m", help="Model id (for tokenizer stats)")
|
|
94
|
+
] = None,
|
|
95
|
+
data: Annotated[str | None, typer.Option("--data", "-d", help="Dataset path or Hub id")] = None,
|
|
96
|
+
overrides: SetOpt = None,
|
|
97
|
+
show: Annotated[int, typer.Option(help="Print this many rendered examples")] = 1,
|
|
98
|
+
) -> None:
|
|
99
|
+
"""Check a dataset: detect its format, convert it, and report token-length stats."""
|
|
100
|
+
from .data import DataError, prepare, render_example, token_stats
|
|
101
|
+
from .model import load_tokenizer
|
|
102
|
+
|
|
103
|
+
model = model or "Qwen/Qwen2.5-0.5B-Instruct"
|
|
104
|
+
cfg = _load_config(config, model, data, None, overrides)
|
|
105
|
+
try:
|
|
106
|
+
nd = prepare(cfg.data)
|
|
107
|
+
except DataError as e:
|
|
108
|
+
console.print(f"[red]data error:[/] {e}")
|
|
109
|
+
raise typer.Exit(1) from None
|
|
110
|
+
|
|
111
|
+
console.print(f"format: [bold]{nd.source_format}[/] -> {nd.kind}")
|
|
112
|
+
n_eval = len(nd.eval) if nd.eval is not None else 0
|
|
113
|
+
console.print(f"rows: {len(nd.train):,} train / {n_eval:,} eval / {nd.dropped} dropped")
|
|
114
|
+
for reason, n in sorted(nd.drop_reasons.items(), key=lambda kv: -kv[1])[:5]:
|
|
115
|
+
console.print(f" [yellow]{n:>6}[/] {reason}")
|
|
116
|
+
|
|
117
|
+
tok = load_tokenizer(cfg.model)
|
|
118
|
+
if nd.kind == "messages" and tok.chat_template is None:
|
|
119
|
+
console.print(f"[red]{cfg.model.name} has no chat template; use an instruct/chat model[/]")
|
|
120
|
+
raise typer.Exit(1)
|
|
121
|
+
st = token_stats(tok, nd.train, nd.kind, cfg.data.max_length)
|
|
122
|
+
table = Table(
|
|
123
|
+
title=f"token lengths (tokenizer: {cfg.model.name}, sampled {st['sampled']:,}/{st['total']:,})"
|
|
124
|
+
)
|
|
125
|
+
for k in ("min", "p50", "p90", "p99", "max", "mean"):
|
|
126
|
+
table.add_column(k, justify="right")
|
|
127
|
+
table.add_row(*(str(st[k]) for k in ("min", "p50", "p90", "p99", "max", "mean")))
|
|
128
|
+
console.print(table)
|
|
129
|
+
console.print(f"~{st['total_tokens_est']:,} tokens per epoch")
|
|
130
|
+
if st["over_max_length"]:
|
|
131
|
+
pct = 100 * st["over_max_length"] / st["sampled"]
|
|
132
|
+
console.print(
|
|
133
|
+
f"[yellow]{st['over_max_length']} of {st['sampled']} sampled rows ({pct:.1f}%) exceed "
|
|
134
|
+
f"max_length={cfg.data.max_length} and will be truncated.[/] Raise data.max_length or shorten them."
|
|
135
|
+
)
|
|
136
|
+
for i in range(min(show, len(nd.train))):
|
|
137
|
+
console.rule(f"[dim]example {i}")
|
|
138
|
+
console.print(render_example(tok, nd.train[i], nd.kind), markup=False, highlight=False)
|
|
139
|
+
console.print("[green]dataset OK[/]")
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
@app.command()
|
|
143
|
+
def train(
|
|
144
|
+
config: Annotated[
|
|
145
|
+
Path | None, typer.Argument(help="YAML run config (see `tunekit init`)")
|
|
146
|
+
] = None,
|
|
147
|
+
model: Annotated[
|
|
148
|
+
str | None, typer.Option("--model", "-m", help="Base model id or path")
|
|
149
|
+
] = None,
|
|
150
|
+
data: Annotated[str | None, typer.Option("--data", "-d", help="Dataset path or Hub id")] = None,
|
|
151
|
+
output: Annotated[str | None, typer.Option("--output", "-o", help="Output directory")] = None,
|
|
152
|
+
overrides: SetOpt = None,
|
|
153
|
+
dry_run: Annotated[
|
|
154
|
+
bool, typer.Option("--dry-run", help="Load everything, train nothing")
|
|
155
|
+
] = False,
|
|
156
|
+
) -> None:
|
|
157
|
+
"""Fine-tune a model. Use a config file, or --model + --data for defaults."""
|
|
158
|
+
from .data import DataError
|
|
159
|
+
from .train import run
|
|
160
|
+
|
|
161
|
+
cfg = _load_config(config, model, data, output, overrides)
|
|
162
|
+
try:
|
|
163
|
+
run(cfg, dry_run=dry_run)
|
|
164
|
+
except DataError as e:
|
|
165
|
+
console.print(f"[red]data error:[/] {e}")
|
|
166
|
+
raise typer.Exit(1) from None
|
|
167
|
+
|
|
168
|
+
|
|
169
|
+
@app.command()
|
|
170
|
+
def chat(
|
|
171
|
+
path: Annotated[str, typer.Argument(help="Adapter dir, merged model dir, or Hub id")],
|
|
172
|
+
system: Annotated[str | None, typer.Option("--system", help="System prompt")] = None,
|
|
173
|
+
max_new_tokens: Annotated[int, typer.Option(help="Max tokens per reply")] = 512,
|
|
174
|
+
temperature: Annotated[float, typer.Option(help="0 = greedy")] = 0.7,
|
|
175
|
+
) -> None:
|
|
176
|
+
"""Chat interactively with a fine-tuned model."""
|
|
177
|
+
from .inference import chat_loop
|
|
178
|
+
|
|
179
|
+
chat_loop(path, system=system, max_new_tokens=max_new_tokens, temperature=temperature)
|
|
180
|
+
|
|
181
|
+
|
|
182
|
+
@app.command()
|
|
183
|
+
def merge(
|
|
184
|
+
adapter: Annotated[str, typer.Argument(help="Adapter directory from `tunekit train`")],
|
|
185
|
+
output: Annotated[str, typer.Argument(help="Where to write the merged model")],
|
|
186
|
+
base: Annotated[
|
|
187
|
+
str | None, typer.Option(help="Override base model (default: from adapter_config.json)")
|
|
188
|
+
] = None,
|
|
189
|
+
dtype: Annotated[str, typer.Option(help="bfloat16 | float16 | float32")] = "bfloat16",
|
|
190
|
+
) -> None:
|
|
191
|
+
"""Merge a LoRA adapter into its base model -> standalone Hugging Face model."""
|
|
192
|
+
from .export import merge as _merge
|
|
193
|
+
|
|
194
|
+
_merge(adapter, output, base=base, dtype=dtype)
|
|
195
|
+
|
|
196
|
+
|
|
197
|
+
@app.command()
|
|
198
|
+
def export(
|
|
199
|
+
model_dir: Annotated[
|
|
200
|
+
str, typer.Argument(help="Merged model directory (run `tunekit merge` first)")
|
|
201
|
+
],
|
|
202
|
+
output: Annotated[str | None, typer.Option("--output", "-o", help="Output .gguf path")] = None,
|
|
203
|
+
quant: Annotated[str, typer.Option(help="f16 | bf16 | q8_0 | f32")] = "q8_0",
|
|
204
|
+
llama_cpp: Annotated[str | None, typer.Option(help="Path to a llama.cpp checkout")] = None,
|
|
205
|
+
ollama: Annotated[
|
|
206
|
+
bool, typer.Option("--ollama", help="Also write an Ollama Modelfile")
|
|
207
|
+
] = False,
|
|
208
|
+
system: Annotated[
|
|
209
|
+
str | None, typer.Option(help="System prompt to bake into the Modelfile")
|
|
210
|
+
] = None,
|
|
211
|
+
) -> None:
|
|
212
|
+
"""Export a merged model to GGUF (llama.cpp / Ollama / LM Studio)."""
|
|
213
|
+
from .export import to_gguf, write_ollama_modelfile
|
|
214
|
+
|
|
215
|
+
try:
|
|
216
|
+
gguf = to_gguf(model_dir, output=output, quant=quant, llama_cpp=llama_cpp)
|
|
217
|
+
except FileNotFoundError as e:
|
|
218
|
+
console.print(f"[red]{e}[/]")
|
|
219
|
+
raise typer.Exit(1) from None
|
|
220
|
+
if ollama:
|
|
221
|
+
write_ollama_modelfile(str(gguf), model_dir, system=system)
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
@app.command()
|
|
225
|
+
def eval( # noqa: A001 - typer command name
|
|
226
|
+
path: Annotated[
|
|
227
|
+
str, typer.Argument(help="Adapter dir from `tunekit train`, merged dir, or Hub id")
|
|
228
|
+
],
|
|
229
|
+
data: Annotated[
|
|
230
|
+
str | None,
|
|
231
|
+
typer.Option(
|
|
232
|
+
"--data", "-d", help="Dataset path or Hub id (default: the run's tunekit.yaml data)"
|
|
233
|
+
),
|
|
234
|
+
] = None,
|
|
235
|
+
config: Annotated[
|
|
236
|
+
Path | None, typer.Option("--config", "-c", help="Run config to take data settings from")
|
|
237
|
+
] = None,
|
|
238
|
+
overrides: SetOpt = None,
|
|
239
|
+
compare_base: Annotated[
|
|
240
|
+
bool,
|
|
241
|
+
typer.Option("--compare-base/--no-compare-base", help="Also score the untuned base model"),
|
|
242
|
+
] = True,
|
|
243
|
+
samples: Annotated[int, typer.Option(help="Sample generations to show")] = 3,
|
|
244
|
+
max_examples: Annotated[int, typer.Option(help="Cap on examples scored")] = 100,
|
|
245
|
+
max_new_tokens: Annotated[int, typer.Option(help="Tokens per sample generation")] = 128,
|
|
246
|
+
output: Annotated[
|
|
247
|
+
Path | None,
|
|
248
|
+
typer.Option("--output", "-o", help="Where to write eval.json (default: <path>/eval.json)"),
|
|
249
|
+
] = None,
|
|
250
|
+
) -> None:
|
|
251
|
+
"""Held-out loss / perplexity and sample generations, tuned vs base."""
|
|
252
|
+
from .data import DataError
|
|
253
|
+
from .eval import print_result, run_eval, save_result
|
|
254
|
+
|
|
255
|
+
run_yaml = Path(path) / "tunekit.yaml"
|
|
256
|
+
if config is None and data is None and run_yaml.exists():
|
|
257
|
+
config = run_yaml
|
|
258
|
+
if config is None and data is None:
|
|
259
|
+
raise typer.BadParameter(
|
|
260
|
+
"Pass --data, or --config, or a run directory containing tunekit.yaml."
|
|
261
|
+
)
|
|
262
|
+
cfg = _load_config(config, "unused" if config is None else None, data, None, overrides)
|
|
263
|
+
try:
|
|
264
|
+
result = run_eval(
|
|
265
|
+
path,
|
|
266
|
+
cfg.data,
|
|
267
|
+
compare_base=compare_base,
|
|
268
|
+
n_samples=samples,
|
|
269
|
+
max_examples=max_examples,
|
|
270
|
+
max_new_tokens=max_new_tokens,
|
|
271
|
+
)
|
|
272
|
+
except DataError as e:
|
|
273
|
+
console.print(f"[red]data error:[/] {e}")
|
|
274
|
+
raise typer.Exit(1) from None
|
|
275
|
+
print_result(result)
|
|
276
|
+
out = output or (Path(path) / "eval.json" if Path(path).is_dir() else Path("eval.json"))
|
|
277
|
+
save_result(result, out)
|
|
278
|
+
console.print(f"[dim]written to {out}[/]")
|
|
279
|
+
|
|
280
|
+
|
|
281
|
+
@app.command()
|
|
282
|
+
def push(
|
|
283
|
+
path: Annotated[str, typer.Argument(help="Adapter or merged model directory")],
|
|
284
|
+
repo_id: Annotated[str, typer.Argument(help="Hub repo, e.g. your-username/my-finetune")],
|
|
285
|
+
private: Annotated[bool, typer.Option("--private/--public")] = True,
|
|
286
|
+
message: Annotated[str, typer.Option("--message", "-m")] = "Upload from tunekit",
|
|
287
|
+
) -> None:
|
|
288
|
+
"""Upload a run directory to the Hugging Face Hub (needs `huggingface-cli login` or HF_TOKEN)."""
|
|
289
|
+
from huggingface_hub import HfApi
|
|
290
|
+
from huggingface_hub.errors import HfHubHTTPError
|
|
291
|
+
|
|
292
|
+
src = Path(path)
|
|
293
|
+
if not src.is_dir():
|
|
294
|
+
console.print(f"[red]{src} is not a directory[/]")
|
|
295
|
+
raise typer.Exit(1)
|
|
296
|
+
api = HfApi()
|
|
297
|
+
try:
|
|
298
|
+
api.create_repo(repo_id, private=private, exist_ok=True)
|
|
299
|
+
url = api.upload_folder(
|
|
300
|
+
folder_path=str(src),
|
|
301
|
+
repo_id=repo_id,
|
|
302
|
+
commit_message=message,
|
|
303
|
+
ignore_patterns=["checkpoint-*", "runs/*", "wandb/*"],
|
|
304
|
+
)
|
|
305
|
+
except HfHubHTTPError as e:
|
|
306
|
+
console.print(f"[red]hub error:[/] {e}")
|
|
307
|
+
console.print("Log in with: huggingface-cli login (or set HF_TOKEN)")
|
|
308
|
+
raise typer.Exit(1) from None
|
|
309
|
+
console.print(f"[green]uploaded[/] {src} -> https://huggingface.co/{repo_id} [dim]{url}[/]")
|
|
310
|
+
|
|
311
|
+
|
|
312
|
+
@app.command()
|
|
313
|
+
def info(
|
|
314
|
+
path: Annotated[str | None, typer.Argument(help="A run output directory to summarise")] = None,
|
|
315
|
+
model: Annotated[
|
|
316
|
+
str | None,
|
|
317
|
+
typer.Option("--model", "-m", help="Check whether a model's weights fit this GPU"),
|
|
318
|
+
] = None,
|
|
319
|
+
) -> None:
|
|
320
|
+
"""Show detected hardware, whether a model fits it, or summarise a finished run."""
|
|
321
|
+
from .model import bitsandbytes_available, detect_hardware, hub_weight_bytes
|
|
322
|
+
|
|
323
|
+
hw = detect_hardware()
|
|
324
|
+
console.print(f"tunekit {__version__}")
|
|
325
|
+
console.print(f"hardware: {hw.summary}")
|
|
326
|
+
console.print(
|
|
327
|
+
f"bitsandbytes (QLoRA): {'available' if bitsandbytes_available() else 'not installed'}"
|
|
328
|
+
)
|
|
329
|
+
if model:
|
|
330
|
+
size = hub_weight_bytes(model)
|
|
331
|
+
if size is None:
|
|
332
|
+
console.print(
|
|
333
|
+
f"[yellow]could not read weight sizes for {model} (offline, gated, or not a model repo)[/]"
|
|
334
|
+
)
|
|
335
|
+
else:
|
|
336
|
+
gb = size / 1e9
|
|
337
|
+
console.print(f"{model}: {gb:.1f} GB of weights on disk")
|
|
338
|
+
console.print(
|
|
339
|
+
f" resident base weights: bf16 ~{gb:.1f} GB | 8-bit ~{gb * 0.55:.1f} GB | 4-bit ~{gb * 0.30:.1f} GB "
|
|
340
|
+
"[dim](LoRA params, optimizer state and activations come on top)[/]"
|
|
341
|
+
)
|
|
342
|
+
if hw.vram_gb:
|
|
343
|
+
for label, ratio in (("bf16", 1.0), ("8-bit", 0.55), ("4-bit", 0.30)):
|
|
344
|
+
need = gb * ratio
|
|
345
|
+
ok = (
|
|
346
|
+
"fits"
|
|
347
|
+
if need < hw.vram_gb * 0.8
|
|
348
|
+
else "tight"
|
|
349
|
+
if need < hw.vram_gb
|
|
350
|
+
else "no"
|
|
351
|
+
)
|
|
352
|
+
console.print(f" {label:>5}: {ok}")
|
|
353
|
+
if path:
|
|
354
|
+
meta_path = Path(path) / "tunekit.json"
|
|
355
|
+
if not meta_path.exists():
|
|
356
|
+
console.print(f"[red]{meta_path} not found[/]")
|
|
357
|
+
raise typer.Exit(1)
|
|
358
|
+
meta = json.loads(meta_path.read_text())
|
|
359
|
+
console.print_json(json.dumps(meta))
|
|
360
|
+
|
|
361
|
+
|
|
362
|
+
if __name__ == "__main__":
|
|
363
|
+
app()
|
tunekit/config.py
ADDED
|
@@ -0,0 +1,250 @@
|
|
|
1
|
+
"""Typed configuration for a fine-tuning run.
|
|
2
|
+
|
|
3
|
+
A run is described by a single YAML file (see ``configs/``). Every field has a
|
|
4
|
+
sensible default so a minimal config only needs ``model.name`` and ``data.path``.
|
|
5
|
+
Any field can be overridden from the CLI with ``--set section.key=value``.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
from pathlib import Path
|
|
11
|
+
from typing import Any, Literal
|
|
12
|
+
|
|
13
|
+
import yaml
|
|
14
|
+
from pydantic import BaseModel, Field, field_validator
|
|
15
|
+
|
|
16
|
+
DataFormat = Literal["auto", "messages", "alpaca", "sharegpt", "prompt_completion", "text"]
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class ModelConfig(BaseModel):
|
|
20
|
+
name: str = Field(
|
|
21
|
+
description="Hugging Face model id or local path, e.g. Qwen/Qwen2.5-0.5B-Instruct"
|
|
22
|
+
)
|
|
23
|
+
load_in_4bit: bool = Field(
|
|
24
|
+
False, description="QLoRA: load base model in 4-bit NF4 (needs bitsandbytes + CUDA)"
|
|
25
|
+
)
|
|
26
|
+
load_in_8bit: bool = Field(
|
|
27
|
+
False, description="Load base model in 8-bit (needs bitsandbytes + CUDA)"
|
|
28
|
+
)
|
|
29
|
+
dtype: Literal["auto", "bfloat16", "float16", "float32"] = Field(
|
|
30
|
+
"auto",
|
|
31
|
+
description="Compute dtype. 'auto' picks bf16 on supported GPUs, fp16 on older GPUs, fp32 on CPU.",
|
|
32
|
+
)
|
|
33
|
+
attn_implementation: str | None = Field(
|
|
34
|
+
None, description="e.g. 'flash_attention_2' or 'sdpa'. None lets transformers choose."
|
|
35
|
+
)
|
|
36
|
+
trust_remote_code: bool = False
|
|
37
|
+
chat_template: str | None = Field(
|
|
38
|
+
None, description="Override the tokenizer chat template (Jinja string). Rarely needed."
|
|
39
|
+
)
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class DataConfig(BaseModel):
|
|
43
|
+
path: str = Field(
|
|
44
|
+
description="Local file (.jsonl/.json/.csv/.parquet), directory, or Hub dataset id"
|
|
45
|
+
)
|
|
46
|
+
format: DataFormat = Field(
|
|
47
|
+
"auto", description="Input schema. 'auto' detects from the first row."
|
|
48
|
+
)
|
|
49
|
+
split: str = Field("train", description="Split to use when loading from the Hub")
|
|
50
|
+
subset: str | None = Field(None, description="Hub dataset config/subset name")
|
|
51
|
+
eval_path: str | None = Field(None, description="Optional separate eval file / dataset id")
|
|
52
|
+
eval_fraction: float = Field(
|
|
53
|
+
0.02,
|
|
54
|
+
ge=0.0,
|
|
55
|
+
lt=1.0,
|
|
56
|
+
description="Hold out this fraction of train for eval if eval_path is not set",
|
|
57
|
+
)
|
|
58
|
+
max_samples: int | None = Field(
|
|
59
|
+
None, description="Truncate the training set (useful for smoke tests)"
|
|
60
|
+
)
|
|
61
|
+
system_prompt: str | None = Field(
|
|
62
|
+
None, description="Prepend a system message to every conversation that lacks one"
|
|
63
|
+
)
|
|
64
|
+
max_length: int = Field(
|
|
65
|
+
2048, ge=64, description="Max tokens per example; longer examples are truncated"
|
|
66
|
+
)
|
|
67
|
+
shuffle_seed: int | None = Field(
|
|
68
|
+
42, description="Seed used to shuffle before splitting. None disables shuffling."
|
|
69
|
+
)
|
|
70
|
+
# Column names, only needed when auto-detection can't find them.
|
|
71
|
+
text_field: str = "text"
|
|
72
|
+
prompt_field: str = "prompt"
|
|
73
|
+
completion_field: str = "completion"
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
class LoraConfigModel(BaseModel):
|
|
77
|
+
enabled: bool = Field(
|
|
78
|
+
True, description="Set false for full fine-tuning (needs a lot more VRAM)"
|
|
79
|
+
)
|
|
80
|
+
r: int = Field(16, ge=1, description="LoRA rank")
|
|
81
|
+
alpha: int = Field(32, ge=1, description="LoRA alpha (scaling = alpha / r)")
|
|
82
|
+
dropout: float = Field(0.05, ge=0.0, le=1.0)
|
|
83
|
+
target_modules: str | list[str] = Field(
|
|
84
|
+
"all-linear", description="'all-linear' or a list like [q_proj, k_proj, v_proj, o_proj]"
|
|
85
|
+
)
|
|
86
|
+
use_rslora: bool = False
|
|
87
|
+
use_dora: bool = False
|
|
88
|
+
modules_to_save: list[str] | None = Field(
|
|
89
|
+
None,
|
|
90
|
+
description="Extra modules to train fully, e.g. [embed_tokens, lm_head] when adding tokens",
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
class TrainConfig(BaseModel):
|
|
95
|
+
output_dir: str = "outputs/run"
|
|
96
|
+
epochs: float = Field(1.0, gt=0)
|
|
97
|
+
max_steps: int = Field(-1, description="Overrides epochs when > 0")
|
|
98
|
+
batch_size: int = Field(2, ge=1, description="Per-device micro batch size")
|
|
99
|
+
grad_accum: int = Field(
|
|
100
|
+
8,
|
|
101
|
+
ge=1,
|
|
102
|
+
description="Gradient accumulation steps. Effective batch = batch_size * grad_accum * n_gpus",
|
|
103
|
+
)
|
|
104
|
+
lr: float = Field(2e-4, gt=0)
|
|
105
|
+
scheduler: str = "cosine"
|
|
106
|
+
warmup_ratio: float = Field(0.03, ge=0.0, le=1.0)
|
|
107
|
+
weight_decay: float = 0.0
|
|
108
|
+
max_grad_norm: float = 1.0
|
|
109
|
+
optimizer: str = Field(
|
|
110
|
+
"adamw_torch", description="e.g. adamw_torch, adamw_8bit, paged_adamw_8bit, adafactor"
|
|
111
|
+
)
|
|
112
|
+
gradient_checkpointing: bool = True
|
|
113
|
+
packing: bool = Field(
|
|
114
|
+
False,
|
|
115
|
+
description="Pack multiple short examples into one sequence. Faster, but can hurt chat models.",
|
|
116
|
+
)
|
|
117
|
+
assistant_only_loss: bool = Field(
|
|
118
|
+
False,
|
|
119
|
+
description="Only compute loss on assistant turns. Needs a chat template with {% generation %} markers.",
|
|
120
|
+
)
|
|
121
|
+
logging_steps: int = 10
|
|
122
|
+
eval_steps: int | None = Field(
|
|
123
|
+
None, description="Evaluate every N steps. None = once per epoch."
|
|
124
|
+
)
|
|
125
|
+
save_steps: int | None = Field(
|
|
126
|
+
None, description="Checkpoint every N steps. None = once per epoch."
|
|
127
|
+
)
|
|
128
|
+
save_total_limit: int = 2
|
|
129
|
+
seed: int = 42
|
|
130
|
+
report_to: list[str] = Field(
|
|
131
|
+
default_factory=lambda: ["none"], description="e.g. [wandb], [tensorboard]"
|
|
132
|
+
)
|
|
133
|
+
run_name: str | None = None
|
|
134
|
+
resume_from_checkpoint: str | bool = Field(
|
|
135
|
+
False, description="Path to a checkpoint dir, or true to pick the latest"
|
|
136
|
+
)
|
|
137
|
+
dataloader_num_workers: int = 0
|
|
138
|
+
extra: dict[str, Any] = Field(
|
|
139
|
+
default_factory=dict, description="Passed straight to SFTConfig (escape hatch)"
|
|
140
|
+
)
|
|
141
|
+
|
|
142
|
+
@field_validator("report_to", mode="before")
|
|
143
|
+
@classmethod
|
|
144
|
+
def _coerce_report_to(cls, v: Any) -> list[str]:
|
|
145
|
+
if isinstance(v, str):
|
|
146
|
+
return [v]
|
|
147
|
+
return v
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
class HubConfig(BaseModel):
|
|
151
|
+
push: bool = False
|
|
152
|
+
repo_id: str | None = Field(None, description="e.g. your-username/my-finetune")
|
|
153
|
+
private: bool = True
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
class RunConfig(BaseModel):
|
|
157
|
+
model: ModelConfig
|
|
158
|
+
data: DataConfig
|
|
159
|
+
lora: LoraConfigModel = Field(default_factory=LoraConfigModel)
|
|
160
|
+
train: TrainConfig = Field(default_factory=TrainConfig)
|
|
161
|
+
hub: HubConfig = Field(default_factory=HubConfig)
|
|
162
|
+
|
|
163
|
+
# ----------------------------------------------------------------- I/O
|
|
164
|
+
@classmethod
|
|
165
|
+
def from_yaml(cls, path: str | Path, overrides: list[str] | None = None) -> RunConfig:
|
|
166
|
+
with open(path) as f:
|
|
167
|
+
raw = yaml.safe_load(f) or {}
|
|
168
|
+
if overrides:
|
|
169
|
+
raw = apply_overrides(raw, overrides)
|
|
170
|
+
return cls.model_validate(raw)
|
|
171
|
+
|
|
172
|
+
@classmethod
|
|
173
|
+
def from_dict(cls, raw: dict[str, Any], overrides: list[str] | None = None) -> RunConfig:
|
|
174
|
+
if overrides:
|
|
175
|
+
raw = apply_overrides(raw, overrides)
|
|
176
|
+
return cls.model_validate(raw)
|
|
177
|
+
|
|
178
|
+
def to_yaml(self, path: str | Path) -> None:
|
|
179
|
+
Path(path).parent.mkdir(parents=True, exist_ok=True)
|
|
180
|
+
with open(path, "w") as f:
|
|
181
|
+
yaml.safe_dump(self.model_dump(mode="json"), f, sort_keys=False)
|
|
182
|
+
|
|
183
|
+
|
|
184
|
+
# ---------------------------------------------------------------------------
|
|
185
|
+
# --set a.b.c=value overrides
|
|
186
|
+
# ---------------------------------------------------------------------------
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
def _parse_scalar(value: str) -> Any:
|
|
190
|
+
"""Parse a CLI string into the most specific YAML scalar (int/float/bool/null/list)."""
|
|
191
|
+
try:
|
|
192
|
+
return yaml.safe_load(value)
|
|
193
|
+
except yaml.YAMLError:
|
|
194
|
+
return value
|
|
195
|
+
|
|
196
|
+
|
|
197
|
+
def apply_overrides(raw: dict[str, Any], overrides: list[str]) -> dict[str, Any]:
|
|
198
|
+
"""Apply ``["train.lr=1e-4", "model.name=foo"]`` onto a nested dict (returns a copy)."""
|
|
199
|
+
import copy
|
|
200
|
+
|
|
201
|
+
out = copy.deepcopy(raw)
|
|
202
|
+
for item in overrides:
|
|
203
|
+
if "=" not in item:
|
|
204
|
+
raise ValueError(f"Override must look like section.key=value, got: {item!r}")
|
|
205
|
+
key, _, value = item.partition("=")
|
|
206
|
+
parts = key.strip().split(".")
|
|
207
|
+
node = out
|
|
208
|
+
for p in parts[:-1]:
|
|
209
|
+
node = node.setdefault(p, {})
|
|
210
|
+
if not isinstance(node, dict):
|
|
211
|
+
raise ValueError(f"Cannot set {key!r}: {p!r} is not a mapping")
|
|
212
|
+
node[parts[-1]] = _parse_scalar(value.strip())
|
|
213
|
+
return out
|
|
214
|
+
|
|
215
|
+
|
|
216
|
+
EXAMPLE_CONFIG = """\
|
|
217
|
+
# tunekit run config. Every field is optional except model.name and data.path.
|
|
218
|
+
# Override anything from the CLI: tunekit train config.yaml --set train.lr=1e-4
|
|
219
|
+
|
|
220
|
+
model:
|
|
221
|
+
name: Qwen/Qwen2.5-0.5B-Instruct # any causal LM on the Hub, or a local path
|
|
222
|
+
load_in_4bit: false # true = QLoRA (needs bitsandbytes + NVIDIA GPU)
|
|
223
|
+
|
|
224
|
+
data:
|
|
225
|
+
path: data/train.jsonl # .jsonl/.json/.csv/.parquet, a directory, or a Hub dataset id
|
|
226
|
+
format: auto # auto | messages | alpaca | sharegpt | prompt_completion | text
|
|
227
|
+
eval_fraction: 0.02 # hold-out for eval loss
|
|
228
|
+
max_length: 2048 # tokens per example
|
|
229
|
+
# system_prompt: "You are a helpful assistant."
|
|
230
|
+
|
|
231
|
+
lora:
|
|
232
|
+
r: 16
|
|
233
|
+
alpha: 32
|
|
234
|
+
dropout: 0.05
|
|
235
|
+
target_modules: all-linear
|
|
236
|
+
|
|
237
|
+
train:
|
|
238
|
+
output_dir: outputs/my-run
|
|
239
|
+
epochs: 1
|
|
240
|
+
batch_size: 2
|
|
241
|
+
grad_accum: 8 # effective batch = 16
|
|
242
|
+
lr: 2.0e-4
|
|
243
|
+
gradient_checkpointing: true
|
|
244
|
+
logging_steps: 10
|
|
245
|
+
report_to: [none] # or [wandb], [tensorboard]
|
|
246
|
+
|
|
247
|
+
hub:
|
|
248
|
+
push: false
|
|
249
|
+
# repo_id: your-username/my-finetune
|
|
250
|
+
"""
|