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 ADDED
@@ -0,0 +1,3 @@
1
+ """tunekit – one-command LoRA/QLoRA fine-tuning for open LLMs."""
2
+
3
+ __version__ = "0.1.0"
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
+ """