#!/usr/bin/env python3
"""
案例 003:长时程衰减里,有多少能靠"上下文卫生"拿回来?

出发点(前沿理论成果):
  《The Illusion of Diminishing Returns: Measuring Long Horizon Execution in LLMs》
  (arXiv 2509.09677)提出 self-conditioning:当上下文里出现模型自己犯过的错,
  它后续出错概率上升;每步错误率随步数升高,且"把模型做大"并不能消除该效应。

本实验要外推的一步(原文未做的解耦 + 一个可部署的干预):
  基线设置里,"上下文含自己的错"与"上下文变长"是绑在一起的。把两者拆开,
  并测量一个**不需要正确答案、可直接部署的运行时干预**能挽回多少。

预注册假设(跑之前写死,负结果照实报):
  H1(复现):条件 A 下,每步准确率随步数显著下降。
  H2(解耦):若 B(长度照样增长、但把历史里的错替换成正确值)不再下降,
     则衰减主因是"错误在场"而非"上下文变长";若 C(长度受限、错误仍在)
     也能恢复,则长度是主因。两者都恢复则两个因素都有份。
  H3(可部署,本案例的落点):纯运行时的上下文卫生(条件 C:只保留最近 k 步 +
     用模型自己报告的状态做一次重述,不使用任何正确答案)能挽回条件 B(上帝修正)
     所恢复损失的大部分(预注册阈值:≥50%)。

任务设计(纯执行,不含推理难度):
  给定一部字典 key->value(两位数),模型逐步执行:
     第 i 步:取出 key_i 的值,加到当前累计和上,只输出新的累计和。
  每一步的正确性可机器判定;难度只来自"稳定地执行很多步",这正是论文关心的量。

四个条件:
  A full        完整历史 + 模型自己的输出(可能含错)   ← 基线
  B oracle      完整历史,但每步把模型输出替换为正确值   ← 隔离"错误在场"
  C hygiene     只保留最近 k 步 + 复述模型自己报的当前和 ← 可部署,无需正确答案
  D both        只保留最近 k 步 + 正确值                ← 上限

用法:
  python run_horizon.py run   --out horizon_runs.json
  python run_horizon.py stats horizon_runs.json
"""
import argparse
import json
import random
import re

N_SEQ = 100          # 并行序列数(同时也是每步的 batch 大小)
N_STEPS = 100        # 每条序列的步数
KEEP_K = 4           # 上下文卫生保留的最近步数
DICT_SIZE = 40
SEED = 20260728
MODEL = "Qwen/Qwen2.5-7B-Instruct"
CONDS = ["restate_only", "truncate_only"]   # v2:只补跑这两个,前四个已在 v1 完成


def make_sequences():
    rng = random.Random(SEED)
    seqs = []
    for _ in range(N_SEQ):
        keys = [f"k{i:02d}" for i in range(DICT_SIZE)]
        table = {k: rng.randint(10, 99) for k in keys}
        order = [rng.choice(keys) for _ in range(N_STEPS)]
        totals, t = [], 0
        for k in order:
            t += table[k]
            totals.append(t)
        seqs.append({"table": table, "order": order, "totals": totals})
    return seqs


def table_text(table):
    return ", ".join(f"{k}={v}" for k, v in table.items())


SYS = ("You execute a running-sum task. At each step you are given a key. "
       "Add the key's value from the table to the running total. "
       "Reply with ONLY the new running total as a plain integer, nothing else.")


def build_prompt(seq, i, history, cond):
    """history: list of (key, shown_total) —— shown_total 依条件为模型自报值或正确值。"""
    TRUNCATE = cond in ("hygiene", "both", "truncate_only")
    RESTATE  = cond in ("hygiene", "both", "restate_only")
    lines = [f"Table: {table_text(seq['table'])}", ""]
    hist = history[-KEEP_K:] if TRUNCATE else history
    if RESTATE and history:
        lines.append(f"Current running total: {history[-1][1]}")   # 状态复述(用已展示的值,非正确答案)
    if hist:
        lines.append("Recent steps:" if TRUNCATE else "Steps so far:")
        for k, t in hist:
            lines.append(f"  {k} -> {t}")
    lines.append("")
    lines.append(f"Step {i+1}: key = {seq['order'][i]}. New running total =")
    return "\n".join(lines)


def parse_int(text):
    m = re.search(r"-?\d+", text.replace(",", ""))
    return int(m.group()) if m else None


