"""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": , "parsed": , } 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": , "slots": {: }, "text": }\n ],\n "abstain_reason": null | \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'.*?\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()