| """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() |
|
|