#!/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 `{...}` 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 `{...}` 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"\s*(\{.*?\})\s*", 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()