Upload eval_baseline.py with huggingface_hub
Browse files- eval_baseline.py +25 -45
eval_baseline.py
CHANGED
|
@@ -91,11 +91,9 @@ def main():
|
|
| 91 |
rows = rows[:args.limit]
|
| 92 |
print(f"total ATC-to-ownship turns: {len(rows)}")
|
| 93 |
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
raw_outputs = []
|
| 98 |
-
|
| 99 |
for i, (scn_id, turn_idx, text, gold_intent, gold_params) in enumerate(rows):
|
| 100 |
msgs = [
|
| 101 |
{"role": "system", "content": SYSTEM_PROMPT},
|
|
@@ -104,47 +102,29 @@ def main():
|
|
| 104 |
prompt = tok.apply_chat_template(msgs, add_generation_prompt=True, tokenize=False)
|
| 105 |
out = generate(model, tok, prompt=prompt, max_tokens=200, verbose=False)
|
| 106 |
parsed = parse(out)
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
pred = {k: str(v) for k, v in parsed["slots"].items()}
|
| 117 |
-
is_slots = all(pred.get(k) == v for k, v in gold_slots.items())
|
| 118 |
-
intent_em += int(is_intent); slots_em += int(is_slots)
|
| 119 |
-
by_intent[gold_intent][0] += 1
|
| 120 |
-
by_intent[gold_intent][1] += int(is_intent)
|
| 121 |
-
by_intent[gold_intent][2] += int(is_slots)
|
| 122 |
-
by_scenario[scn_id][0] += 1
|
| 123 |
-
by_scenario[scn_id][1] += int(is_intent)
|
| 124 |
-
by_scenario[scn_id][2] += int(is_slots)
|
| 125 |
-
if i < 5:
|
| 126 |
-
raw_outputs.append({"scn": scn_id, "turn": turn_idx, "text": text, "gold": {"intent": gold_intent, "slots": gold_slots}, "pred": parsed})
|
| 127 |
if (i + 1) % 50 == 0:
|
| 128 |
-
print(f" ...{i+1}/{len(rows)}
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
if args.out:
|
| 142 |
-
Path(args.out).write_text(json.dumps({
|
| 143 |
-
"adapter": args.adapter, "n": n, "intent_em": intent_em, "slots_em": slots_em,
|
| 144 |
-
"schema_valid": schema_ok,
|
| 145 |
-
"by_intent": dict(by_intent), "by_scenario": dict(by_scenario),
|
| 146 |
-
"samples": raw_outputs,
|
| 147 |
-
}, indent=2))
|
| 148 |
|
| 149 |
|
| 150 |
if __name__ == "__main__":
|
|
|
|
| 91 |
rows = rows[:args.limit]
|
| 92 |
print(f"total ATC-to-ownship turns: {len(rows)}")
|
| 93 |
|
| 94 |
+
# Single generation pass: dump every row's gold + raw output + parsed.
|
| 95 |
+
# Scoring is offline via score_baseline.py — reusable for multiple metric variants.
|
| 96 |
+
rows_out = []
|
|
|
|
|
|
|
| 97 |
for i, (scn_id, turn_idx, text, gold_intent, gold_params) in enumerate(rows):
|
| 98 |
msgs = [
|
| 99 |
{"role": "system", "content": SYSTEM_PROMPT},
|
|
|
|
| 102 |
prompt = tok.apply_chat_template(msgs, add_generation_prompt=True, tokenize=False)
|
| 103 |
out = generate(model, tok, prompt=prompt, max_tokens=200, verbose=False)
|
| 104 |
parsed = parse(out)
|
| 105 |
+
rows_out.append({
|
| 106 |
+
"scenario_id": scn_id,
|
| 107 |
+
"turn_index": turn_idx,
|
| 108 |
+
"text": text,
|
| 109 |
+
"gold_intent": gold_intent,
|
| 110 |
+
"gold_parameters": gold_params,
|
| 111 |
+
"raw_output": out,
|
| 112 |
+
"parsed": parsed,
|
| 113 |
+
})
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 114 |
if (i + 1) % 50 == 0:
|
| 115 |
+
print(f" ...{i+1}/{len(rows)} generated")
|
| 116 |
+
|
| 117 |
+
out_path = Path(args.out) if args.out else Path(f"poc/llm-finetune/training/eval_dump_{Path(args.adapter).name}.json")
|
| 118 |
+
out_path.parent.mkdir(parents=True, exist_ok=True)
|
| 119 |
+
out_path.write_text(json.dumps({
|
| 120 |
+
"adapter": args.adapter,
|
| 121 |
+
"base": args.base,
|
| 122 |
+
"scenario_dir": str(SCENARIO_DIR),
|
| 123 |
+
"n": len(rows_out),
|
| 124 |
+
"rows": rows_out,
|
| 125 |
+
}, indent=2))
|
| 126 |
+
print(f"\nDUMP COMPLETE: {len(rows_out)} rows -> {out_path}")
|
| 127 |
+
print(f"Score with: uv run python poc/llm-finetune/training/score_baseline.py {out_path}")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 128 |
|
| 129 |
|
| 130 |
if __name__ == "__main__":
|