troy-cli 0.2.2__tar.gz → 0.4.0__tar.gz

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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: troy-cli
3
- Version: 0.2.2
3
+ Version: 0.4.0
4
4
  Summary: Fine-tune LLMs on your MacBook with one YAML file. Built for Apple Silicon.
5
5
  Author: Troy
6
6
  License: Apache-2.0
@@ -11,6 +11,8 @@ Requires-Dist: pydantic>=2.5
11
11
  Requires-Dist: pyyaml>=6.0
12
12
  Requires-Dist: rich>=13.0
13
13
  Requires-Dist: typer>=0.12
14
+ Provides-Extra: vision
15
+ Requires-Dist: mlx-vlm[train]>=0.7; extra == 'vision'
14
16
  Description-Content-Type: text/markdown
15
17
 
16
18
  # Troy
@@ -79,7 +81,7 @@ output: ./output
79
81
  |---|---|
80
82
  | `troy init` | Create a config from a template (`chat`, `dpo`, `orpo`) |
81
83
  | `troy doctor` | Hardware + dependency check, with model-size guidance |
82
- | `troy train` | LoRA fine-tuning: SFT, DPO, or ORPO |
84
+ | `troy train` | LoRA fine-tuning: SFT, DPO, or ORPO — text, or vision with `[vision]` extra |
83
85
  | `troy chat` | Interactive REPL (or `-p` for one-shot) with your adapter |
84
86
  | `troy eval` | Base-vs-tuned val loss, perplexity, side-by-side samples |
85
87
  | `troy serve` | OpenAI-compatible API server for your model |
@@ -64,7 +64,7 @@ output: ./output
64
64
  |---|---|
65
65
  | `troy init` | Create a config from a template (`chat`, `dpo`, `orpo`) |
66
66
  | `troy doctor` | Hardware + dependency check, with model-size guidance |
67
- | `troy train` | LoRA fine-tuning: SFT, DPO, or ORPO |
67
+ | `troy train` | LoRA fine-tuning: SFT, DPO, or ORPO — text, or vision with `[vision]` extra |
68
68
  | `troy chat` | Interactive REPL (or `-p` for one-shot) with your adapter |
69
69
  | `troy eval` | Base-vs-tuned val loss, perplexity, side-by-side samples |
70
70
  | `troy serve` | OpenAI-compatible API server for your model |
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "troy-cli"
3
- version = "0.2.2"
3
+ version = "0.4.0"
4
4
  description = "Fine-tune LLMs on your MacBook with one YAML file. Built for Apple Silicon."
5
5
  readme = "README.md"
6
6
  requires-python = ">=3.10"
@@ -15,6 +15,9 @@ dependencies = [
15
15
  "rich>=13.0",
16
16
  ]
17
17
 
18
+ [project.optional-dependencies]
19
+ vision = ["mlx-vlm[train]>=0.7"]
20
+
18
21
  [project.scripts]
19
22
  troy = "troy.cli:app"
20
23
 
@@ -1,3 +1,3 @@
1
1
  """Troy: fine-tune LLMs on your MacBook with one YAML file."""
2
2
 
