#!/usr/bin/env python3 """Compare base and merged LoRA model outputs on held-out TenaOS tasks. The script is intentionally sidecar-friendly: it never talks to, restarts, or mutates the running TenaOS demo container. It can either call OpenAI-compatible HTTP endpoints, call ``llama-cli`` directly, or score previously generated prediction JSONL files. """ from __future__ import annotations import argparse import json import math import re import subprocess import sys import time import urllib.error import urllib.request from collections import Counter, defaultdict from dataclasses import dataclass from datetime import datetime, timezone from pathlib import Path from typing import Any DEFAULT_TEST_JSONL = Path("lora_training/artifacts/sft/test.jsonl") DEFAULT_OUT_DIR = Path("lora_training/artifacts/ab_eval") TASK_LABELS = { "cds": "Clinical decision support", "form": "Form builder", "patient_education": "Patient education", "report": "Report builder", "scribe_text_amharic": "Amharic text scribe", "scribe_text_english": "English text scribe", "voice_scribe_audio": "Voice scribe", } CDS_HEADINGS = ( "## Clinical Assessment", "## Evidence-Based Considerations", "## Suggested Actions", "## Safety Alerts", "## Key Points", ) EDU_HEADINGS = ( "## What You Have", "## Why It Matters", "## What To Do", "## Your Medications", "## What to Avoid", "## Follow-Up Schedule", "## When To Seek Help", ) SOAP_KEYS = ("subjective", "objective", "assessment", "plan") @dataclass(frozen=True) class Example: id: str kind: str task_tag: str prompt: str reference: str request: dict[str, Any] reference_json: dict[str, Any] | None @dataclass(frozen=True) class ModelSpec: name: str endpoint: str | None = None model: str | None = None llama_cli: Path | None = None gguf: Path | None = None mmproj: Path | None = None def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--test-jsonl", type=Path, default=DEFAULT_TEST_JSONL) parser.add_argument("--out-dir", type=Path, default=DEFAULT_OUT_DIR) parser.add_argument("--limit-per-kind", type=int, default=2) parser.add_argument("--kinds", nargs="*", default=sorted(TASK_LABELS)) parser.add_argument("--max-tokens", type=int, default=1536) parser.add_argument("--temperature", type=float, default=0.0) parser.add_argument("--timeout-seconds", type=int, default=900) parser.add_argument("--ctx-size", type=int, default=8192) parser.add_argument("--base-endpoint", help="OpenAI-compatible /v1/chat/completions base URL.") parser.add_argument("--base-model", default="base") parser.add_argument("--lora-endpoint", help="OpenAI-compatible /v1/chat/completions base URL.") parser.add_argument("--lora-model", default="lora") parser.add_argument("--llama-cli", type=Path, help="Path to llama-cli for direct GGUF inference.") parser.add_argument("--base-gguf", type=Path, help="Base GGUF path for direct llama-cli inference.") parser.add_argument("--lora-gguf", type=Path, help="Merged LoRA GGUF path for direct llama-cli inference.") parser.add_argument("--mmproj", type=Path, help="Optional multimodal projector for llama-cli.") parser.add_argument("--base-predictions", type=Path, help="Existing base prediction JSONL to score.") parser.add_argument("--lora-predictions", type=Path, help="Existing LoRA prediction JSONL to score.") parser.add_argument("--dry-run", action="store_true", help="Only sample examples and write the eval plan.") return parser.parse_args() 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. Check --test-jsonl, --kinds, and --limit-per-kind.") plan_path = args.out_dir / "eval_plan.json" write_json( plan_path, { "schema_version": "tenaos_lora_ab_eval_plan_v1", "created_at": now(), "test_jsonl": str(args.test_jsonl), "limit_per_kind": args.limit_per_kind, "selected_counts": dict(sorted(Counter(example.kind for example in examples).items())), "examples": [{"id": e.id, "kind": e.kind, "task_tag": e.task_tag} for e in examples], }, ) if args.dry_run: print(f"Wrote dry-run eval plan: {plan_path}") return if args.base_predictions and args.lora_predictions: base_results = score_prediction_file(args.base_predictions, examples, "base") lora_results = score_prediction_file(args.lora_predictions, examples, "lora") else: base_spec, lora_spec = build_model_specs(args) base_results = run_model(base_spec, examples, args, args.out_dir / "base_predictions.jsonl") lora_results = run_model(lora_spec, examples, args, args.out_dir / "lora_predictions.jsonl") summary = summarize(base_results, lora_results) summary_path = args.out_dir / "summary.json" write_json(summary_path, summary) print(json.dumps(summary, indent=2, ensure_ascii=False, sort_keys=True)) print(f"Wrote summary: {summary_path}") 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 handle: for line in handle: 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 conversations = raw.get("conversations") if isinstance(raw.get("conversations"), list) else [] if len(conversations) < 2: continue prompt = str(conversations[0].get("content") or "") reference = str(conversations[1].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 ""), prompt=prompt, reference=reference, request=parse_prompt_request(prompt), reference_json=extract_json_object(reference), ) ) counts[kind] += 1 if wanted and all(counts[kind] >= limit_per_kind for kind in wanted): break return examples def build_model_specs(args: argparse.Namespace) -> tuple[ModelSpec, ModelSpec]: if args.base_endpoint and args.lora_endpoint: return ( ModelSpec("base", endpoint=args.base_endpoint.rstrip("/"), model=args.base_model), ModelSpec("lora", endpoint=args.lora_endpoint.rstrip("/"), model=args.lora_model), ) if args.llama_cli and args.base_gguf and args.lora_gguf: return ( ModelSpec("base", llama_cli=args.llama_cli, gguf=args.base_gguf, mmproj=args.mmproj), ModelSpec("lora", llama_cli=args.llama_cli, gguf=args.lora_gguf, mmproj=args.mmproj), ) raise SystemExit( "Provide either --base-endpoint/--lora-endpoint, --llama-cli with both GGUF paths, " "or --base-predictions/--lora-predictions." ) def run_model( spec: ModelSpec, examples: list[Example], args: argparse.Namespace, predictions_path: Path, ) -> list[dict[str, Any]]: results: list[dict[str, Any]] = [] with predictions_path.open("w", encoding="utf-8") as handle: for index, example in enumerate(examples, 1): started = time.time() error = "" try: output = generate(spec, example.prompt, args) except Exception as exc: # noqa: BLE001 - the eval should record failures and continue. output = "" error = f"{type(exc).__name__}: {exc}" elapsed = time.time() - started result = score_output(example, output, spec.name, error=error, elapsed_seconds=elapsed) handle.write(json.dumps(result, ensure_ascii=False, sort_keys=True) + "\n") handle.flush() print(f"[{spec.name}] {index}/{len(examples)} {example.kind}/{example.id}: {result['metrics']['total']:.3f}") results.append(result) return results def generate(spec: ModelSpec, prompt: str, args: argparse.Namespace) -> str: if spec.endpoint: return generate_http(spec, prompt, args) if spec.llama_cli and spec.gguf: return generate_llama_cli(spec, prompt, args) raise RuntimeError(f"Model spec {spec.name!r} has no runnable backend.") def generate_http(spec: ModelSpec, prompt: str, args: argparse.Namespace) -> str: payload = { "model": spec.model or spec.name, "messages": [{"role": "user", "content": prompt}], "temperature": args.temperature, "max_tokens": args.max_tokens, } request = urllib.request.Request( f"{spec.endpoint}/v1/chat/completions", data=json.dumps(payload).encode("utf-8"), headers={"Content-Type": "application/json"}, method="POST", ) try: with urllib.request.urlopen(request, timeout=args.timeout_seconds) as response: data = json.loads(response.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[:1000]}") from exc return str(data["choices"][0]["message"]["content"]) def generate_llama_cli(spec: ModelSpec, prompt: str, args: argparse.Namespace) -> str: command = [ str(spec.llama_cli), "-m", str(spec.gguf), "-p", prompt, "-n", str(args.max_tokens), "--ctx-size", str(args.ctx_size), "--temp", str(args.temperature), "--no-display-prompt", ] if spec.mmproj: command.extend(["--mmproj", str(spec.mmproj)]) completed = subprocess.run( command, check=False, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, timeout=args.timeout_seconds, ) if completed.returncode: raise RuntimeError(completed.stderr.strip()[:2000]) return completed.stdout.strip() def score_prediction_file(path: Path, examples: list[Example], model_name: str) -> list[dict[str, Any]]: by_id = {example.id: example for example in examples} results: list[dict[str, Any]] = [] with path.open("r", encoding="utf-8") as handle: for line in handle: if not line.strip(): continue raw = json.loads(line) example_id = str(raw.get("id") or raw.get("example_id") or "") if example_id not in by_id: continue output = str(raw.get("output") or raw.get("completion") or raw.get("response") or "") results.append(score_output(by_id[example_id], output, model_name)) return results def score_output( example: Example, output: str, model_name: str, *, error: str = "", elapsed_seconds: float | None = None, ) -> dict[str, Any]: parsed = extract_json_object(output) metrics = score_metrics(example, output, parsed) if error: metrics = {**metrics, "total": 0.0, "inference_ok": 0.0} return { "schema_version": "tenaos_lora_ab_prediction_v1", "created_at": now(), "model": model_name, "id": example.id, "kind": example.kind, "task_tag": example.task_tag, "error": error, "elapsed_seconds": elapsed_seconds, "metrics": metrics, "output": output, } def score_metrics(example: Example, output: str, parsed: dict[str, Any] | None) -> dict[str, float]: reference = example.reference_json or {} output_text = json.dumps(parsed, ensure_ascii=False, sort_keys=True) if parsed else output reference_text = json.dumps(reference, ensure_ascii=False, sort_keys=True) if reference else example.reference metrics: dict[str, float] = { "inference_ok": 1.0, "valid_json": 1.0 if parsed else 0.0, "schema_match": schema_match(reference, parsed), "top_level_key_f1": key_f1(reference, parsed), "token_f1": token_f1(reference_text, output_text), } metrics.update(task_metrics(example, parsed, output_text)) component_keys = [key for key in metrics if key not in {"total", "inference_ok"}] metrics["total"] = sum(metrics[key] for key in component_keys) / max(1, len(component_keys)) return {key: round(value, 6) for key, value in metrics.items()} def task_metrics(example: Example, parsed: dict[str, Any] | None, output_text: str) -> dict[str, float]: if example.kind == "cds": content = nested_text(parsed, ("structured_cds", "content")) return { "required_section_recall": phrase_recall(CDS_HEADINGS, content or output_text), "content_length_ok": 1.0 if len(content) >= 1200 else 0.0, "request_anchor_recall": request_anchor_recall(example.request, output_text, ("case_id",)), } if example.kind == "patient_education": content = nested_text(parsed, ("material", "content")) return { "required_section_recall": phrase_recall(EDU_HEADINGS, content or output_text), "content_length_ok": 1.0 if len(content) >= 1600 else 0.0, "request_anchor_recall": request_anchor_recall(example.request, output_text, ("case_id",)), } if example.kind == "report": metadata = example.request.get("metadata") if isinstance(example.request.get("metadata"), dict) else {} expected = list(metadata.get("expected_filters") or []) + list(metadata.get("expected_group_by") or []) if metadata.get("report_type"): expected.append(str(metadata["report_type"])) if metadata.get("date_range"): expected.append(str(metadata["date_range"])) return { "expected_metadata_recall": phrase_recall(expected, output_text), "has_draft": 1.0 if parsed and isinstance(parsed.get("draft"), dict) else 0.0, "has_summary": 1.0 if parsed and isinstance(parsed.get("summary"), dict) else 0.0, } if example.kind == "form": metadata = example.request.get("metadata") if isinstance(example.request.get("metadata"), dict) else {} expected = list(metadata.get("expected_sections") or []) return { "expected_section_recall": phrase_recall(expected, output_text), "has_draft": 1.0 if parsed and isinstance(parsed.get("draft"), dict) else 0.0, "has_summary": 1.0 if parsed and isinstance(parsed.get("summary"), dict) else 0.0, } if example.kind in {"scribe_text_english", "scribe_text_amharic", "voice_scribe_audio"}: expected = example.request.get("expected") if isinstance(example.request.get("expected"), dict) else {} soap = find_soap(parsed) return { "soap_completeness": sum(1 for key in SOAP_KEYS if str(soap.get(key) or "").strip()) / len(SOAP_KEYS), "expected_extraction_recall": expected_extraction_recall(expected, output_text), "forbidden_extra_avoidance": forbidden_extra_avoidance(expected, output_text), } return {} def parse_prompt_request(prompt: str) -> dict[str, Any]: start = prompt.find("{") if start < 0: return {} parsed = extract_json_object(prompt[start:]) return parsed or {} def extract_json_object(text: str) -> dict[str, Any] | None: decoder = json.JSONDecoder() for match in re.finditer(r"\{", text): try: parsed, _ = decoder.raw_decode(text[match.start() :]) except json.JSONDecodeError: continue if isinstance(parsed, dict): return parsed return None def schema_match(reference: dict[str, Any], parsed: dict[str, Any] | None) -> float: if not reference or not parsed: return 0.0 expected = reference.get("schema_version") if not expected: return 1.0 return 1.0 if parsed.get("schema_version") == expected else 0.0 def key_f1(reference: dict[str, Any], parsed: dict[str, Any] | None) -> float: if not reference or not parsed: return 0.0 expected = set(reference) actual = set(parsed) return f1(len(expected & actual), len(actual - expected), len(expected - actual)) def token_f1(expected: str, actual: str) -> float: expected_tokens = Counter(tokens(expected)) actual_tokens = Counter(tokens(actual)) if not expected_tokens or not actual_tokens: return 0.0 overlap = sum((expected_tokens & actual_tokens).values()) precision = overlap / sum(actual_tokens.values()) recall = overlap / sum(expected_tokens.values()) return harmonic(precision, recall) def expected_extraction_recall(expected: dict[str, Any], output_text: str) -> float: targets: list[str] = [] for group in ("concepts", "observations", "medications"): for item in expected.get(group) or []: if not isinstance(item, dict): continue for key in ("label", "value", "dose", "drug", "name"): value = str(item.get(key) or "").strip() if value: targets.append(value) break return phrase_recall(targets, output_text) def forbidden_extra_avoidance(expected: dict[str, Any], output_text: str) -> float: forbidden = expected.get("forbiddenExtractions") or [] if not forbidden: return 1.0 lowered = normalize(output_text) hits = 0 for item in forbidden: phrase = item if isinstance(item, str) else json.dumps(item, ensure_ascii=False) if normalize(str(phrase)) in lowered: hits += 1 return 1.0 - (hits / len(forbidden)) def request_anchor_recall(request: dict[str, Any], output_text: str, keys: tuple[str, ...]) -> float: anchors = [str(request[key]) for key in keys if request.get(key)] return phrase_recall(anchors, output_text) def phrase_recall(phrases: list[str] | tuple[str, ...], text: str) -> float: cleaned = [normalize(phrase) for phrase in phrases if str(phrase).strip()] if not cleaned: return 1.0 lowered = normalize(text) return sum(1 for phrase in cleaned if phrase in lowered) / len(cleaned) def find_soap(parsed: dict[str, Any] | None) -> dict[str, Any]: if not parsed: return {} candidates = [ parsed.get("soap"), (parsed.get("result") or {}).get("soap") if isinstance(parsed.get("result"), dict) else None, ((parsed.get("audio_trace") or {}).get("result") or {}).get("soap") if isinstance(parsed.get("audio_trace"), dict) and isinstance((parsed.get("audio_trace") or {}).get("result"), dict) else None, ((parsed.get("amharic_trace") or {}).get("result") or {}).get("soap") if isinstance(parsed.get("amharic_trace"), dict) and isinstance((parsed.get("amharic_trace") or {}).get("result"), dict) else None, ] for candidate in candidates: if isinstance(candidate, dict): return candidate return {} def nested_text(parsed: dict[str, Any] | None, path: tuple[str, ...]) -> str: current: Any = parsed for key in path: if not isinstance(current, dict): return "" current = current.get(key) return str(current or "") def f1(tp: int, fp: int, fn: int) -> float: precision = tp / (tp + fp) if tp + fp else 0.0 recall = tp / (tp + fn) if tp + fn else 0.0 return harmonic(precision, recall) def harmonic(precision: float, recall: float) -> float: if precision + recall == 0: return 0.0 return 2 * precision * recall / (precision + recall) def tokens(text: str) -> list[str]: return re.findall(r"[a-z0-9_]+", normalize(text)) def normalize(text: str) -> str: return re.sub(r"\s+", " ", str(text).casefold()).strip() def summarize(base_results: list[dict[str, Any]], lora_results: list[dict[str, Any]]) -> dict[str, Any]: base_by_id = {str(result["id"]): result for result in base_results} lora_by_id = {str(result["id"]): result for result in lora_results} shared_ids = sorted(set(base_by_id) & set(lora_by_id)) by_kind: dict[str, dict[str, Any]] = {} wins = Counter() for example_id in shared_ids: base = base_by_id[example_id] lora = lora_by_id[example_id] base_total = float(base["metrics"]["total"]) lora_total = float(lora["metrics"]["total"]) if math.isclose(base_total, lora_total, abs_tol=1e-9): wins["tie"] += 1 elif lora_total > base_total: wins["lora"] += 1 else: wins["base"] += 1 for kind in sorted({result["kind"] for result in base_results + lora_results}): base_kind = [result for result in base_results if result["kind"] == kind] lora_kind = [result for result in lora_results if result["kind"] == kind] by_kind[kind] = { "label": TASK_LABELS.get(kind, kind), "base_count": len(base_kind), "lora_count": len(lora_kind), "base_avg_total": average_total(base_kind), "lora_avg_total": average_total(lora_kind), "delta_lora_minus_base": round(average_total(lora_kind) - average_total(base_kind), 6), } return { "schema_version": "tenaos_lora_ab_eval_summary_v1", "created_at": now(), "shared_example_count": len(shared_ids), "wins": dict(sorted(wins.items())), "by_kind": by_kind, "base_avg_total": average_total(base_results), "lora_avg_total": average_total(lora_results), "delta_lora_minus_base": round(average_total(lora_results) - average_total(base_results), 6), } def average_total(results: list[dict[str, Any]]) -> float: if not results: return 0.0 return round(sum(float(result["metrics"]["total"]) for result in results) / len(results), 6) def write_json(path: Path, data: dict[str, Any]) -> None: path.write_text(json.dumps(data, indent=2, ensure_ascii=False, sort_keys=True) + "\n", encoding="utf-8") def now() -> str: return datetime.now(timezone.utc).isoformat() if __name__ == "__main__": try: main() except KeyboardInterrupt: sys.exit(130)