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