def cmd_run(args):
    from vllm import LLM, SamplingParams
    seqs = make_sequences()
    llm = LLM(model=MODEL, gpu_memory_utilization=0.90, max_model_len=8192,
              enable_prefix_caching=True)
    tok = llm.get_tokenizer()
    sp = SamplingParams(temperature=0.0, max_tokens=12)

    results = {"config": {"n_seq": N_SEQ, "n_steps": N_STEPS, "keep_k": KEEP_K,
                          "dict_size": DICT_SIZE, "model": MODEL, "seed": SEED},
               "conditions": {}}

    for cond in CONDS:
        histories = [[] for _ in range(N_SEQ)]
        # per_step_correct[i][s] = 该序列第 s 步是否正确
        correct = [[None] * N_STEPS for _ in range(N_SEQ)]
        preds = [[None] * N_STEPS for _ in range(N_SEQ)]
        for s in range(N_STEPS):
            prompts = [tok.apply_chat_template(
                [{"role": "system", "content": SYS},
                 {"role": "user", "content": build_prompt(seqs[j], s, histories[j], cond)}],
                tokenize=False, add_generation_prompt=True) for j in range(N_SEQ)]
            outs = sorted(llm.generate(prompts, sp), key=lambda o: int(o.request_id))
            for j, o in enumerate(outs):
                p = parse_int(o.outputs[0].text)
                gold = seqs[j]["totals"][s]
                preds[j][s] = p
                correct[j][s] = int(p == gold)
                shown = gold if cond in ("oracle", "both") else (p if p is not None else gold)
                histories[j].append((seqs[j]["order"][s], shown))
            if (s + 1) % 20 == 0:
                acc = sum(correct[j][s] for j in range(N_SEQ)) / N_SEQ
                print(f"[{cond}] step {s+1}/{N_STEPS} 本步准确率={acc:.3f}", flush=True)
        results["conditions"][cond] = {"correct": correct, "preds": preds}
        overall = sum(sum(r) for r in correct) / (N_SEQ * N_STEPS)
        print(f"[{cond}] 全程平均每步准确率 = {overall:.4f}", flush=True)

    results["gold"] = [s["totals"] for s in seqs]
    with open(args.out, "w") as f:
        json.dump(results, f)
    print("saved", args.out, flush=True)


def boot_ci(vals, n=10000, seed=SEED):
    import random as _r
    rng = _r.Random(seed)
    n_ = len(vals)
    ms = sorted(sum(vals[rng.randrange(n_)] for _ in range(n_)) / n_ for _ in range(n))
    return ms[250], ms[9750]


def cmd_stats(args):
    d = json.load(open(args.results))
    C = d["conditions"]
    n_steps = d["config"]["n_steps"]
    rep = {"config": d["config"], "per_condition": {}}

    def block_acc(correct, lo, hi):
        v = [c for row in correct for c in row[lo:hi]]
        return sum(v) / len(v)

    for cond in CONDS:
        cor = C[cond]["correct"]
        per_seq = [sum(r) / len(r) for r in cor]
        lo, hi = boot_ci(per_seq)
        first = block_acc(cor, 0, 20)
        last = block_acc(cor, n_steps - 20, n_steps)
        # 每步(跨序列)准确率曲线
        curve = [sum(cor[j][s] for j in range(len(cor))) / len(cor) for s in range(n_steps)]
        # 平均能连续正确执行多少步(首次出错前的步数)
        horizon = [next((s for s, c in enumerate(r) if c == 0), n_steps) for r in cor]
        rep["per_condition"][cond] = {
            "mean_step_acc": sum(per_seq) / len(per_seq), "ci95": [lo, hi],
            "first20": first, "last20": last, "drop": first - last,
            "curve": curve, "mean_first_error_step": sum(horizon) / len(horizon),
        }

    A, B, Cc, D = (rep["per_condition"][k] for k in CONDS)
    loss = A["drop"]                                   # 基线的衰减幅度
    rec_oracle = loss - B["drop"]                      # 上帝修正挽回多少
    rec_hyg = loss - Cc["drop"]                        # 上下文卫生挽回多少
    rep["verdicts"] = {
        "H1_baseline_degrades": {"first20": A["first20"], "last20": A["last20"],
                                 "drop": loss, "holds": loss > 0.02},
        "H2_decoupling": {"oracle_drop": B["drop"], "hygiene_drop": Cc["drop"],
                          "both_drop": D["drop"],
                          "error_presence_share": (rec_oracle / loss) if loss else None,
                          "length_share": (rec_hyg / loss) if loss else None},
        "H3_deployable_recovery": {
            "recovered_by_hygiene_over_oracle": (rec_hyg / rec_oracle) if rec_oracle else None,
            "threshold": 0.5,
            "holds": bool(rec_oracle and (rec_hyg / rec_oracle) >= 0.5)},
    }
    print(json.dumps({k: v for k, v in rep.items() if k != "per_condition"}, indent=2, ensure_ascii=False))
    for cond in CONDS:
        p = rep["per_condition"][cond]
        print(f"{cond:<9} 每步准确率 {p['mean_step_acc']:.4f} CI[{p['ci95'][0]:.4f},{p['ci95'][1]:.4f}] "
              f"| 前20步 {p['first20']:.4f} → 后20步 {p['last20']:.4f} (降 {p['drop']:.4f}) "
              f"| 平均首次出错步 {p['mean_first_error_step']:.1f}")
    json.dump(rep, open("horizon_stats.json", "w"), ensure_ascii=False)


if __name__ == "__main__":
    ap = argparse.ArgumentParser()
    sub = ap.add_subparsers(dest="cmd", required=True)
    r = sub.add_parser("run"); r.add_argument("--out", required=True); r.set_defaults(func=cmd_run)
    s = sub.add_parser("stats"); s.add_argument("results"); s.set_defaults(func=cmd_stats)
    a = ap.parse_args(); a.func(a)
