File size: 6,333 Bytes
df22e77
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
"""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()