TenaOS / training_code /evaluate_lora_ab_v2.py
beza4588's picture
Add synthetic LoRA training corpus and scripts
58a59f8 verified
Raw History Blame Contribute Delete
12.5 kB
#!/usr/bin/env python3
"""Base vs. LoRA A/B eval, schema-correct for the current multi-turn traces.
`evaluate_lora_ab.py` predates the two-system-message, multi-turn agentic
`<tena_call>{...}</tena_call>` trace format used by artifacts_v2 (verified
during the training-data/production alignment audit). Its `load_examples`
treats `conversations[0]` as the prompt and `conversations[1]` as the
reference completion -- but conversations[0]/[1] are both system messages in
the current schema, so that comparison is meaningless for this dataset.
This script instead:
1. For each example, finds the LAST assistant-authored turn (the target
completion the model was trained to imitate at that point).
2. Builds the prompt as the exact chat-message prefix before that turn
(all system/user/assistant turns, in order, verbatim) -- i.e. a
teacher-forced continuation, using the real tool-result turns already
recorded in the trace instead of executing tools live.
3. Sends that prefix to an OpenAI-compatible /v1/chat/completions endpoint
for both the base and LoRA model and compares each completion against
the real reference turn.
Scoring (per example):
- tena_call_present: did the completion contain a `<tena_call>{...}</tena_call>` wrapper at all
- tool_name_match: does the called tool name match the reference's
- arg_key_f1: key-set F1 between reference and completion tool-call arguments
- text_field_token_f1: token-level F1 between reference and completion on
the long-form text fields we actually care about ("content"/"summary"/
free text), falling back to whole-output token F1 if no tool call parsed
- json_valid: could the tool-call arguments be parsed as JSON at all
This is intentionally NOT a full agentic rollout (no live KB/tool execution)
-- it measures next-turn imitation quality on real held-out traces, which is
the right question for judging whether a checkpoint has degraded or improved
relative to base, without requiring a running TenaOS/KB stack.
"""
from __future__ import annotations
import argparse
import json
import re
import time
import urllib.error
import urllib.request
from collections import Counter, defaultdict
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
TENA_CALL_RE = re.compile(r"<tena_call>\s*(\{.*?\})\s*</tena_call>", re.DOTALL)
TEXT_FIELD_KEYS = ("content", "summary", "text")
@dataclass
class Example:
id: str
kind: str
task_tag: str
prefix_messages: list[dict[str, str]]
reference_text: str
reference_call: dict[str, Any] | None
@dataclass
class ModelResult:
id: str
kind: str
output: str
error: str
elapsed_seconds: float
metrics: dict[str, float] = field(default_factory=dict)
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--test-jsonl", type=Path, required=True)
p.add_argument("--out-dir", type=Path, required=True)
p.add_argument("--limit-per-kind", type=int, default=6)
p.add_argument("--kinds", nargs="*", default=["cds", "form", "patient_education", "report", "scribe"])
p.add_argument("--max-tokens", type=int, default=1024)
p.add_argument("--temperature", type=float, default=0.0)
p.add_argument("--timeout-seconds", type=int, default=300)
p.add_argument("--seed", type=int, default=0)
p.add_argument("--base-endpoint", required=True, help="OpenAI-compatible base URL, e.g. http://host:8090")
p.add_argument("--lora-endpoint", required=True, help="OpenAI-compatible base URL, e.g. http://host:8091")
p.add_argument("--base-model", default="base")
p.add_argument("--lora-model", default="lora")
return p.parse_args()
def load_examples(path: Path, kinds: list[str], limit_per_kind: int) -> list[Example]:
wanted = set(kinds)
counts: Counter[str] = Counter()
examples: list[Example] = []
with path.open("r", encoding="utf-8") as fh:
for line in fh:
if not line.strip():
continue
raw = json.loads(line)
kind = str(raw.get("kind") or "")
if kind not in wanted or counts[kind] >= limit_per_kind:
continue
convs = raw.get("conversations") or []
if not isinstance(convs, list) or len(convs) < 2:
continue
last_assistant_idx = None
for i in range(len(convs) - 1, -1, -1):
if convs[i].get("role") == "assistant":
last_assistant_idx = i
break
if last_assistant_idx is None or last_assistant_idx == 0:
continue
prefix = [
{"role": m.get("role", "user"), "content": str(m.get("content") or "")}
for m in convs[:last_assistant_idx]
]
reference_text = str(convs[last_assistant_idx].get("content") or "")
examples.append(
Example(
id=str(raw.get("id") or f"{kind}_{counts[kind] + 1}"),
kind=kind,
task_tag=str(raw.get("task_tag") or ""),
prefix_messages=prefix,
reference_text=reference_text,
reference_call=extract_tena_call(reference_text),
)
)
counts[kind] += 1
if wanted and all(counts[k] >= limit_per_kind for k in wanted):
break
return examples
def extract_tena_call(text: str) -> dict[str, Any] | None:
m = TENA_CALL_RE.search(text)
if not m:
return None
try:
return json.loads(m.group(1))
except json.JSONDecodeError:
return None
def generate(endpoint: str, model: str, messages: list[dict[str, str]], args: argparse.Namespace) -> str:
payload = {
"model": model,
"messages": messages,
"temperature": args.temperature,
"max_tokens": args.max_tokens,
"seed": args.seed,
}
req = urllib.request.Request(
f"{endpoint.rstrip('/')}/v1/chat/completions",
data=json.dumps(payload).encode("utf-8"),
headers={"Content-Type": "application/json"},
method="POST",
)
try:
with urllib.request.urlopen(req, timeout=args.timeout_seconds) as resp:
data = json.loads(resp.read().decode("utf-8"))
except urllib.error.HTTPError as exc:
body = exc.read().decode("utf-8", errors="replace")
raise RuntimeError(f"HTTP {exc.code}: {body[:800]}") from exc
return str(data["choices"][0]["message"]["content"])
def tokens(text: str) -> list[str]:
return re.findall(r"[a-z0-9_]+", text.lower())
def token_f1(expected: str, actual: str) -> float:
exp_c = Counter(tokens(expected))
act_c = Counter(tokens(actual))
if not exp_c or not act_c:
return 0.0
overlap = sum((exp_c & act_c).values())
precision = overlap / sum(act_c.values())
recall = overlap / sum(exp_c.values())
if precision + recall == 0:
return 0.0
return 2 * precision * recall / (precision + recall)
def key_f1(expected: dict[str, Any], actual: dict[str, Any]) -> float:
exp_keys, act_keys = set(expected), set(actual)
if not exp_keys and not act_keys:
return 1.0
if not exp_keys or not act_keys:
return 0.0
tp = len(exp_keys & act_keys)
precision = tp / len(act_keys)
recall = tp / len(exp_keys)
if precision + recall == 0:
return 0.0
return 2 * precision * recall / (precision + recall)
def score(example: Example, output: str, error: str) -> dict[str, float]:
if error:
return {"tena_call_present": 0.0, "tool_name_match": 0.0, "arg_key_f1": 0.0, "text_field_token_f1": 0.0, "json_valid": 0.0, "error": 1.0}
call = extract_tena_call(output)
ref_call = example.reference_call
metrics: dict[str, float] = {"error": 0.0}
metrics["tena_call_present"] = 1.0 if call is not None else 0.0
metrics["json_valid"] = 1.0 if call is not None else 0.0
if ref_call is None:
# Reference itself has no tool call (e.g. a plain tool-result echo) --
# only meaningful signal left is raw text overlap.
metrics["tool_name_match"] = 1.0 if call is None else 0.0
metrics["arg_key_f1"] = 1.0 if call is None else 0.0
metrics["text_field_token_f1"] = token_f1(example.reference_text, output)
return metrics
ref_name = str(ref_call.get("name") or "")
act_name = str(call.get("name") or "") if call else ""
metrics["tool_name_match"] = 1.0 if call is not None and act_name == ref_name else 0.0
ref_args = ref_call.get("arguments") if isinstance(ref_call.get("arguments"), dict) else {}
act_args = (call.get("arguments") if call and isinstance(call.get("arguments"), dict) else {}) or {}
metrics["arg_key_f1"] = key_f1(ref_args, act_args)
ref_text_parts = [str(ref_args.get(k) or "") for k in TEXT_FIELD_KEYS if ref_args.get(k)]
act_text_parts = [str(act_args.get(k) or "") for k in TEXT_FIELD_KEYS if act_args.get(k)]
if ref_text_parts:
metrics["text_field_token_f1"] = token_f1(" ".join(ref_text_parts), " ".join(act_text_parts) or output)
else:
metrics["text_field_token_f1"] = token_f1(example.reference_text, output)
return metrics
def run_model(name: str, endpoint: str, model: str, examples: list[Example], args: argparse.Namespace, out_path: Path) -> list[ModelResult]:
results: list[ModelResult] = []
with out_path.open("w", encoding="utf-8") as fh:
for i, ex in enumerate(examples, 1):
started = time.time()
error = ""
output = ""
try:
output = generate(endpoint, model, ex.prefix_messages, args)
except Exception as exc: # noqa: BLE001
error = f"{type(exc).__name__}: {exc}"
elapsed = time.time() - started
metrics = score(ex, output, error)
result = ModelResult(id=ex.id, kind=ex.kind, output=output, error=error, elapsed_seconds=elapsed, metrics=metrics)
results.append(result)
fh.write(json.dumps({
"id": result.id, "kind": result.kind, "model": name, "error": result.error,
"elapsed_seconds": round(result.elapsed_seconds, 2), "metrics": result.metrics, "output": result.output,
}, ensure_ascii=False) + "\n")
fh.flush()
print(f"[{name}] {i}/{len(examples)} {ex.kind}/{ex.id}: "
f"tena_call={metrics['tena_call_present']:.0f} tool_match={metrics['tool_name_match']:.2f} "
f"arg_f1={metrics['arg_key_f1']:.2f} text_f1={metrics['text_field_token_f1']:.2f} err={bool(error)}")
return results
def summarize(name: str, results: list[ModelResult]) -> dict[str, Any]:
by_kind: dict[str, list[ModelResult]] = defaultdict(list)
for r in results:
by_kind[r.kind].append(r)
metric_keys = ["tena_call_present", "tool_name_match", "arg_key_f1", "text_field_token_f1", "json_valid", "error"]
def agg(rs: list[ModelResult]) -> dict[str, float]:
return {k: round(sum(r.metrics.get(k, 0.0) for r in rs) / max(1, len(rs)), 4) for k in metric_keys}
return {
"model": name,
"n": len(results),
"overall": agg(results),
"by_kind": {k: agg(v) for k, v in sorted(by_kind.items())},
"mean_elapsed_seconds": round(sum(r.elapsed_seconds for r in results) / max(1, len(results)), 2),
}
def main() -> None:
args = parse_args()
args.out_dir.mkdir(parents=True, exist_ok=True)
examples = load_examples(args.test_jsonl, args.kinds, args.limit_per_kind)
if not examples:
raise SystemExit("No examples selected.")
print(f"Selected {len(examples)} examples: {dict(sorted(Counter(e.kind for e in examples).items()))}")
base_results = run_model("base", args.base_endpoint, args.base_model, examples, args, args.out_dir / "base_predictions.jsonl")
lora_results = run_model("lora", args.lora_endpoint, args.lora_model, examples, args, args.out_dir / "lora_predictions.jsonl")
summary = {"base": summarize("base", base_results), "lora": summarize("lora", lora_results)}
(args.out_dir / "summary.json").write_text(json.dumps(summary, indent=2, ensure_ascii=False), encoding="utf-8")
print(json.dumps(summary, indent=2, ensure_ascii=False))
if __name__ == "__main__":
main()