- 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
224 lines
6.4 KiB
Python
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())
|