rudeparis commited on
Commit
6dcffbd
·
verified ·
1 Parent(s): f77f964

Upload eval_baseline.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. 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
- schema_ok = intent_em = slots_em = 0
95
- by_intent = defaultdict(lambda: [0, 0, 0]) # n, intent_hits, slots_hits
96
- by_scenario = defaultdict(lambda: [0, 0, 0])
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
- if parsed:
109
- schema_ok += 1
110
- gold_slots = {k: str(v) for k, v in gold_params.items() if v is not None}
111
- is_intent = bool(parsed and parsed.get("intent") == gold_intent)
112
- # Slot match: check predicted slots includes all gold keys with matching values
113
- # (tolerates extra slot keys; harsh exact-match like eval_holdout would be too strict here)
114
- is_slots = False
115
- if parsed and isinstance(parsed.get("slots"), dict):
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)}: intent={intent_em/(i+1):.1%} slots={slots_em/(i+1):.1%}")
129
-
130
- n = len(rows)
131
- print(f"\n=== CANONICAL BASELINE (134 scenarios × 553 turns equivalent) ===")
132
- print(f"adapter: {args.adapter}")
133
- print(f"total turns: {n}")
134
- print(f"schema_valid: {schema_ok}/{n} = {schema_ok/n:.1%}")
135
- print(f"intent_em: {intent_em}/{n} = {intent_em/n:.1%} (V32 baseline: 87.4%)")
136
- print(f"slots_em: {slots_em}/{n} = {slots_em/n:.1%} (V32 baseline: 80.6%)")
137
- print(f"\n=== PER-INTENT (top 20 by count) ===")
138
- for intent, (n_i, i_em, s_em) in sorted(by_intent.items(), key=lambda kv: -kv[1][0])[:20]:
139
- print(f" {intent:30s} n={n_i:4d} intent={i_em}/{n_i} ({i_em/n_i:.0%}) slots={s_em}/{n_i} ({s_em/n_i:.0%})")
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__":