#!/usr/bin/env python3
"""
合并 v1(full/oracle/hygiene/both)与 v2(restate_only/truncate_only)的结果,
用修正后的度量重算,并做 2×2 拆解(截断 × 状态复述)。

修正后的度量:
  第 s 步是否正确 = (模型这一步输出 − 它上一步被展示的值) 是否等于本步该加的数值。
  也就是"有没有把该加的加对",与是否背着历史偏移无关。
  ——第一版把"累计和是否等于正确答案"当作判分,一旦早期算错就永久判错,
    那不是自我条件化,是度量本身坏掉。原始坏结果一并保留公开。
"""
import json
import sys

import numpy as np

V1 = "horizon_runs.json"
V2 = "horizon_v2.json"
ORACLE_CONDS = ("oracle", "both")      # 这两条腿的历史被替换为正确值,其余全部是模型自报值


def load_all():
    d1 = json.load(open(V1))
    conds = dict(d1["conditions"])
    gold = d1["gold"]
    cfg = d1["config"]
    try:
        d2 = json.load(open(V2))
        conds.update(d2["conditions"])
        assert d2["gold"] == gold, "两轮的题目序列不一致,不能合并"
    except FileNotFoundError:
        print("(未找到 v2,只分析 v1)", file=sys.stderr)
    return cfg, gold, conds


def stepwise(cond, preds, gold):
    """返回 (n_seq, n_steps) 的 0/1 矩阵:每步增量是否加对。"""
    n_steps = len(gold[0])
    inc = [[g[0]] + [g[s] - g[s - 1] for s in range(1, n_steps)] for g in gold]
    out = []
    for j, pr in enumerate(preds):
        row, prev = [], 0
        for s in range(n_steps):
            p = pr[s]
            row.append(int(p is not None and p - prev == inc[j][s]))
            prev = gold[j][s] if cond in ORACLE_CONDS else (p if p is not None else gold[j][s])
        out.append(row)
    return np.array(out)


def main():
    cfg, gold, conds = load_all()
    A = {c: stepwise(c, conds[c]["preds"], gold) for c in conds}
    rng = np.random.default_rng(20260728)
    n = len(gold)
    idx = rng.integers(0, n, size=(10000, n))

    def ci(vec):
        v = np.sort(vec[idx].mean(axis=1))
        return float(v[250]), float(v[9750])

    drops = {c: a[:, -20:].mean(1) - a[:, :20].mean(1) for c, a in A.items()}
    rep = {"config": cfg, "metric": "per-step increment correctness", "conditions": {}}
    for c, a in A.items():
        lo, hi = ci(a.mean(1))
        dlo, dhi = ci(drops[c])
        rep["conditions"][c] = {
            "mean_step_acc": float(a.mean()), "acc_ci95": [lo, hi],
            "first20": float(a[:, :20].mean()), "last20": float(a[:, -20:].mean()),
            "drop": float(drops[c].mean()), "drop_ci95": [dlo, dhi],
            "curve": a.mean(0).tolist(),
        }

    def contrast(name, c1, c2):
        diff = drops[c1] - drops[c2]
        lo, hi = ci(diff)
        return {"name": name, "delta": float(diff.mean()), "ci95": [lo, hi],
                "significant": bool(lo > 0 or hi < 0)}

    rep["contrasts"] = {}
    # 自我条件化(去掉自己的错)在长/短上下文下各值多少
    if {"oracle", "full"} <= set(A):
        rep["contrasts"]["self_cond_long"] = contrast("长上下文:去掉自己的错", "oracle", "full")
    if {"both", "hygiene"} <= set(A):
        rep["contrasts"]["self_cond_short"] = contrast("短上下文:去掉自己的错", "both", "hygiene")
    # 2×2:截断 × 状态复述(全部使用模型自报值,可部署)
    quad = {"full": ("长", "无复述"), "restate_only": ("长", "有复述"),
            "truncate_only": ("短", "无复述"), "hygiene": ("短", "有复述")}
    if set(quad) <= set(A):
        rep["contrasts"]["restate_at_long"] = contrast("长上下文下加状态复述", "restate_only", "full")
        rep["contrasts"]["restate_at_short"] = contrast("短上下文下加状态复述", "hygiene", "truncate_only")
        rep["contrasts"]["truncate_at_norestate"] = contrast("无复述时截断上下文", "truncate_only", "full")
        rep["contrasts"]["truncate_at_restate"] = contrast("有复述时截断上下文", "hygiene", "restate_only")
        base = -drops["full"].mean()
        rep["recovery_share"] = {
            c: float((base + drops[c].mean()) / base) for c in
            ("restate_only", "truncate_only", "hygiene", "oracle", "both") if c in A
        }

    json.dump(rep, open("horizon_final.json", "w"), ensure_ascii=False, indent=1)

    print(f"度量:每步增量是否加对 | 序列数 {n} | 步数 {len(gold[0])}\n")
    print(f"{'条件':<15}{'平均每步正确率':>14}{'前20步':>9}{'后20步':>9}{'衰减':>9}{'衰减CI95':>18}")
    order = [c for c in ("full", "restate_only", "truncate_only", "hygiene", "oracle", "both") if c in A]
    for c in order:
        p = rep["conditions"][c]
        print(f"{c:<15}{p['mean_step_acc']*100:>13.1f}%{p['first20']*100:>8.1f}%{p['last20']*100:>8.1f}%"
              f"{p['drop']*100:>+8.1f}  [{p['drop_ci95'][0]*100:+.1f},{p['drop_ci95'][1]*100:+.1f}]")
    print("\n对比(正值 = 衰减更小,即该干预有效):")
    for k, v in rep["contrasts"].items():
        print(f"  {v['name']:<22}{v['delta']*100:+7.1f} 点  CI95[{v['ci95'][0]*100:+.1f},{v['ci95'][1]*100:+.1f}]"
              f"  {'显著' if v['significant'] else '不显著'}")
    if "recovery_share" in rep:
        print("\n各干预挽回基线衰减的比例:")
        for c, v in rep["recovery_share"].items():
            print(f"  {c:<15}{v*100:6.0f}%")


if __name__ == "__main__":
    main()
