File size: 12,834 Bytes
544e392
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
#!/usr/bin/env python3
"""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  # noqa: E402


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