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