atc-parser-scripts / eval_v1.py
rudeparis's picture
Upload eval_v1.py with huggingface_hub
df22e77 verified
Raw
History Blame Contribute Delete
6.33 kB
"""V1 multi-segment eval: walks canonical scenarios, dumps generation results
with full V1 schema (segments + abstain_reason) for offline scoring.
Output dump shape per row:
{
"scenario_id", "turn_index", "text",
"compound": bool,
"gold": {"segments": [...], "abstain_reason": null},
"raw_output": <raw LLM text>,
"parsed": <parsed payload or null>,
}
Usage:
uv run python -u poc/llm-finetune/training/eval_v1.py \\
--adapter poc/llm-finetune/training/q3b_adapter_v9 \\
--out poc/llm-finetune/training/eval_dump_v9.json
"""
from __future__ import annotations
import argparse
import json
import re
import sys
from pathlib import Path
from mlx_lm import load, generate
ROOT = Path(__file__).resolve().parents[3]
INTENTS_50 = sorted(json.loads((ROOT / "poc/deberta_intent/checkpoints-base/label_mapping.json").read_text())["intent2id"].keys())
INTENTS_51 = INTENTS_50 + ["unknown"]
SLOTS_LOWER = ["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",
"sequence"]
SYSTEM_PROMPT = (
"You are an ATC parser. Parse the air traffic control transmission into a JSON object.\n\n"
"OUTPUT SCHEMA (always emit `segments` list, even single-intent turns):\n"
'{\n "segments": [\n {"intent": <one-of-enum>, "slots": {<lowercase_key>: <value>}, "text": <segment substring>}\n ],\n "abstain_reason": null | <short reason if unsure>\n}\n\n'
"Intent enum (51 values, includes `unknown` for ambiguous/garbled):\n"
+ ", ".join(INTENTS_51) + "\n\n"
"Slot keys (lowercase only): " + ", ".join(SLOTS_LOWER) + "\n\n"
"Rules:\n"
"1. Compound transmissions get multiple segments — split by comma, period, or new clause.\n"
"2. Single-intent transmissions still emit one segment in the list.\n"
"3. Use `unknown` intent + `abstain_reason` when transcript is garbled, partial, or doesn't match any enum.\n"
"4. Slots is FMM-relevant subset only — do NOT include callsign or facility unless required for the action.\n"
"5. Output ONLY the JSON object."
)
SCENARIO_DIR = ROOT / "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 normalize_segment(seg: dict) -> dict | None:
intent = seg.get("intent")
if intent not in INTENTS_50:
return None
slots_raw = seg.get("slots") or {}
slots = {k.lower(): str(v) for k, v in slots_raw.items() if v is not None and str(v) != ""}
return {"intent": intent, "slots": slots, "text": (seg.get("text") or "").strip()}
def walk_canonical_with_gold():
"""Yield (scn_id, turn_idx, text, compound, gold_payload) for every ATC-spoken turn."""
for path in sorted(SCENARIO_DIR.glob("*.json")):
try:
scn = json.loads(path.read_text())
except Exception:
continue
scn_id = scn.get("scenario_id", path.stem)
for turn in scn.get("turns", []):
if turn.get("speaker", "").lower() != "atc":
continue
text = (turn.get("expected_transcript") or "").strip()
if not text:
continue
compound = bool(turn.get("compound") and turn.get("expected_segments"))
if compound:
segs = []
for s in turn["expected_segments"]:
n = normalize_segment(s)
if n:
segs.append(n)
if not segs:
continue
else:
intent = turn.get("expected_intent")
if intent not in INTENTS_50:
continue
params = turn.get("expected_parameters") or {}
slots = {k.lower(): str(v) for k, v in params.items() if v is not None and str(v) != ""}
segs = [{"intent": intent, "slots": slots, "text": text}]
gold = {"segments": segs, "abstain_reason": None}
yield scn_id, turn.get("turn_index", 0), text, compound, gold
def main():
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--adapter", required=True)
p.add_argument("--base", default="mlx-community/Qwen3-4B-Instruct-2507-4bit")
p.add_argument("--out", required=True)
p.add_argument("--limit", type=int, default=0)
args = p.parse_args()
print(f"loading {args.base} + {args.adapter}", flush=True)
model, tok = load(args.base, adapter_path=args.adapter)
rows_iter = list(walk_canonical_with_gold())
if args.limit:
rows_iter = rows_iter[:args.limit]
print(f"total turns: {len(rows_iter)}", flush=True)
rows_out = []
for i, (scn_id, turn_idx, text, compound, gold) in enumerate(rows_iter):
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=512, verbose=False)
parsed = parse(out)
rows_out.append({
"scenario_id": scn_id,
"turn_index": turn_idx,
"text": text,
"compound": compound,
"gold": gold,
"raw_output": out,
"parsed": parsed,
})
if (i + 1) % 25 == 0:
print(f" ...{i+1}/{len(rows_iter)} generated", flush=True)
out_path = Path(args.out)
out_path.parent.mkdir(parents=True, exist_ok=True)
out_path.write_text(json.dumps({
"adapter": args.adapter, "base": args.base, "n": len(rows_out),
"rows": rows_out,
}, indent=2))
print(f"\nDUMP COMPLETE: {len(rows_out)} -> {out_path}", flush=True)
print(f"Score: uv run python poc/llm-finetune/training/score_v1.py {out_path} --use-verifier", flush=True)
if __name__ == "__main__":
main()