3
- __version__ = "0.2.2"
3
+ __version__ = "0.4.0"
@@ -142,6 +142,19 @@ def train(
142
142
  from .data import load_and_prepare
143
143
 
144
144
  cfg = load_config(config)
145
+
146
+ if Path(cfg.data.train).expanduser().is_dir(): # folder of images => vision
147
+ if cfg.task != "sft":
148
+ console.print("[red]Vision fine-tuning supports task: sft (for now).[/red]")
149
+ raise typer.Exit(1)
150
+ from .train_vision import run_vision_sft
151
+
152
+ run_vision_sft(cfg)
153
+ console.print(
154
+ '\nTry it: [bold]troy chat --image photo.png -p "your question"[/bold]'
155
+ )
156
+ return
157
+
145
158
  train_records, valid_records, fmt = load_and_prepare(
146
159
  cfg.data, cfg.task, cfg.training.seed
147
160
  )
@@ -179,6 +192,7 @@ def chat(
179
192
  max_tokens: int = typer.Option(512),
180
193
  temperature: float = typer.Option(0.7),
181
194
  prompt: Optional[str] = typer.Option(None, "--prompt", "-p", help="One-shot prompt (no REPL)."),
195
+ image: Optional[Path] = typer.Option(None, help="Image for a vision model (one-shot; needs -p)."),
182
196
  ) -> None:
183
197
  """Chat with your fine-tuned model."""
184
198
  _require_apple_silicon()
@@ -198,6 +212,14 @@ def chat(
198
212
  console.print(
199
213
  "[yellow]No trained adapter found — chatting with the base model.[/yellow]"
200
214
  )
215
+ if image is not None:
216
+ if prompt is None:
217
+ console.print("[red]--image needs a one-shot prompt: -p \"your question\"[/red]")
218
+ raise typer.Exit(1)
219
+ from .train_vision import run_vision_chat
220
+
221
+ run_vision_chat(model, adapter, str(image), prompt, max_tokens, temperature)
222
+ return
201
223
  run_chat(model, adapter, max_tokens, temperature, prompt)
202
224
 
203
225
 
@@ -297,15 +319,19 @@ def push(
297
319
  run_push(folder, repo, private=not public)
298
320
 
299
321
 
300
- @app.command()
301
- def data(
302
- action: str = typer.Argument(help="Action: inspect"),
322
+ data_app = typer.Typer(
323
+ name="data",
324
+ help="Inspect, validate, and synthesize datasets.",
325
+ no_args_is_help=True,
326
+ )
327
+ app.add_typer(data_app)
328
+
329
+
330
+ @data_app.command()
331
+ def inspect(
303
332
  path: Path = typer.Argument(help="Dataset file (.jsonl, .json, .csv)."),
304
333
  ) -> None:
305
334
  """Inspect a dataset: record count, detected format, sizes."""
306
- if action != "inspect":
307
- console.print("[red]Only `troy data inspect <path>` is supported.[/red]")
308
- raise typer.Exit(1)
309
335
  from .data import inspect_stats
310
336
 
311
337
  stats = inspect_stats(str(path))
@@ -319,5 +345,92 @@ def data(
319
345
  console.print(table)
320
346
 
321
347
 
348
+ @data_app.command()
349
+ def validate(
350
+ path: Path = typer.Argument(help="Dataset file (.jsonl, .json, .csv)."),
351
+ ) -> None:
352
+ """Lint a dataset: broken records, mixed formats, empty fields, duplicates."""
353
+ from .data import validate_records
354
+
355
+ report = validate_records(str(path))
356
+ console.print(
357
+ f"{report['records']} records, format: [bold]{report['format']}[/bold]"
358
+ )
359
+ if not report["issues"]:
360
+ console.print("[green]No issues found.[/green]")
361
+ return
362
+ for issue in report["issues"]:
363
+ console.print(f" [yellow]•[/yellow] {issue}")
364
+ if report["truncated"]:
365
+ console.print(" [dim]... more issues not shown[/dim]")
366
+ console.print(f"[red]{len(report['issues'])}{'+' if report['truncated'] else ''} issue(s).[/red]")
367
+ raise typer.Exit(1)
368
+
369
+
370
+ @data_app.command()
371
+ def synth(
372
+ source: Optional[Path] = typer.Option(
373
+ None, "--from", help="Ground examples in a file or folder of docs/code."
374
+ ),
375
+ seed: Optional[str] = typer.Option(
376
+ None, "--seed", help='Task description, e.g. "customer support bot for Acme".'
377
+ ),
378
+ n: int = typer.Option(100, "--n", help="Number of examples to generate."),
379
+ fmt: str = typer.Option(
380
+ "chat", "--format", "-f",
381
+ help="Output format: chat (SFT) or preference (DPO/ORPO).",
382
+ ),
383
+ teacher: str = typer.Option(
384
+ "auto", help="Teacher model (auto = sized to this Mac's memory)."
385
+ ),
386
+ out: Optional[Path] = typer.Option(
387
+ None, "--out", "-o",
388
+ help="Output file (default: data/train.jsonl or data/preferences.jsonl).",
389
+ ),
390
+ max_tokens: int = typer.Option(2048, help="Max tokens per teacher call."),
391
+ temperature: float = typer.Option(0.8),
392
+ ) -> None:
393
+ """Synthesize a training dataset with a local teacher model."""
394
+ _require_apple_silicon()
395
+ if source is None and seed is None:
396
+ console.print(
397
+ '[red]Give the teacher something to work from:[/red] '
398
+ '--from ./docs and/or --seed "task description".'
399
+ )
400
+ raise typer.Exit(1)
401
+ if fmt not in ("chat", "preference"):
402
+ console.print("[red]--format must be `chat` or `preference`.[/red]")
403
+ raise typer.Exit(1)
404
+
405
+ from .synth import pick_teacher, run_synth
406
+
407
+ if teacher == "auto":
408
+ teacher = pick_teacher()
409
+ console.print(f"Teacher: [bold]{teacher}[/bold] (picked for this Mac's memory)")
410
+ out = out or Path("data") / ("preferences.jsonl" if fmt == "preference" else "train.jsonl")
411
+ if out.exists():
412
+ console.print(f"[red]{out} already exists[/red] — pass -o to write elsewhere.")
413
+ raise typer.Exit(1)
414
+
415
+ stats = run_synth(
416
+ n=n, out_path=out, fmt=fmt, teacher=teacher,
417
+ seed_task=seed, source=source,
418
+ max_tokens=max_tokens, temperature=temperature,
419
+ )
420
+ console.print(
421
+ f"\nWrote [bold]{stats['records']}[/bold] examples to [bold]{stats['out']}[/bold] "
422
+ f"({stats['teacher_calls']} teacher calls)"
423
+ )
424
+ if stats["records"] < stats["requested"]:
425
+ console.print(
426
+ f"[yellow]Stopped at {stats['records']}/{stats['requested']} — "
427
+ "try a larger --teacher, higher --max-tokens, or more source material.[/yellow]"
428
+ )
429
+ console.print(
430
+ "Review the data before training — spot-check a dozen examples, then: "
431
+ "[bold]troy data validate " + str(out) + "[/bold] and [bold]troy train[/bold]."
432
+ )
433
+
434
+
322
435
  if __name__ == "__main__":
323
436
  app()
@@ -152,6 +152,120 @@ def load_and_prepare(
152
152
  return train, valid, fmt
153
153
 
154
154
 
155
+ # Required non-empty string fields per format (nested fields checked separately).
156
+ _REQUIRED_FIELDS = {
157
+ "alpaca": ("instruction", "output"),
158
+ "completions": ("prompt", "completion"),
159
+ "preference": ("prompt", "chosen", "rejected"),
160
+ "text": ("text",),
161
+ }
162
+
163
+
164
+ def _record_issues(record: Dict[str, Any], fmt: str, line: int) -> List[str]:
165
+ issues = []
166
+
167
+ def empty(v: Any) -> bool:
168
+ return not (isinstance(v, str) and v.strip())
169
+
170
+ for field in _REQUIRED_FIELDS.get(fmt, ()):
171
+ if empty(record.get(field)):
172
+ issues.append(f"line {line}: empty or missing `{field}`")
173
+ if fmt == "chat":
174
+ msgs = record.get("messages")
175
+ if not isinstance(msgs, list) or not msgs:
176
+ issues.append(f"line {line}: `messages` is not a non-empty list")
177
+ else:
178
+ roles = [m.get("role") for m in msgs if isinstance(m, dict)]
179
+ if "assistant" not in roles:
180
+ issues.append(f"line {line}: no assistant message")
181
+ if any(
182
+ not isinstance(m, dict) or empty(m.get("content"))
183
+ for m in msgs
184
+ ):
185
+ issues.append(f"line {line}: message with empty content")
186
+ if fmt == "sharegpt":
187
+ convs = record.get("conversations")
188
+ if not isinstance(convs, list) or not convs:
189
+ issues.append(f"line {line}: `conversations` is not a non-empty list")
190
+ else:
191
+ bad = [
192
+ m.get("from")
193
+ for m in convs
194
+ if not isinstance(m, dict)
195
+ or str(m.get("from", "")).lower() not in _SHAREGPT_ROLES
196
+ ]
197
+ if bad:
198
+ issues.append(f"line {line}: unknown speaker role(s): {bad}")
199
+ if fmt == "preference" and not issues:
200
+ if record["chosen"].strip() == record["rejected"].strip():
201
+ issues.append(f"line {line}: chosen == rejected")
202
+ return issues
203
+
204
+
205
+ def validate_records(path: str, max_issues: int = 50) -> Dict[str, Any]:
206
+ """Lint a dataset: per-record issues, mixed formats, duplicates.
207
+
208
+ Returns {"records", "format", "issues", "duplicates", "truncated"}.
209
+ """
210
+ p = Path(path).expanduser()
211
+ issues: List[str] = []
212
+ records: List[Tuple[int, Dict[str, Any]]] = []
213
+
214
+ if p.suffix.lower() == ".jsonl":
215
+ with open(p) as f:
216
+ for i, line in enumerate(f, 1):
217
+ if not line.strip():
218
+ continue
219
+ try:
220
+ records.append((i, json.loads(line)))
221
+ except json.JSONDecodeError as e:
222
+ issues.append(f"line {i}: invalid JSON ({e.msg})")
223
+ else:
224
+ records = [(i, r) for i, r in enumerate(_read_records(p), 1)]
225
+
226
+ if not records:
227
+ return {
228
+ "records": 0, "format": "unknown", "issues": issues or ["file has no records"],
229
+ "duplicates": 0, "truncated": False,
230
+ }
231
+
232
+ try:
233
+ fmt = detect_format(records[0][1])
234
+ except ValueError as e:
235
+ return {
236
+ "records": len(records), "format": "unknown",
237
+ "issues": issues + [str(e)], "duplicates": 0, "truncated": False,
238
+ }
239
+
240
+ seen: Dict[str, int] = {}
241
+ duplicates = 0
242
+ for line, record in records:
243
+ try:
244
+ rec_fmt = detect_format(record)
245
+ except ValueError:
246
+ issues.append(f"line {line}: keys match no known format")
247
+ continue
248
+ if rec_fmt != fmt:
249
+ issues.append(f"line {line}: format `{rec_fmt}` (file is `{fmt}`)")
250
+ continue
251
+ issues.extend(_record_issues(record, fmt, line))
252
+ key = json.dumps(record, sort_keys=True)
253
+ if key in seen:
254
+ duplicates += 1
255
+ issues.append(f"line {line}: exact duplicate of line {seen[key]}")
256
+ else:
257
+ seen[key] = line
258
+
259
+ truncated = len(issues) > max_issues
260
+ return {
261
+ "records": len(records),
262
+ "format": fmt,
263
+ "issues": issues[:max_issues],
264
+ "duplicates": duplicates,
265
+ "truncated": truncated,
266
+ }
267
+
268
+
155
269
  def inspect_stats(path: str) -> Dict[str, Any]:
156
270
  """Lightweight dataset statistics for `troy data inspect`-style output."""
157
271
  records = _read_records(Path(path).expanduser())
@@ -0,0 +1,247 @@
1
+ """Synthesize training data with a local teacher model.
2
+
3
+ `troy data synth` turns your documents (or a task description) into a
4
+ train.jsonl — the teacher runs locally via mlx-lm, so nothing leaves the Mac.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import json
10
+ import re
11
+ from pathlib import Path
12
+ from typing import Any, Dict, List, Optional, Tuple
13
+
14
+ from .hardware import detect
15
+
16
+ # Source files worth mining for Q&A pairs.
17
+ SOURCE_SUFFIXES = {
18
+ ".md", ".mdx", ".txt", ".rst", ".adoc",
19
+ ".py", ".js", ".ts", ".go", ".rs", ".java", ".rb", ".swift",
20
+ ".yaml", ".yml", ".toml", ".json", ".html",
21
+ }
22
+
23
+ # (min unified memory GB, teacher) — instruct models, 4-bit; the teacher only
24
+ # runs inference, so it can be larger than what the same Mac can train.
25
+ TEACHER_GUIDANCE = [
26
+ (64, "mlx-community/Qwen3-30B-A3B-4bit"),
27
+ (36, "mlx-community/Qwen3-14B-4bit"),
28
+ (24, "mlx-community/Qwen3-8B-4bit"),
29
+ (16, "mlx-community/Qwen3-4B-4bit"),
30
+ (0, "mlx-community/Qwen3-1.7B-4bit"),
31
+ ]
32
+
33
+ PAIRS_PER_CALL = 10
34
+
35
+
36
+ def pick_teacher(memory_gb: Optional[float] = None) -> str:
37
+ if memory_gb is None:
38
+ memory_gb = detect().memory_gb
39
+ for min_gb, model in TEACHER_GUIDANCE:
40
+ if memory_gb >= min_gb:
41
+ return model
42
+ return TEACHER_GUIDANCE[-1][1]
43
+
44
+
45
+ def collect_sources(path: Path) -> List[Tuple[str, str]]:
46
+ """Read (name, text) from a file or a folder of text-ish files."""
47
+ path = path.expanduser()
48
+ if path.is_file():
49
+ return [(path.name, path.read_text(errors="ignore"))]
50
+ if not path.is_dir():
51
+ raise ValueError(f"Source not found: {path}")
52
+ sources = []
53
+ for f in sorted(path.rglob("*")):
54
+ if not f.is_file() or f.suffix.lower() not in SOURCE_SUFFIXES:
55
+ continue
56
+ if any(part.startswith(".") for part in f.relative_to(path).parts):
57
+ continue # skip hidden dirs/files (.git, .venv, ...)
58
+ text = f.read_text(errors="ignore").strip()
59
+ if text:
60
+ sources.append((str(f.relative_to(path)), text))
61
+ if not sources:
62
+ raise ValueError(
63
+ f"No readable text files under {path} "
64
+ f"(looked for {', '.join(sorted(SOURCE_SUFFIXES))})"
65
+ )
66
+ return sources
67
+
68
+
69
+ def chunk_text(text: str, size: int = 4000, overlap: int = 200) -> List[str]:
70
+ """Split on paragraph boundaries where possible, hard-split otherwise."""
71
+ if len(text) <= size:
72
+ return [text]
73
+ chunks = []
74
+ start = 0
75
+ while start < len(text):
76
+ end = start + size
77
+ if end < len(text):
78
+ cut = text.rfind("\n\n", start + size // 2, end)
79
+ if cut != -1:
80
+ end = cut
81
+ chunks.append(text[start:end].strip())
82
+ if end >= len(text):
83
+ break
84
+ start = max(end - overlap, start + 1)
85
+ return [c for c in chunks if c]
86
+
87
+
88
+ _THINK_RE = re.compile(r"<think>.*?</think>", re.DOTALL)
89
+
90
+
91
+ def parse_pairs(raw: str, keys: Tuple[str, ...]) -> List[Dict[str, str]]:
92
+ """Extract JSON objects with the given string keys from teacher output.
93
+
94
+ Tolerates thinking blocks, code fences, prose between objects, and
95
+ objects that span multiple lines.
96
+ """
97
+ raw = _THINK_RE.sub("", raw)
98
+ pairs = []
99
+ decoder = json.JSONDecoder()
100
+ pos = 0
101
+ while True:
102
+ brace = raw.find("{", pos)
103
+ if brace == -1:
104
+ break
105
+ try:
106
+ obj, consumed = decoder.raw_decode(raw[brace:])
107
+ pos = brace + consumed
108
+ except json.JSONDecodeError:
109
+ pos = brace + 1
110
+ continue
111
+ if not isinstance(obj, dict):
112
+ continue
113
+ values = {k: obj.get(k) for k in keys}
114
+ if all(isinstance(v, str) and v.strip() for v in values.values()):
115
+ pairs.append({k: v.strip() for k, v in values.items()})
116
+ return pairs
117
+
118
+
119
+ def _chat_prompt(task: str, material: Optional[str], n: int) -> str:
120
+ lines = [
121
+ f"You are creating fine-tuning data for this task: {task}",
122
+ "",
123
+ f"Write {n} diverse training examples. Output ONLY JSON objects, "
124
+ 'one per line, each exactly: {"user": "...", "assistant": "..."}',
125
+ "Vary phrasing, length, and difficulty. No numbering, no commentary.",
126
+ ]
127
+ if material:
128
+ lines += [
129
+ "",
130
+ "Ground every example in this source material — questions a reader "
131
+ "would ask about it, answered faithfully from it:",
132
+ "---",
133
+ material,
134
+ "---",
135
+ ]
136
+ return "\n".join(lines)
137
+
138
+
139
+ def _preference_prompt(task: str, material: Optional[str], n: int) -> str:
140
+ lines = [
141
+ f"You are creating preference-tuning data for this task: {task}",
142
+ "",
143
+ f"Write {n} diverse training examples. Output ONLY JSON objects, one "
144
+ 'per line, each exactly: '
145
+ '{"prompt": "...", "chosen": "...", "rejected": "..."}',
146
+ "`chosen` is a genuinely good response; `rejected` is plausible but "
147
+ "clearly worse for the task (vague, bloated, off-style, or subtly "
148
+ "wrong). No numbering, no commentary.",
149
+ ]
150
+ if material:
151
+ lines += [
152
+ "",
153
+ "Ground every example in this source material:",
154
+ "---",
155
+ material,
156
+ "---",
157
+ ]
158
+ return "\n".join(lines)
159
+
160
+
161
+ def run_synth(
162
+ n: int,
163
+ out_path: Path,
164
+ fmt: str,
165
+ teacher: str,
166
+ seed_task: Optional[str],
167
+ source: Optional[Path],
168
+ max_tokens: int,
169
+ temperature: float,
170
+ ) -> Dict[str, Any]:
171
+ """Generate n examples; returns stats. Writes JSONL to out_path."""
172
+ from mlx_lm.generate import generate
173
+ from mlx_lm.sample_utils import make_sampler
174
+ from mlx_lm.utils import load
175
+
176
+ chunks: List[Optional[str]]
177
+ if source is not None:
178
+ texts = collect_sources(source)
179
+ chunks = [c for _, text in texts for c in chunk_text(text)]
180
+ else:
181
+ chunks = [None] # pure seed-description mode
182
+
183
+ task = seed_task or (
184
+ "answering questions about the source material accurately and concisely"
185
+ )
186
+ keys = ("prompt", "chosen", "rejected") if fmt == "preference" else ("user", "assistant")
187
+ build = _preference_prompt if fmt == "preference" else _chat_prompt
188
+
189
+ print(f"Loading teacher {teacher} ...")
190
+ model, tokenizer = load(teacher)
191
+ sampler = make_sampler(temp=temperature)
192
+
193
+ seen = set()
194
+ records: List[Dict[str, Any]] = []
195
+ calls = failures = 0
196
+ chunk_i = 0
197
+ max_calls = (n // PAIRS_PER_CALL + 1) * 4 # generous retry budget
198
+
199
+ while len(records) < n and calls < max_calls:
200
+ want = min(PAIRS_PER_CALL, n - len(records))
201
+ prompt = build(task, chunks[chunk_i % len(chunks)], want)
202
+ chunk_i += 1
203
+ calls += 1
204
+ messages = [{"role": "user", "content": prompt}]
205
+ templated = tokenizer.apply_chat_template(
206
+ messages, add_generation_prompt=True, return_dict=False
207
+ )
208
+ raw = generate(
209
+ model, tokenizer, templated, max_tokens=max_tokens, sampler=sampler
210
+ )
211
+ pairs = parse_pairs(raw, keys)
212
+ if not pairs:
213
+ failures += 1
214
+ continue
215
+ for pair in pairs:
216
+ key = pair[keys[0]].lower()
217
+ if key in seen:
218
+ continue
219
+ seen.add(key)
220
+ if fmt == "preference":
221
+ records.append(pair)
222
+ else:
223
+ records.append(
224
+ {
225
+ "messages": [
226
+ {"role": "user", "content": pair["user"]},
227
+ {"role": "assistant", "content": pair["assistant"]},
228
+ ]
229
+ }
230
+ )
231
+ if len(records) >= n:
232
+ break
233
+ print(f" {len(records)}/{n} examples ({calls} teacher calls)")
234
+
235
+ out_path.parent.mkdir(parents=True, exist_ok=True)
236
+ with open(out_path, "w") as f:
237
+ for r in records:
238
+ f.write(json.dumps(r, ensure_ascii=False) + "\n")
239
+
240
+ return {
241
+ "records": len(records),
242
+ "requested": n,
243
+ "teacher_calls": calls,
244
+ "empty_responses": failures,
245
+ "chunks": 0 if chunks == [None] else len(chunks),
246
+ "out": str(out_path),
247
+ }
@@ -0,0 +1,94 @@
1
+ """Vision-language fine-tuning (wraps mlx-vlm's LoRA trainer).
2
+
3
+ Contract: `data.train` is a FOLDER containing images plus a metadata.jsonl
4
+ with {"file_name", "question", "answer"} rows (HF imagefolder format).
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import json
10
+ import math
11
+ import subprocess
12
+ import sys
13
+ from pathlib import Path
14
+
15
+ from .config import TroyConfig
16
+
17
+
18
+ def ensure_mlx_vlm() -> None:
19
+ try:
20
+ import mlx_vlm # noqa: F401
21
+ except ImportError:
22
+ raise SystemExit(
23
+ "Vision fine-tuning needs mlx-vlm. Install with:\n"
24
+ " pip install 'troy-cli[vision]'"
25
+ )
26
+
27
+
28
+ def count_records(data_dir: Path) -> int:
29
+ meta = data_dir / "metadata.jsonl"
30
+ if not meta.exists():
31
+ raise SystemExit(
32
+ f"{data_dir} has no metadata.jsonl. Vision data is a folder of images "
33
+ 'plus metadata.jsonl rows: {"file_name": ..., "question": ..., "answer": ...}'
34
+ )
35
+ with open(meta) as f:
36
+ rows = [json.loads(line) for line in f if line.strip()]
37
+ if rows and not {"file_name", "question", "answer"} <= set(rows[0]):
38
+ raise SystemExit(
39
+ "metadata.jsonl rows need file_name, question, and answer fields."
40
+ )
41
+ return len(rows)
42
+
43
+
44
+ def run_vision_sft(config: TroyConfig) -> None:
45
+ ensure_mlx_vlm()
46
+ data_dir = Path(config.data.train).expanduser()
47
+ n = count_records(data_dir)
48
+ print(f"Vision data: {n} images (folder: {data_dir})")
49
+
50
+ t = config.training
51
+ batch_size = 1 if t.batch_size == "auto" else int(t.batch_size)
52
+ iters = t.iters or max(1, math.ceil((t.epochs or 3) * n / batch_size))
53
+
54
+ config.adapter_path.mkdir(parents=True, exist_ok=True)
55
+ cmd = [
56
+ sys.executable, "-m", "mlx_vlm.lora",
57
+ "--model-path", config.base,
58
+ "--dataset", str(data_dir),
59
+ "--split", "train",
60
+ "--batch-size", str(batch_size),
61
+ "--iters", str(iters),
62
+ "--learning-rate", str(t.lr),
63
+ "--lora-rank", str(t.lora.r),
64
+ "--lora-alpha", str(t.lora.alpha),
65
+ "--lora-dropout", str(t.lora.dropout),
66
+ "--max-seq-length", str(t.seq_len),
67
+ "--steps-per-save", str(t.save_every),
68
+ "--output-path", str(config.adapter_path),
69
+ ]
70
+ if t.grad_checkpoint:
71
+ cmd.append("--grad-checkpoint")
72
+
73
+ print(f"Training: task=sft (vision) batch_size={batch_size} iters={iters} lr={t.lr}")
74
+ result = subprocess.run(cmd)
75
+ if result.returncode != 0:
76
+ raise SystemExit(result.returncode)
77
+ print(f"\nDone. Adapter saved to {config.adapter_path}")
78
+
79
+
80
+ def run_vision_chat(
81
+ model: str, adapter: str | None, image: str, prompt: str, max_tokens: int, temperature: float
82
+ ) -> None:
83
+ ensure_mlx_vlm()
84
+ cmd = [
85
+ sys.executable, "-m", "mlx_vlm", "generate",
86
+ "--model", model,
87
+ "--image", image,
88
+ "--prompt", prompt,
89
+ "--max-tokens", str(max_tokens),
90
+ "--temperature", str(temperature),
91
+ ]
92
+ if adapter:
93
+ cmd += ["--adapter-path", adapter]
94
+ raise SystemExit(subprocess.run(cmd).returncode)
@@ -74,3 +74,50 @@ def test_explicit_valid_file(tmp_path):
74
74
  write_jsonl(v, [{"text": "b"}] * 2)
75
75
  train, valid, _ = load_and_prepare(DataConfig(train=str(t), valid=str(v)), "sft")
76
76
  assert len(train) == 6 and len(valid) == 2
77
+
78
+
79
+ def test_validate_clean_file(tmp_path):
80
+ from troy.data import validate_records
81
+
82
+ f = tmp_path / "t.jsonl"
83
+ write_jsonl(f, [{"instruction": f"q{i}", "output": f"a{i}"} for i in range(5)])
84
+ report = validate_records(str(f))
85
+ assert report["format"] == "alpaca"
86
+ assert report["records"] == 5
87
+ assert report["issues"] == []
88
+
89
+
90
+ def test_validate_catches_problems(tmp_path):
91
+ from troy.data import validate_records
92
+
93
+ f = tmp_path / "t.jsonl"
94
+ rows = [
95
+ {"instruction": "q", "output": "a"}, # ok
96
+ {"instruction": "", "output": "a"}, # empty field
97
+ {"prompt": "p", "completion": "c"}, # mixed format
98
+ {"instruction": "q", "output": "a"}, # duplicate of line 1
99
+ ]
100
+ with open(f, "w") as fh:
101
+ for r in rows:
102
+ fh.write(json.dumps(r) + "\n")
103
+ fh.write("{broken json\n")
104
+
105
+ report = validate_records(str(f))
106
+ text = "\n".join(report["issues"])
107
+ assert "empty or missing `instruction`" in text
108
+ assert "format `completions`" in text
109
+ assert "duplicate of line 1" in text
110
+ assert "invalid JSON" in text
111
+ assert report["duplicates"] == 1
112
+
113
+
114
+ def test_validate_preference_and_chat(tmp_path):
115
+ from troy.data import validate_records
116
+
117
+ f = tmp_path / "p.jsonl"
118
+ write_jsonl(f, [{"prompt": "p", "chosen": "same", "rejected": "same"}])
119
+ assert "chosen == rejected" in validate_records(str(f))["issues"][0]
120
+
121
+ f2 = tmp_path / "c.jsonl"
122
+ write_jsonl(f2, [{"messages": [{"role": "user", "content": "hi"}]}])
123
+ assert "no assistant message" in validate_records(str(f2))["issues"][0]
@@ -0,0 +1,82 @@
1
+ import pytest
2
+
3
+ from troy.synth import chunk_text, collect_sources, parse_pairs, pick_teacher
4
+
5
+
6
+ def test_pick_teacher_scales_with_memory():
7
+ assert "1.7B" in pick_teacher(8)
8
+ assert "4B" in pick_teacher(16)
9
+ assert "8B" in pick_teacher(24)
10
+ assert "14B" in pick_teacher(36)
11
+ assert "30B" in pick_teacher(64)
12
+
13
+
14
+ def test_chunk_text_short_passthrough():
15
+ assert chunk_text("hello", size=100) == ["hello"]
16
+
17
+
18
+ def test_chunk_text_splits_on_paragraphs():
19
+ text = ("para one " * 50 + "\n\n" + "para two " * 50).strip()
20
+ chunks = chunk_text(text, size=500, overlap=50)
21
+ assert len(chunks) >= 2
22
+ assert all(len(c) <= 500 for c in chunks)
23
+ # nothing lost beyond whitespace at the seams
24
+ assert "para two" in chunks[-1]
25
+
26
+
27
+ def test_parse_pairs_jsonl():
28
+ raw = (
29
+ '{"user": "What is Troy?", "assistant": "A fine-tuning CLI."}\n'
30
+ '{"user": "Which Macs?", "assistant": "Apple Silicon, M1+."}\n'
31
+ )
32
+ pairs = parse_pairs(raw, ("user", "assistant"))
33
+ assert len(pairs) == 2
34
+ assert pairs[0]["user"] == "What is Troy?"
35
+
36
+
37
+ def test_parse_pairs_tolerates_noise():
38
+ raw = (
39
+ "<think>let me write some examples</think>\n"
40
+ "Here are the examples:\n"
41
+ "```json\n"
42
+ '{"user": "q1",\n "assistant": "a1"}\n'
43
+ "```\n"
44
+ "not json at all {broken\n"
45
+ '{"user": "q2", "assistant": "a2", "extra": 1}\n'
46
+ '{"user": "", "assistant": "empty user skipped"}\n'
47
+ '{"user": "no assistant key"}\n'
48
+ )
49
+ pairs = parse_pairs(raw, ("user", "assistant"))
50
+ assert [p["user"] for p in pairs] == ["q1", "q2"]
51
+ assert "extra" not in pairs[1]
52
+
53
+
54
+ def test_parse_pairs_preference_keys():
55
+ raw = '{"prompt": "p", "chosen": "good", "rejected": "bad"}'
56
+ pairs = parse_pairs(raw, ("prompt", "chosen", "rejected"))
57
+ assert pairs == [{"prompt": "p", "chosen": "good", "rejected": "bad"}]
58
+
59
+
60
+ def test_collect_sources_file_and_folder(tmp_path):
61
+ (tmp_path / "a.md").write_text("# Doc A")
62
+ (tmp_path / "sub").mkdir()
63
+ (tmp_path / "sub" / "b.txt").write_text("Doc B")
64
+ (tmp_path / ".git").mkdir()
65
+ (tmp_path / ".git" / "c.md").write_text("hidden")
66
+ (tmp_path / "img.png").write_bytes(b"\x89PNG")
67
+
68
+ sources = collect_sources(tmp_path)
69
+ names = [n for n, _ in sources]
70
+ assert names == ["a.md", "sub/b.txt"]
71
+
72
+ single = collect_sources(tmp_path / "a.md")
73
+ assert single == [("a.md", "# Doc A")]
74
+
75
+
76
+ def test_collect_sources_errors(tmp_path):
77
+ with pytest.raises(ValueError, match="not found"):
78
+ collect_sources(tmp_path / "missing")
79
+ empty = tmp_path / "empty"
80
+ empty.mkdir()
81
+ with pytest.raises(ValueError, match="No readable"):
82
+ collect_sources(empty)
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes