Files
meeting-lab/scripts/run_gold_test.py
admin f7ad9ba51f Establish prompt engineering baseline with Gold Standard tests
- introduce Gold Standard evaluation corpus
- document decision taxonomy
- define prompt-engineering methodology
- add regression workflow
- establish Prompt Version 2 baseline
- validate decision_simple, decision_deferred and decision_none
2026-07-30 12:13:10 +02:00

224 lines
6.4 KiB
Python

#!/usr/bin/env python3
"""Run one gold-corpus scenario through the existing extraction flow."""
from __future__ import annotations
import argparse
import json
import sys
import time
from pathlib import Path
from typing import Any
import requests
REPO_ROOT = Path(__file__).resolve().parents[1]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
from src.meeting_lab.extraction.extract_chunks import (
DEFAULT_ENDPOINT,
EXTRACTION_CATEGORIES,
build_prompt,
call_ollama,
normalize_current_schema,
parse_json_response,
)
REQUIRED_KEYS = set(EXTRACTION_CATEGORIES)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Run one gold-standard transcript through extraction."
)
parser.add_argument("scenario", type=Path, help="Gold scenario directory")
parser.add_argument("--model", required=True, help="Ollama model name")
parser.add_argument(
"--endpoint",
default=DEFAULT_ENDPOINT,
help=f"Ollama generate endpoint (default: {DEFAULT_ENDPOINT})",
)
parser.add_argument(
"--timeout",
type=int,
default=1800,
help="HTTP timeout in seconds (default: 1800)",
)
parser.add_argument(
"--temperature",
type=float,
default=0.0,
help="Sampling temperature (default: 0.0)",
)
parser.add_argument(
"--num-predict",
type=int,
default=8192,
help="Maximum generated tokens (default: 8192)",
)
parser.add_argument(
"--num-ctx",
type=int,
default=32768,
help="Context window tokens (default: 32768)",
)
return parser.parse_args()
def scenario_paths(scenario_dir: Path) -> tuple[Path, Path, Path]:
if not scenario_dir.is_dir():
raise FileNotFoundError(f"Scenario directory not found: {scenario_dir}")
transcript_path = scenario_dir / "transcript.txt"
expected_path = scenario_dir / "expected.json"
actual_path = scenario_dir / "actual.json"
if not transcript_path.is_file():
raise FileNotFoundError(f"Missing transcript.txt: {transcript_path}")
if not expected_path.is_file():
raise FileNotFoundError(f"Missing expected.json: {expected_path}")
return transcript_path, expected_path, actual_path
def read_json_object(path: Path) -> dict[str, Any]:
try:
data = json.loads(path.read_text(encoding="utf-8"))
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid JSON in {path}: {exc}") from exc
if not isinstance(data, dict):
raise ValueError(f"JSON file must contain an object: {path}")
return data
def validate_required_keys(data: dict[str, Any], path: Path) -> None:
missing = sorted(REQUIRED_KEYS - set(data))
if missing:
raise ValueError(f"Missing required keys in {path}: {', '.join(missing)}")
def format_items(items: Any) -> list[str]:
if not isinstance(items, list):
return [f"<invalid non-list value: {items!r}>"]
if not items:
return ["<none>"]
return [str(item) for item in items]
def print_decision_comparison(
scenario_dir: Path,
model: str,
expected: dict[str, Any],
actual: dict[str, Any],
runtime: float,
) -> None:
expected_decisions = expected.get("decisions", [])
actual_decisions = actual.get("decisions", [])
print(f"Scenario: {scenario_dir}")
print(f"Model: {model}")
print(f"Expected decision count: {len(expected_decisions)}")
print(f"Actual decision count: {len(actual_decisions)}")
print("Expected decisions:")
for item in format_items(expected_decisions):
print(f"- {item}")
print("Actual decisions:")
for item in format_items(actual_decisions):
print(f"- {item}")
print(f"Runtime: {runtime:.2f}s")
def run_gold_test(
scenario_dir: Path,
model: str,
endpoint: str,
timeout: int,
temperature: float,
num_predict: int | None,
num_ctx: int | None,
) -> Path:
transcript_path, expected_path, actual_path = scenario_paths(scenario_dir)
expected = read_json_object(expected_path)
validate_required_keys(expected, expected_path)
transcript = transcript_path.read_text(encoding="utf-8-sig").strip()
if not transcript:
raise ValueError(f"The transcript is empty: {transcript_path}")
prompt = build_prompt(transcript_path.name, transcript)
started = time.perf_counter()
raw_text, _metadata = call_ollama(
endpoint=endpoint,
model=model,
prompt=prompt,
timeout=timeout,
temperature=temperature,
num_predict=num_predict,
num_ctx=num_ctx,
)
try:
parsed = parse_json_response(raw_text)
except (json.JSONDecodeError, ValueError) as exc:
raw_path = actual_path.with_suffix(".raw.txt")
raw_path.write_text(raw_text + "\n", encoding="utf-8")
raise ValueError(f"Model output was not valid JSON. Raw output: {raw_path}") from exc
actual = normalize_current_schema(parsed)
actual_path.write_text(
json.dumps(actual, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
actual_from_disk = read_json_object(actual_path)
validate_required_keys(actual_from_disk, actual_path)
runtime = time.perf_counter() - started
print_decision_comparison(
scenario_dir=scenario_dir,
model=model,
expected=expected,
actual=actual_from_disk,
runtime=runtime,
)
return actual_path
def main() -> int:
args = parse_args()
try:
actual_path = run_gold_test(
scenario_dir=args.scenario,
model=args.model,
endpoint=args.endpoint,
timeout=args.timeout,
temperature=args.temperature,
num_predict=args.num_predict,
num_ctx=args.num_ctx,
)
except requests.ConnectionError:
print(
"Error: Ollama is not reachable. Is `ollama serve` running?",
file=sys.stderr,
)
return 1
except requests.Timeout:
print("Error: The Ollama request timed out.", file=sys.stderr)
return 1
except requests.HTTPError as exc:
print(f"Error: Ollama returned an HTTP error: {exc}", file=sys.stderr)
return 1
except (OSError, UnicodeError, ValueError) as exc:
print(f"Error: {exc}", file=sys.stderr)
return 1
print(f"Actual JSON: {actual_path}")
return 0
if __name__ == "__main__":
raise SystemExit(main())