File size: 5,504 Bytes
da1f428 f77f964 da1f428 f77f964 da1f428 6dcffbd da1f428 6dcffbd da1f428 6dcffbd da1f428 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 | """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()
|