| |
| """Evaluate whether L2 rationales support and ground their decisions. |
| |
| This script consumes row-level L2 prediction JSONL files produced by |
| `scripts/eval_l2_decision.py` or compatible runners. It intentionally starts |
| with deterministic rule checks so the metric is reproducible for paper tables. |
| Embedding/NLI/LLM judges can be added later as auxiliary metrics. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import math |
| import os |
| import re |
| import sys |
| from collections import Counter |
| from typing import Any, Dict, Iterable, List, Sequence, Set |
|
|
| sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src")) |
| from layered_belief import CATEGORY, macro_recall |
|
|
|
|
| DISTANCE_PATTERNS = { |
| "near": [ |
| r"\bnear distance\b", |
| r"\bdistance is near\b", |
| r"\bpoint[- ]blank\b", |
| r"\bvery close\b", |
| r"\badjacent\b", |
| ], |
| "mid": [ |
| r"\bmid distance\b", |
| r"\bmedium distance\b", |
| r"\bdistance is mid\b", |
| r"\bdistance is medium\b", |
| ], |
| "close_range": [ |
| r"\bclose range\b", |
| r"\bwithin range\b", |
| r"\bshort range\b", |
| r"\b2\s*-\s*4\b", |
| r"\b0\s*-\s*2\b", |
| ], |
| "far": [ |
| r"\bfar distance\b", |
| r"\bdistance is far\b", |
| r"\blong range\b", |
| r"\bat range\b", |
| r"\bdistant\b", |
| r"\bkeep distance\b", |
| r"\bkite\b", |
| r"\b4\s*-\s*6\b", |
| r"\b6\+\b", |
| ], |
| } |
|
|
| FAMILY_PATTERNS = { |
| "close": [ |
| r"\bclose skill\b", |
| r"\bmelee\b", |
| r"\bgrab\b", |
| r"\bslam\b", |
| r"\bpoint[- ]blank\b", |
| r"\badjacent\b", |
| ], |
| "far": [ |
| r"\bfar skill\b", |
| r"\branged\b", |
| r"\bprojectile\b", |
| r"\bgap[- ]close\b", |
| r"\bleap\b", |
| r"\bpull\b", |
| ], |
| "cd_aoe": [ |
| r"\bspecial skill\b", |
| r"\baoe\b", |
| r"\barea\b", |
| r"\bshockwave\b", |
| r"\bspin\b", |
| ], |
| "summon": [ |
| r"\bsummon\b", |
| r"\badds\b", |
| ], |
| } |
|
|
|
|
| def load_jsonl(path: str) -> List[Dict[str, Any]]: |
| rows = [] |
| with open(path, encoding="utf-8") as f: |
| for line in f: |
| if line.strip(): |
| rows.append(json.loads(line)) |
| return rows |
|
|
|
|
| def rate(values: Iterable[bool]) -> float: |
| vals = list(values) |
| return sum(vals) / max(1, len(vals)) |
|
|
|
|
| def safe_round(value: float) -> float: |
| if value is None or (isinstance(value, float) and math.isnan(value)): |
| return float("nan") |
| return round(float(value), 4) |
|
|
|
|
| def bool_or_none(value: Any) -> bool | None: |
| if isinstance(value, bool): |
| return value |
| if value is None: |
| return None |
| if isinstance(value, str): |
| v = value.strip().lower() |
| if v in {"true", "yes", "1"}: |
| return True |
| if v in {"false", "no", "0"}: |
| return False |
| return None |
|
|
|
|
| def normalized_belief(row: Dict[str, Any]) -> Dict[str, Any]: |
| belief = row.get("belief") or {} |
| nested = row.get("l1_belief") or {} |
| geom = nested.get("geometry", {}) if isinstance(nested, dict) else {} |
| boss_state = nested.get("boss_state", {}) if isinstance(nested, dict) else {} |
| resource = nested.get("resource_state", {}) if isinstance(nested, dict) else {} |
| return { |
| "distance_bin": geom.get("distance_bin", belief.get("player_distance_bin")), |
| "dp_bin": geom.get("dp_bin", belief.get("dp_bin")), |
| "front_cone": bool_or_none(geom.get("front_cone", belief.get("front_cone"))), |
| "decision_zone": geom.get("decision_zone", belief.get("decision_zone")), |
| "tactical_sector": geom.get("tactical_sector", belief.get("tactical_sector")), |
| "behind": bool_or_none(geom.get("behind", belief.get("behind"))), |
| "cooldown_ready": resource.get("cooldown_ready", belief.get("cooldown_ready", {})) or {}, |
| "hp_phase": resource.get("hp_phase", belief.get("hp_phase")), |
| "prev_boss_skill": boss_state.get("prev_skill", belief.get("prev_boss_skill")), |
| } |
|
|
|
|
| def matches_any(text: str, patterns: Sequence[str]) -> bool: |
| return any(re.search(pattern, text) for pattern in patterns) |
|
|
|
|
| def extract_claims(reason: str) -> Dict[str, Any]: |
| text = str(reason or "").lower() |
| distance_claims = { |
| label for label, patterns in DISTANCE_PATTERNS.items() |
| if matches_any(text, patterns) |
| } |
| family_claims = { |
| label for label, patterns in FAMILY_PATTERNS.items() |
| if matches_any(text, patterns) |
| } |
| front_claim = None |
| if re.search(r"\bfront cone\b|\bin front\b|\bfront\b", text): |
| front_claim = True |
| if re.search(r"\bnot in front\b|\bside\b|\bflank\b", text): |
| front_claim = False |
| behind_claim = None |
| if re.search(r"\bbehind\b", text): |
| behind_claim = True |
| cooldown_claim = bool(re.search(r"\bcooldown\b|\bready\b|\busable\b|\boff cooldown\b", text)) |
| hp_claims = { |
| label for label in ("high", "mid", "low") |
| if re.search(rf"\b{label} hp\b|\b{label} health\b|\bhp is {label}\b", text) |
| } |
| return { |
| "distance": distance_claims, |
| "family": family_claims, |
| "front_cone": front_claim, |
| "behind": behind_claim, |
| "cooldown": cooldown_claim, |
| "hp_phase": hp_claims, |
| "has_factor": bool(distance_claims or family_claims or front_claim is not None |
| or behind_claim is not None or cooldown_claim or hp_claims), |
| } |
|
|
|
|
| def dp_distance_bucket(dp_bin: str | None) -> str | None: |
| if dp_bin in {"0-2", "0_2"}: |
| return "near" |
| if dp_bin == "2-4": |
| return "mid" |
| if dp_bin in {"4-6", "6+"}: |
| return "far" |
| return None |
|
|
|
|
| def distance_claim_grounded(claim: str, belief: Dict[str, Any]) -> bool: |
| bins: Set[str] = set() |
| if belief.get("distance_bin"): |
| bins.add(str(belief["distance_bin"])) |
| dp_bucket = dp_distance_bucket(belief.get("dp_bin")) |
| if dp_bucket: |
| bins.add(dp_bucket) |
| if claim == "near": |
| return "near" in bins |
| if claim == "mid": |
| return "mid" in bins |
| if claim == "close_range": |
| return bool(bins.intersection({"near", "mid"})) |
| if claim == "far": |
| return "far" in bins |
| return False |
|
|
|
|
| def decision_alignment(skill: str | None, claims: Dict[str, Any]) -> bool: |
| if not skill: |
| return False |
| cat = CATEGORY.get(skill, "other") |
| distance_claims = claims["distance"] |
| family_claims = claims["family"] |
| if cat == "close": |
| return bool(family_claims.intersection({"close"}) or |
| distance_claims.intersection({"near", "mid", "close_range"})) |
| if cat == "far": |
| return bool(family_claims.intersection({"far"}) or "far" in distance_claims) |
| if cat == "cd_aoe": |
| return "cd_aoe" in family_claims |
| if cat == "summon": |
| return "summon" in family_claims |
| return False |
|
|
|
|
| def belief_grounding(skill: str | None, claims: Dict[str, Any], belief: Dict[str, Any]) -> Dict[str, Any]: |
| checked = [] |
| contradictions = [] |
| for claim in claims["distance"]: |
| checked.append(f"distance:{claim}") |
| if not distance_claim_grounded(claim, belief): |
| contradictions.append(f"distance:{claim}") |
| if claims["front_cone"] is not None: |
| checked.append("front_cone") |
| if belief.get("front_cone") is not None and claims["front_cone"] != belief["front_cone"]: |
| contradictions.append("front_cone") |
| if claims["behind"] is not None: |
| checked.append("behind") |
| if belief.get("behind") is not None and claims["behind"] != belief["behind"]: |
| contradictions.append("behind") |
| if claims["cooldown"] and skill: |
| checked.append("cooldown") |
| ready = belief.get("cooldown_ready", {}).get(skill) |
| if ready is False: |
| contradictions.append("cooldown") |
| for claim in claims["hp_phase"]: |
| checked.append(f"hp_phase:{claim}") |
| if belief.get("hp_phase") and claim != belief["hp_phase"]: |
| contradictions.append(f"hp_phase:{claim}") |
| return { |
| "checked": checked, |
| "contradictions": contradictions, |
| "grounded": bool(checked) and not contradictions, |
| } |
|
|
|
|
| def evaluate_row(row: Dict[str, Any]) -> Dict[str, Any]: |
| decision = row.get("decision") or {} |
| skill = decision.get("skill") |
| reason = decision.get("reason", "") |
| claims = extract_claims(reason) |
| belief = normalized_belief(row) |
| grounding = belief_grounding(skill, claims, belief) |
| valid = row.get("valid") or {} |
| target = row.get("target_skill") |
| return { |
| "boss": row.get("boss"), |
| "fight": row.get("fight"), |
| "index": row.get("index"), |
| "target_skill": target, |
| "pred_skill": skill, |
| "pred_family": CATEGORY.get(skill, "other"), |
| "target_family": CATEGORY.get(target, "other"), |
| "reason": reason, |
| "reason_nonempty": bool(str(reason).strip()), |
| "rationale_informative": bool(claims["has_factor"]), |
| "decision_rationale_aligned": decision_alignment(skill, claims), |
| "rationale_belief_grounded": grounding["grounded"], |
| "rationale_contradiction": bool(grounding["contradictions"]), |
| "checked_claims": grounding["checked"], |
| "contradictions": grounding["contradictions"], |
| "json_valid": bool(valid.get("json_valid", True)), |
| "schema_valid": bool(valid.get("schema_valid", True)), |
| "legal": bool(valid.get("legal", skill in (row.get("legal_skills") or []))), |
| "cooldown_ok": bool(valid.get("cooldown_ok", True)), |
| "rule_grounded": bool(valid.get("grounded", True)), |
| } |
|
|
|
|
| def summarize(rows: List[Dict[str, Any]], source: str) -> Dict[str, Any]: |
| details = [evaluate_row(row) for row in rows] |
| y_true = [d["target_skill"] for d in details] |
| y_pred = [d["pred_skill"] for d in details] |
| labels = sorted({s for row in rows for s in (row.get("legal_skills") or [])}) |
| valid_reason = [ |
| d["schema_valid"] and d["legal"] and d["cooldown_ok"] and |
| d["decision_rationale_aligned"] and d["rationale_belief_grounded"] |
| for d in details |
| ] |
| return { |
| "source": source, |
| "n": len(details), |
| "skill_match_rate": safe_round(rate(t == p for t, p in zip(y_true, y_pred))), |
| "skill_macro_recall": safe_round(macro_recall(y_true, y_pred, labels)) if labels else float("nan"), |
| "skill_family_match_rate": safe_round(rate( |
| d["target_family"] == d["pred_family"] for d in details |
| )), |
| "json_valid_rate": safe_round(rate(d["json_valid"] for d in details)), |
| "schema_valid_rate": safe_round(rate(d["schema_valid"] for d in details)), |
| "legal_rate": safe_round(rate(d["legal"] for d in details)), |
| "cooldown_ok_rate": safe_round(rate(d["cooldown_ok"] for d in details)), |
| "rule_grounded_rate": safe_round(rate(d["rule_grounded"] for d in details)), |
| "reason_nonempty_rate": safe_round(rate(d["reason_nonempty"] for d in details)), |
| "rationale_informative_rate": safe_round(rate(d["rationale_informative"] for d in details)), |
| "decision_rationale_alignment_rate": safe_round(rate(d["decision_rationale_aligned"] for d in details)), |
| "rationale_belief_grounding_rate": safe_round(rate(d["rationale_belief_grounded"] for d in details)), |
| "rationale_contradiction_rate": safe_round(rate(d["rationale_contradiction"] for d in details)), |
| "valid_reasoned_decision_rate": safe_round(rate(valid_reason)), |
| "prediction_counts": dict(Counter(y_pred).most_common()), |
| "family_counts": dict(Counter(d["pred_family"] for d in details).most_common()), |
| } |
|
|
|
|
| def main() -> None: |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--predictions", required=True, help="L2 row-level prediction JSONL.") |
| ap.add_argument("--out", default=None, help="Summary JSON output path.") |
| ap.add_argument("--details_out", default=None, help="Optional per-row JSONL details.") |
| args = ap.parse_args() |
|
|
| rows = load_jsonl(args.predictions) |
| details = [evaluate_row(row) for row in rows] |
| summary = summarize(rows, args.predictions) |
| print(json.dumps(summary, ensure_ascii=False, indent=2)) |
|
|
| out = args.out or os.path.splitext(args.predictions)[0] + "_rationale_eval.json" |
| os.makedirs(os.path.dirname(out) or ".", exist_ok=True) |
| with open(out, "w", encoding="utf-8") as f: |
| json.dump(summary, f, ensure_ascii=False, indent=2) |
| if args.details_out: |
| os.makedirs(os.path.dirname(args.details_out) or ".", exist_ok=True) |
| with open(args.details_out, "w", encoding="utf-8") as f: |
| for row in details: |
| f.write(json.dumps(row, ensure_ascii=False) + "\n") |
| print(f"wrote {out}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|