ftgate 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.
- ftgate/__init__.py +2 -0
- ftgate/assets/assert_tools.js +77 -0
- ftgate/cli.py +125 -0
- ftgate/compare.py +119 -0
- ftgate/conversations.py +84 -0
- ftgate/data.py +318 -0
- ftgate/pytest_plugin.py +43 -0
- ftgate/reference.py +37 -0
- ftgate/render.py +37 -0
- ftgate/runtimes/__init__.py +19 -0
- ftgate/runtimes/llama.py +25 -0
- ftgate/runtimes/ollama.py +67 -0
- ftgate/testing.py +31 -0
- ftgate/tools_eval.py +220 -0
- ftgate-0.1.0.dist-info/METADATA +167 -0
- ftgate-0.1.0.dist-info/RECORD +20 -0
- ftgate-0.1.0.dist-info/WHEEL +5 -0
- ftgate-0.1.0.dist-info/entry_points.txt +5 -0
- ftgate-0.1.0.dist-info/licenses/LICENSE +202 -0
- ftgate-0.1.0.dist-info/top_level.txt +1 -0
ftgate/__init__.py
ADDED
|
@@ -0,0 +1,77 @@
|
|
|
1
|
+
// ftgate: judge one tool-calling response against the expected calls.
|
|
2
|
+
//
|
|
3
|
+
// Written once by ftgate, run by promptfoo for every case. It reads the
|
|
4
|
+
// model's output in whatever shape the provider produced — OpenAI
|
|
5
|
+
// tool_calls JSON, a list of function_call items, or raw text with
|
|
6
|
+
// <tool_call> blocks — and reports named scores rather than one pass/fail,
|
|
7
|
+
// so the table can separate "wrong tool", "wrong arguments" and "the
|
|
8
|
+
// call was there but the provider's parser lost it".
|
|
9
|
+
module.exports = (output, context) => {
|
|
10
|
+
const expected = context.vars.__expected_calls; // [{name, arguments}]
|
|
11
|
+
const norm = (v) => {
|
|
12
|
+
if (Array.isArray(v)) return v.map(norm);
|
|
13
|
+
if (v && typeof v === "object") { const o = {}; for (const k of Object.keys(v).sort()) o[k] = norm(v[k]); return o; }
|
|
14
|
+
if (typeof v === "number") return Number(v);
|
|
15
|
+
if (typeof v === "boolean") return v;
|
|
16
|
+
return String(v).trim().toLowerCase();
|
|
17
|
+
};
|
|
18
|
+
// Expected arguments must be present and equal; extra (optional) arguments
|
|
19
|
+
// are allowed and reported separately, so a model that adds unit="celsius"
|
|
20
|
+
// to a weather call is not marked wrong for being thorough.
|
|
21
|
+
const subsetKey = (c, exp) => {
|
|
22
|
+
const got = c.arguments || {}, want = exp.arguments || {};
|
|
23
|
+
const picked = {}; for (const k of Object.keys(want)) picked[k] = got[k];
|
|
24
|
+
return JSON.stringify([c.name, norm(picked)]);
|
|
25
|
+
};
|
|
26
|
+
const wantKey = (e) => JSON.stringify([e.name, norm(e.arguments || {})]);
|
|
27
|
+
|
|
28
|
+
let structured = [];
|
|
29
|
+
let text = typeof output === "string" ? output : JSON.stringify(output);
|
|
30
|
+
const takeCall = (c) => {
|
|
31
|
+
const fn = c.function || c;
|
|
32
|
+
let args = fn.arguments ?? fn.input ?? {};
|
|
33
|
+
if (typeof args === "string") { try { args = JSON.parse(args); } catch (e) { return { name: fn.name, arguments: null, malformed: true }; } }
|
|
34
|
+
return { name: fn.name, arguments: args };
|
|
35
|
+
};
|
|
36
|
+
try {
|
|
37
|
+
const o = typeof output === "string" ? JSON.parse(output) : output;
|
|
38
|
+
const list = Array.isArray(o) ? o : (o && o.tool_calls) ? o.tool_calls : (o && o.name) ? [o] : [];
|
|
39
|
+
structured = list.filter(c => c && (c.function || c.name)).map(takeCall);
|
|
40
|
+
} catch (e) { /* not JSON: raw text */ }
|
|
41
|
+
|
|
42
|
+
let fromText = [];
|
|
43
|
+
const re = /<tool_call>\s*([\s\S]*?)\s*<\/tool_call>/g; let m;
|
|
44
|
+
while ((m = re.exec(text)) !== null) {
|
|
45
|
+
try { const d = JSON.parse(m[1]); fromText.push({ name: d.name, arguments: d.arguments || {} }); }
|
|
46
|
+
catch (e) { fromText.push({ name: null, arguments: null, malformed: true }); }
|
|
47
|
+
}
|
|
48
|
+
|
|
49
|
+
const calls = structured.length ? structured : fromText;
|
|
50
|
+
const parserLost = structured.length === 0 && fromText.length > 0; // the model called; the provider did not surface it
|
|
51
|
+
const malformed = calls.some(c => c.malformed);
|
|
52
|
+
const names = (a) => a.map(c => c.name).sort().join("|");
|
|
53
|
+
const namesOk = names(calls) === names(expected);
|
|
54
|
+
// Pair each expected call with a produced call of the same name (greedy).
|
|
55
|
+
let argsOk = namesOk && !malformed, extra = false;
|
|
56
|
+
if (argsOk) {
|
|
57
|
+
const pool = calls.slice();
|
|
58
|
+
for (const e of expected) {
|
|
59
|
+
const i = pool.findIndex(c => c.name === e.name && subsetKey(c, e) === wantKey(e));
|
|
60
|
+
if (i < 0) { argsOk = false; break; }
|
|
61
|
+
const got = pool[i].arguments || {};
|
|
62
|
+
if (Object.keys(got).some(k => !(k in (e.arguments || {})))) extra = true;
|
|
63
|
+
pool.splice(i, 1);
|
|
64
|
+
}
|
|
65
|
+
}
|
|
66
|
+
|
|
67
|
+
return {
|
|
68
|
+
pass: argsOk,
|
|
69
|
+
score: argsOk ? 1 : namesOk ? 0.5 : 0,
|
|
70
|
+
reason: argsOk ? (extra ? "ok, extra optional arguments" : "exact")
|
|
71
|
+
: !calls.length ? "no call"
|
|
72
|
+
: malformed ? "malformed arguments: " + text.slice(0, 160)
|
|
73
|
+
: !namesOk ? `tools ${names(calls) || "-"} vs expected ${names(expected)}`
|
|
74
|
+
: "arguments differ: got " + JSON.stringify(calls.map(c => c.arguments)) + " expected " + JSON.stringify(expected.map(e => e.arguments)),
|
|
75
|
+
namedScores: { called: calls.length ? 1 : 0, names: namesOk ? 1 : 0, args: argsOk ? 1 : 0, extra_args: extra ? 1 : 0, parser_lost: parserLost ? 1 : 0, malformed: malformed ? 1 : 0 },
|
|
76
|
+
};
|
|
77
|
+
};
|
ftgate/cli.py
ADDED
|
@@ -0,0 +1,125 @@
|
|
|
1
|
+
"""ftgate command line."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
import argparse
|
|
5
|
+
import sys
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
|
|
8
|
+
from ftgate import __version__, reference, render
|
|
9
|
+
from ftgate.compare import compare
|
|
10
|
+
from ftgate.conversations import BUILTIN, load_dataset
|
|
11
|
+
from ftgate.runtimes import llama as llama_rt
|
|
12
|
+
from ftgate.runtimes import ollama as ollama_rt
|
|
13
|
+
from ftgate import data as data_mod
|
|
14
|
+
import json as _json
|
|
15
|
+
from ftgate import tools_eval
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def cmd_tools(args) -> int:
|
|
19
|
+
tools = _json.loads(Path(args.tools).read_text(encoding="utf-8")) if args.tools else []
|
|
20
|
+
if args.from_dataset:
|
|
21
|
+
cases, ds_tools = tools_eval.cases_from_dataset(Path(args.from_dataset), args.holdout)
|
|
22
|
+
tools = tools or ds_tools
|
|
23
|
+
else:
|
|
24
|
+
cases = tools_eval.load_cases(Path(args.cases))
|
|
25
|
+
if not cases:
|
|
26
|
+
print("no cases", file=sys.stderr); return 2
|
|
27
|
+
if not tools:
|
|
28
|
+
print("no tools: pass --tools tools.json (or a dataset whose rows carry `tools`)", file=sys.stderr); return 2
|
|
29
|
+
arms = tools_eval.run(cases, tools, args.provider, keep=Path(args.keep) if args.keep else None, concurrency=args.concurrency)
|
|
30
|
+
if args.format == "json":
|
|
31
|
+
print(_json.dumps([a.__dict__ for a in arms], indent=2, ensure_ascii=False))
|
|
32
|
+
else:
|
|
33
|
+
print(tools_eval.to_text(arms), end="")
|
|
34
|
+
if args.min_exact is not None:
|
|
35
|
+
worst = min((a.args / a.n) for a in arms if a.n) if arms else 0
|
|
36
|
+
return 1 if worst < args.min_exact else 0
|
|
37
|
+
return 0
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def cmd_data(args) -> int:
|
|
41
|
+
reps = [data_mod.lint(Path(p), model=args.model, max_seq_len=args.max_seq_len,
|
|
42
|
+
eval_path=Path(args.eval) if args.eval else None, limit=args.limit) for p in args.path]
|
|
43
|
+
if args.format == "json":
|
|
44
|
+
print(_json.dumps([r.as_dict() for r in reps] if len(reps) > 1 else reps[0].as_dict(), indent=2, ensure_ascii=False))
|
|
45
|
+
else:
|
|
46
|
+
print("".join(data_mod.to_text(r) for r in reps), end="")
|
|
47
|
+
errors = sum(r.count("error") for r in reps); warnings = sum(r.count("warning") for r in reps)
|
|
48
|
+
if args.fail_on == "error":
|
|
49
|
+
return 1 if errors else 0
|
|
50
|
+
if args.fail_on == "warning":
|
|
51
|
+
return 1 if errors or warnings else 0
|
|
52
|
+
return 0
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def cmd_template(args) -> int:
|
|
56
|
+
convs = load_dataset(Path(args.dataset), args.sample) if args.dataset else BUILTIN
|
|
57
|
+
if not args.llama and not args.ollama:
|
|
58
|
+
print("nothing to compare against: pass --llama URL and/or --ollama MODEL", file=sys.stderr); return 2
|
|
59
|
+
|
|
60
|
+
results = []
|
|
61
|
+
for conv in convs:
|
|
62
|
+
msgs = conv.inference_messages
|
|
63
|
+
ref = reference.render(args.model, msgs, conv.tools, add_generation_prompt=True)
|
|
64
|
+
if args.llama:
|
|
65
|
+
rt = llama_rt.render(args.llama.rstrip("/"), msgs, conv.tools)
|
|
66
|
+
results.append(compare(conv.name, ref, rt, reference.tokenize(args.model, rt.text)))
|
|
67
|
+
if args.ollama:
|
|
68
|
+
rt = ollama_rt.render(args.ollama_url.rstrip("/"), args.ollama, msgs, conv.tools, verify=not args.no_verify)
|
|
69
|
+
results.append(compare(conv.name, ref, rt, reference.tokenize(args.model, rt.text)))
|
|
70
|
+
|
|
71
|
+
print(render.to_json(results, args.model) if args.format == "json" else render.to_text(results, args.model, args.diff), end="")
|
|
72
|
+
if args.fail_on == "mismatch":
|
|
73
|
+
return 1 if any(r.verdict == "mismatch" for r in results) else 0
|
|
74
|
+
if args.fail_on == "drift":
|
|
75
|
+
return 1 if any(r.verdict != "identical" for r in results) else 0
|
|
76
|
+
return 0
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def main(argv: list[str] | None = None) -> int:
|
|
80
|
+
ap = argparse.ArgumentParser(prog="ftgate", description="Do runtimes send the bytes the model was trained on?")
|
|
81
|
+
ap.add_argument("--version", action="version", version=__version__)
|
|
82
|
+
sub = ap.add_subparsers(dest="command", required=True)
|
|
83
|
+
|
|
84
|
+
t = sub.add_parser("template", help="compare runtime prompts and tokens against the training-side rendering")
|
|
85
|
+
t.add_argument("--model", required=True, help="HF model id or path whose tokenizer and chat template were used for training")
|
|
86
|
+
t.add_argument("--llama", metavar="URL", help="llama-server base URL, e.g. http://localhost:8080")
|
|
87
|
+
t.add_argument("--ollama", metavar="MODEL", help="Ollama model name, e.g. qwen2.5:0.5b-instruct")
|
|
88
|
+
t.add_argument("--ollama-url", default="http://localhost:11434")
|
|
89
|
+
t.add_argument("--no-verify", action="store_true", help="skip the live cross-check of Ollama's prompt token count")
|
|
90
|
+
t.add_argument("--dataset", metavar="JSONL", help="your own conversations (messages[, tools] per line) instead of the built-in set")
|
|
91
|
+
t.add_argument("--sample", type=int, default=20, help="rows to take from --dataset (default 20)")
|
|
92
|
+
t.add_argument("--diff", action="store_true", help="print unified diffs for findings")
|
|
93
|
+
t.add_argument("--format", choices=["text", "json"], default="text")
|
|
94
|
+
t.add_argument("--fail-on", choices=["never", "drift", "mismatch"], default="never")
|
|
95
|
+
t.set_defaults(func=cmd_template)
|
|
96
|
+
|
|
97
|
+
d = sub.add_parser("data", help="lint a fine-tuning dataset with the target model's tokenizer and template")
|
|
98
|
+
d.add_argument("path", nargs="+", help="JSONL file(s): messages / ShareGPT / Alpaca / preference rows")
|
|
99
|
+
d.add_argument("--model", help="HF model id; enables tokenizer and template checks")
|
|
100
|
+
d.add_argument("--max-seq-len", type=int, help="the trainer's cutoff; rows beyond it are reported")
|
|
101
|
+
d.add_argument("--eval", metavar="JSONL", help="eval set to check for leakage (near-duplicates)")
|
|
102
|
+
d.add_argument("--limit", type=int)
|
|
103
|
+
d.add_argument("--format", choices=["text", "json"], default="text")
|
|
104
|
+
d.add_argument("--fail-on", choices=["never", "warning", "error"], default="never")
|
|
105
|
+
d.set_defaults(func=cmd_data)
|
|
106
|
+
|
|
107
|
+
tl = sub.add_parser("tools", help="evaluate tool calling on your tools and cases, base vs tuned vs quantised, via promptfoo")
|
|
108
|
+
tl.add_argument("--tools", metavar="JSON", help="OpenAI-style tool list")
|
|
109
|
+
tl.add_argument("--cases", metavar="FILE", help="cases: {user, calls[, system]} per row (json/jsonl/yaml)")
|
|
110
|
+
tl.add_argument("--from-dataset", metavar="JSONL", help="hold out cases from your training rows instead")
|
|
111
|
+
tl.add_argument("--holdout", type=int, default=50)
|
|
112
|
+
tl.add_argument("--provider", action="append", required=True,
|
|
113
|
+
help="repeatable: ollama:MODEL | openai:MODEL@http://host:port/v1 | any promptfoo provider id")
|
|
114
|
+
tl.add_argument("--keep", metavar="DIR", help="keep the generated promptfoo config and results here")
|
|
115
|
+
tl.add_argument("--concurrency", type=int, default=1, help="requests in flight (default 1: local servers answer one at a time)")
|
|
116
|
+
tl.add_argument("--min-exact", type=float, help="exit 1 if any arm's exact-args rate is below this")
|
|
117
|
+
tl.add_argument("--format", choices=["text", "json"], default="text")
|
|
118
|
+
tl.set_defaults(func=cmd_tools)
|
|
119
|
+
|
|
120
|
+
args = ap.parse_args(argv)
|
|
121
|
+
return args.func(args)
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
if __name__ == "__main__":
|
|
125
|
+
sys.exit(main())
|
ftgate/compare.py
ADDED
|
@@ -0,0 +1,119 @@
|
|
|
1
|
+
"""Compare a runtime's bytes with the reference and say what kind of difference it is.
|
|
2
|
+
|
|
3
|
+
Findings are ordered by how much they matter to a model that was trained
|
|
4
|
+
on the reference:
|
|
5
|
+
|
|
6
|
+
schema_not_json a tool definition reached the model as something other
|
|
7
|
+
than JSON — the schema itself is gone
|
|
8
|
+
injected_system the runtime added a system prompt the training bytes
|
|
9
|
+
did not have (Ollama's Modelfile SYSTEM)
|
|
10
|
+
special_prefix the runtime prepends a special token the training
|
|
11
|
+
bytes did not have (the BOS case)
|
|
12
|
+
structural text differs beyond whitespace/JSON spacing
|
|
13
|
+
spacing same content, different whitespace or JSON separators —
|
|
14
|
+
different tokens, same meaning
|
|
15
|
+
tokenizer identical text, different ids — the runtime's
|
|
16
|
+
tokenizer disagrees with the training tokenizer
|
|
17
|
+
identical byte-for-byte and id-for-id
|
|
18
|
+
"""
|
|
19
|
+
from __future__ import annotations
|
|
20
|
+
|
|
21
|
+
import difflib
|
|
22
|
+
import json
|
|
23
|
+
import re
|
|
24
|
+
from dataclasses import dataclass, field
|
|
25
|
+
|
|
26
|
+
from ftgate.reference import Rendering
|
|
27
|
+
from ftgate.runtimes import RuntimeRendering
|
|
28
|
+
|
|
29
|
+
_TOOLS_BLOCK = re.compile(r"<tools>(.*?)</tools>", re.S)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
@dataclass
|
|
33
|
+
class Finding:
|
|
34
|
+
kind: str
|
|
35
|
+
detail: str
|
|
36
|
+
diff: str = ""
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
@dataclass
|
|
40
|
+
class Comparison:
|
|
41
|
+
conversation: str
|
|
42
|
+
runtime: str
|
|
43
|
+
reference_tokens: int
|
|
44
|
+
runtime_tokens: int | None
|
|
45
|
+
findings: list[Finding] = field(default_factory=list)
|
|
46
|
+
notes: list[str] = field(default_factory=list)
|
|
47
|
+
|
|
48
|
+
@property
|
|
49
|
+
def verdict(self) -> str:
|
|
50
|
+
kinds = {f.kind for f in self.findings}
|
|
51
|
+
if kinds & {"schema_not_json", "structural", "special_prefix", "injected_system"}:
|
|
52
|
+
return "mismatch"
|
|
53
|
+
if kinds & {"spacing", "tokenizer"}:
|
|
54
|
+
return "drift"
|
|
55
|
+
return "identical"
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def _canon(text: str) -> str:
|
|
59
|
+
"""Collapse whitespace and JSON separators so spacing-only differences vanish."""
|
|
60
|
+
t = re.sub(r"[ \t]+", " ", text)
|
|
61
|
+
t = re.sub(r"\s*([,:{}\[\]])\s*", r"\1", t)
|
|
62
|
+
return t.strip()
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def _tool_schemas_valid(text: str) -> list[str]:
|
|
66
|
+
"""Every line inside <tools> must parse as JSON with a function name."""
|
|
67
|
+
bad = []
|
|
68
|
+
for block in _TOOLS_BLOCK.findall(text):
|
|
69
|
+
for line in block.strip().splitlines():
|
|
70
|
+
line = line.strip()
|
|
71
|
+
if not line:
|
|
72
|
+
continue
|
|
73
|
+
try:
|
|
74
|
+
obj = json.loads(line)
|
|
75
|
+
fn = obj.get("function", obj)
|
|
76
|
+
if not isinstance(fn, dict) or "name" not in fn:
|
|
77
|
+
bad.append(line[:120])
|
|
78
|
+
except json.JSONDecodeError:
|
|
79
|
+
bad.append(line[:120])
|
|
80
|
+
return bad
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def compare(name: str, ref: Rendering, rt: RuntimeRendering, ref_ids_of_runtime_text: list[int]) -> Comparison:
|
|
84
|
+
c = Comparison(name, rt.runtime, len(ref.ids), len(rt.ids) if rt.ids is not None else None, notes=list(rt.notes))
|
|
85
|
+
|
|
86
|
+
bad = _tool_schemas_valid(rt.text)
|
|
87
|
+
if bad and not _tool_schemas_valid(ref.text):
|
|
88
|
+
c.findings.append(Finding("schema_not_json",
|
|
89
|
+
f"{len(bad)} tool definition(s) are not JSON in the runtime prompt", "\n".join(bad)))
|
|
90
|
+
|
|
91
|
+
if any(n.startswith("Modelfile SYSTEM injected") for n in rt.notes):
|
|
92
|
+
c.findings.append(Finding("injected_system", "runtime prepends a system prompt the training conversation does not have"))
|
|
93
|
+
|
|
94
|
+
if rt.special_prefix and not ref.ids[:len(rt.special_prefix)] == rt.special_prefix:
|
|
95
|
+
c.findings.append(Finding("special_prefix",
|
|
96
|
+
f"runtime prepends special token(s) {rt.special_prefix} that the training bytes do not start with"))
|
|
97
|
+
|
|
98
|
+
if rt.text != ref.text:
|
|
99
|
+
diff = "\n".join(difflib.unified_diff(ref.text.splitlines(), rt.text.splitlines(), "training", rt.runtime, lineterm="", n=1))
|
|
100
|
+
if _canon(rt.text) == _canon(ref.text):
|
|
101
|
+
c.findings.append(Finding("spacing", "same content, different whitespace or JSON spacing — different tokens", diff))
|
|
102
|
+
elif not bad and not any(f.kind == "injected_system" for f in c.findings):
|
|
103
|
+
c.findings.append(Finding("structural", "prompt text differs from the training bytes", diff))
|
|
104
|
+
else:
|
|
105
|
+
c.findings[0].diff = diff
|
|
106
|
+
elif rt.ids is not None and rt.ids != ref.ids:
|
|
107
|
+
c.findings.append(Finding("tokenizer", f"identical text, different ids ({len(rt.ids)} vs {len(ref.ids)})"))
|
|
108
|
+
|
|
109
|
+
if rt.ids is None:
|
|
110
|
+
c.runtime_tokens = len(ref_ids_of_runtime_text)
|
|
111
|
+
c.notes.append("runtime cannot tokenize; ids computed with the training tokenizer on the runtime's text")
|
|
112
|
+
for n in rt.notes:
|
|
113
|
+
if n.startswith("__prompt_eval_count="):
|
|
114
|
+
reported = int(n.split("=", 1)[1])
|
|
115
|
+
c.notes = [x for x in c.notes if not x.startswith("__")]
|
|
116
|
+
c.notes.append("cross-check: matches the runtime's reported prompt token count" if reported == c.runtime_tokens
|
|
117
|
+
else f"cross-check FAILED: runtime reported {reported} prompt tokens, render has {c.runtime_tokens}")
|
|
118
|
+
c.notes = [x for x in c.notes if not x.startswith("__")]
|
|
119
|
+
return c
|
ftgate/conversations.py
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
"""Conversations to compare: built-in fixtures, or rows from the user's own dataset.
|
|
2
|
+
|
|
3
|
+
The built-in set covers the shapes a chat template has to get right —
|
|
4
|
+
plain chat, a system prompt, a tool call with a response, several tools.
|
|
5
|
+
The user's dataset is the one that matters: those are the rows the model
|
|
6
|
+
was actually trained on, so a mismatch there is a mismatch that counts.
|
|
7
|
+
"""
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import json
|
|
11
|
+
from dataclasses import dataclass, field
|
|
12
|
+
from pathlib import Path
|
|
13
|
+
|
|
14
|
+
WEATHER = {"type": "function", "function": {"name": "get_weather", "description": "Get the current weather for a city.",
|
|
15
|
+
"parameters": {"type": "object", "properties": {"city": {"type": "string", "description": "City name"},
|
|
16
|
+
"unit": {"type": "string", "enum": ["celsius", "fahrenheit"]}}, "required": ["city"]}}}
|
|
17
|
+
CONVERT = {"type": "function", "function": {"name": "convert_currency", "description": "Convert an amount between currencies.",
|
|
18
|
+
"parameters": {"type": "object", "properties": {"amount": {"type": "number"}, "from_currency": {"type": "string"},
|
|
19
|
+
"to_currency": {"type": "string"}}, "required": ["amount", "from_currency", "to_currency"]}}}
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
@dataclass
|
|
23
|
+
class Conversation:
|
|
24
|
+
name: str
|
|
25
|
+
messages: list[dict]
|
|
26
|
+
tools: list[dict] = field(default_factory=list)
|
|
27
|
+
|
|
28
|
+
@property
|
|
29
|
+
def inference_messages(self) -> list[dict]:
|
|
30
|
+
"""Everything up to the final assistant turn — what a runtime is sent."""
|
|
31
|
+
msgs = self.messages
|
|
32
|
+
if msgs and msgs[-1].get("role") == "assistant":
|
|
33
|
+
msgs = msgs[:-1]
|
|
34
|
+
return msgs
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
BUILTIN = [
|
|
38
|
+
Conversation("plain-chat", [
|
|
39
|
+
{"role": "user", "content": "Name three primary colours."},
|
|
40
|
+
{"role": "assistant", "content": "Red, blue and yellow."},
|
|
41
|
+
]),
|
|
42
|
+
Conversation("system-prompt", [
|
|
43
|
+
{"role": "system", "content": "You are a terse assistant. Answer in one sentence."},
|
|
44
|
+
{"role": "user", "content": "Why is the sky blue?"},
|
|
45
|
+
{"role": "assistant", "content": "Air scatters short wavelengths more than long ones."},
|
|
46
|
+
]),
|
|
47
|
+
Conversation("tool-call-and-response", [
|
|
48
|
+
{"role": "system", "content": "You are a helpful assistant."},
|
|
49
|
+
{"role": "user", "content": "What's the weather in Tbilisi?"},
|
|
50
|
+
{"role": "assistant", "content": "", "tool_calls": [{"type": "function", "function": {"name": "get_weather", "arguments": {"city": "Tbilisi", "unit": "celsius"}}}]},
|
|
51
|
+
{"role": "tool", "content": "{\"temp_c\": 21, \"sky\": \"clear\"}"},
|
|
52
|
+
{"role": "assistant", "content": "It's 21°C and clear in Tbilisi."},
|
|
53
|
+
], tools=[WEATHER]),
|
|
54
|
+
Conversation("two-tools-three-params", [
|
|
55
|
+
{"role": "user", "content": "Convert 100 USD to EUR."},
|
|
56
|
+
{"role": "assistant", "content": "", "tool_calls": [{"type": "function", "function": {"name": "convert_currency", "arguments": {"amount": 100, "from_currency": "USD", "to_currency": "EUR"}}}]},
|
|
57
|
+
], tools=[WEATHER, CONVERT]),
|
|
58
|
+
]
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def load_dataset(path: Path, limit: int | None = None) -> list[Conversation]:
|
|
62
|
+
"""Rows from a JSONL file in the common `messages` (+ optional `tools`) shape.
|
|
63
|
+
|
|
64
|
+
ShareGPT-style `conversations` with `from`/`value` is mapped onto it; any
|
|
65
|
+
other shape is reported rather than guessed at.
|
|
66
|
+
"""
|
|
67
|
+
out: list[Conversation] = []
|
|
68
|
+
with path.open(encoding="utf-8") as fh:
|
|
69
|
+
for n, line in enumerate(fh):
|
|
70
|
+
line = line.strip()
|
|
71
|
+
if not line:
|
|
72
|
+
continue
|
|
73
|
+
row = json.loads(line)
|
|
74
|
+
if "messages" in row:
|
|
75
|
+
msgs = row["messages"]
|
|
76
|
+
elif "conversations" in row:
|
|
77
|
+
role = {"human": "user", "gpt": "assistant", "system": "system", "tool": "tool"}
|
|
78
|
+
msgs = [{"role": role.get(m.get("from"), m.get("from")), "content": m.get("value", "")} for m in row["conversations"]]
|
|
79
|
+
else:
|
|
80
|
+
raise ValueError(f"{path}:{n+1}: no `messages` or `conversations` key; keys are {sorted(row)}")
|
|
81
|
+
out.append(Conversation(f"{path.name}:{n+1}", msgs, row.get("tools") or []))
|
|
82
|
+
if limit and len(out) >= limit:
|
|
83
|
+
break
|
|
84
|
+
return out
|