TenaOS / training_code /evaluate_lora_ab.py
beza4588's picture
Add synthetic LoRA training corpus and scripts
cc6ba10 verified
Raw History Blame Contribute Delete
22.8 kB
#!/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)