troy-cli 0.3.0__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.
- {troy_cli-0.3.0 → troy_cli-0.4.0}/PKG-INFO +1 -1
- {troy_cli-0.3.0 → troy_cli-0.4.0}/pyproject.toml +1 -1
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/__init__.py +1 -1
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/cli.py +97 -6
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/data.py +114 -0
- troy_cli-0.4.0/src/troy/synth.py +247 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/tests/test_data.py +47 -0
- troy_cli-0.4.0/tests/test_synth.py +82 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/.gitignore +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/README.md +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/chat.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/config.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/evaluate.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/export.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/hardware.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/push.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/serve.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/templates.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/train_dpo.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/train_orpo.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/train_sft.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/src/troy/train_vision.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/tests/test_config.py +0 -0
- {troy_cli-0.3.0 → troy_cli-0.4.0}/tests/test_hardware.py +0 -0
|
@@ -319,15 +319,19 @@ def push(
|
|
|
319
319
|
run_push(folder, repo, private=not public)
|
|
320
320
|
|
|
321
321
|
|
|
322
|
-
|
|
323
|
-
|
|
324
|
-
|
|
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(
|
|
325
332
|
path: Path = typer.Argument(help="Dataset file (.jsonl, .json, .csv)."),
|
|
326
333
|
) -> None:
|
|
327
334
|
"""Inspect a dataset: record count, detected format, sizes."""
|
|
328
|
-
if action != "inspect":
|
|
329
|
-
console.print("[red]Only `troy data inspect <path>` is supported.[/red]")
|
|
330
|
-
raise typer.Exit(1)
|
|
331
335
|
from .data import inspect_stats
|
|
332
336
|
|
|
333
337
|
stats = inspect_stats(str(path))
|
|
@@ -341,5 +345,92 @@ def data(
|
|
|
341
345
|
console.print(table)
|
|
342
346
|
|
|
343
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
|
+
|
|
344
435
|
if __name__ == "__main__":
|
|
345
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
|
+
}
|
|
@@ -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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|