atc-parser-scripts / eval_baseline.py
rudeparis's picture
Upload eval_baseline.py with huggingface_hub
6dcffbd verified
Raw
History Blame Contribute Delete
5.5 kB
"""Eval LLM parser against canonical 134 scenarios × 553 turns from `output/scenarios/*.json`.
Apples-to-apples vs production V32 baseline (`just baseline` reports).
Walks the same gold the harness consumes, measures intent_em + slots_em + per-intent breakdown.
V32 production reference (from `just baseline`):
intent: 87.4% slots: 80.6% readback: 76.3% compound_seg: 51.3% pass: 76.3%
on 134 scenarios / 553 turns / 253 ATC-to-ownship
Usage:
uv run python -u poc/llm-finetune/training/eval_baseline.py \\
--adapter poc/llm-finetune/training/q3b_adapter_v7_r16
"""
from __future__ import annotations
import argparse
import json
import re
from collections import Counter, defaultdict
from pathlib import Path
from mlx_lm import load, generate
INTENTS = sorted(json.loads(Path("/Users/jean-patricksmith/digital/kingly/apps/production/naac/poc/deberta_intent/checkpoints-base/label_mapping.json").read_text())["intent2id"].keys())
SLOTS = ["ALTIMETER_SETTING", "ALTITUDE", "APPROACH_TYPE", "CALL_SIGN", "CLOCK_POSITION",
"DIRECTION", "DISTANCE", "FACILITY", "FIX", "FREQUENCY", "HEADING", "PATTERN_LEG",
"ROUTE", "RUNWAY", "SPEED", "TAXIWAY", "TIME", "TRANSPONDER_CODE", "TURN_DIRECTION"]
SYSTEM_PROMPT = (
"You are an ATC parser. Parse the air traffic control transmission into a JSON object.\n"
"OUTPUT FORMAT: {\"intent\": <one-of-enum>, \"slots\": {<SLOT_TYPE>: <value>}}\n"
f"\nINTENT MUST BE ONE OF: {', '.join(INTENTS)}\n"
f"\nSLOT TYPES (UPPERCASE): {', '.join(SLOTS)}\n"
"\nOutput ONLY the JSON object. No prose, no code fences."
)
SCENARIO_DIR = Path("/Users/jean-patricksmith/digital/kingly/apps/production/naac/data/src/partner_graph/output/scenarios")
def parse(text: str):
text = text.strip()
text = re.sub(r'<think>.*?</think>\s*', '', text, flags=re.DOTALL).strip()
if text.startswith("```"):
text = text.split("```")[1]
if text.startswith("json"):
text = text[4:]
try:
return json.loads(text.strip())
except Exception:
return None
def walk_scenarios(only_ownship: bool = True):
"""Yield (scenario_id, turn_idx, atc_text, expected_intent, expected_params) per turn.
Schema: top-level `turns[]` with {speaker, target, expected_transcript, expected_intent, expected_parameters}.
ATC-to-ownship subset matches `just baseline`'s 253 scored turns when only_ownship=True.
"""
for path in sorted(SCENARIO_DIR.glob("*.json")):
try:
scn = json.loads(path.read_text())
except Exception:
continue
scn_id = scn.get("scenario_id", scn.get("id", path.stem))
for turn in scn.get("turns", []):
speaker = turn.get("speaker", "").lower()
target = turn.get("target", "").lower()
if speaker != "atc":
continue
if only_ownship and target != "ownship":
continue
text = turn.get("expected_transcript", "").strip()
exp_intent = turn.get("expected_intent")
exp_params = turn.get("expected_parameters") or {}
if not text or not exp_intent:
continue
yield scn_id, turn.get("turn_index", 0), text, exp_intent, exp_params
def main():
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--adapter", default="poc/llm-finetune/training/q3b_adapter_v7_r16")
p.add_argument("--base", default="mlx-community/Qwen3-4B-Instruct-2507-4bit")
p.add_argument("--out", default=None, help="optional JSON output path")
p.add_argument("--limit", type=int, default=0, help="limit turns for quick test")
args = p.parse_args()
print(f"loading {args.base} + {args.adapter}")
model, tok = load(args.base, adapter_path=args.adapter)
rows = list(walk_scenarios())
if args.limit:
rows = rows[:args.limit]
print(f"total ATC-to-ownship turns: {len(rows)}")
# Single generation pass: dump every row's gold + raw output + parsed.
# Scoring is offline via score_baseline.py — reusable for multiple metric variants.
rows_out = []
for i, (scn_id, turn_idx, text, gold_intent, gold_params) in enumerate(rows):
msgs = [
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": text},
]
prompt = tok.apply_chat_template(msgs, add_generation_prompt=True, tokenize=False)
out = generate(model, tok, prompt=prompt, max_tokens=200, verbose=False)
parsed = parse(out)
rows_out.append({
"scenario_id": scn_id,
"turn_index": turn_idx,
"text": text,
"gold_intent": gold_intent,
"gold_parameters": gold_params,
"raw_output": out,
"parsed": parsed,
})
if (i + 1) % 50 == 0:
print(f" ...{i+1}/{len(rows)} generated")
out_path = Path(args.out) if args.out else Path(f"poc/llm-finetune/training/eval_dump_{Path(args.adapter).name}.json")
out_path.parent.mkdir(parents=True, exist_ok=True)
out_path.write_text(json.dumps({
"adapter": args.adapter,
"base": args.base,
"scenario_dir": str(SCENARIO_DIR),
"n": len(rows_out),
"rows": rows_out,
}, indent=2))
print(f"\nDUMP COMPLETE: {len(rows_out)} rows -> {out_path}")
print(f"Score with: uv run python poc/llm-finetune/training/score_baseline.py {out_path}")
if __name__ == "__main__":
main()