#!/usr/bin/env python3
"""
Pilot:量化收益的任务依赖性(FP16 vs INT4-AWQ)。来自一位提交者研究方向的试点实验(已按提交者隐私脱敏)。

预注册假设(跑之前写死,负结果照实报):
  H1: INT4-AWQ 相对 FP16 有显著吞吐收益(两类任务都有)。
  H2: 量化的质量损失是任务依赖的——数学推理(GSM8K)损失 > 抽取式问答(SQuAD)损失
      (交互效应 bootstrap 95% CI 不含 0)。
  若 H2 成立 → 支持"最优 inference 配置应按任务在运行时选择"(agent-aware inference 的核心前提)。

用法(Colab A100 上,分两次跑避免双模型显存/清理问题):
  python run_pilot.py run --config fp16 --out results_fp16.json
  python run_pilot.py run --config awq  --out results_awq.json
  python run_pilot.py stats results_fp16.json results_awq.json

诚实性约束:greedy 解码(确定性);同一批题目、同一 prompt、同一 max_tokens;
逐题结果全部落盘;CI 用逐题 bootstrap(10k 次);不挑窗口、不删离群点。
"""
import argparse
import json
import re
import string
import time

MODELS = {
    "fp16": "Qwen/Qwen2.5-7B-Instruct",
    "awq": "Qwen/Qwen2.5-7B-Instruct-AWQ",
}
N_ITEMS = 200
SEED = 20260726  # 只用于抽样题目子集,两配置共享同一子集


# ---------- 数据 ----------

def load_tasks():
    from datasets import load_dataset
    import random
    rng = random.Random(SEED)

    gsm = load_dataset("openai/gsm8k", "main", split="test")
    gsm_idx = rng.sample(range(len(gsm)), N_ITEMS)
    gsm_items = [{"id": f"gsm8k-{i}", "question": gsm[i]["question"],
                  "gold": gsm[i]["answer"].split("####")[-1].strip()} for i in gsm_idx]

    sq = load_dataset("rajpurkar/squad", split="validation")
    sq_idx = rng.sample(range(len(sq)), N_ITEMS)
    sq_items = [{"id": f"squad-{i}", "context": sq[i]["context"],
                 "question": sq[i]["question"],
                 "gold_answers": sq[i]["answers"]["text"]} for i in sq_idx]
    return gsm_items, sq_items


def gsm_prompt(q):
    return (q + "\n\nPlease reason step by step, and put your final numeric answer after '####'.")


def squad_prompt(context, question):
    return ("Answer the question using ONLY a short span copied from the context. "
            "Output the span and nothing else.\n\n"
            f"Context: {context}\nQuestion: {question}\nAnswer:")


# ---------- 评分 ----------

def extract_gsm_answer(text):
    if "####" in text:
        tail = text.split("####")[-1]
    else:
        tail = text
    nums = re.findall(r"-?\d[\d,]*\.?\d*", tail.replace("$", ""))
    return nums[-1].replace(",", "").rstrip(".") if nums else None


def gsm_correct(pred_text, gold):
    p = extract_gsm_answer(pred_text)
    g = gold.replace(",", "").strip()
    if p is None:
        return 0.0
    try:
        return float(abs(float(p) - float(g)) < 1e-6)
    except ValueError:
        return float(p == g)


def _norm(s):
    s = s.lower()
    s = "".join(ch for ch in s if ch not in set(string.punctuation))
    s = re.sub(r"\b(a|an|the)\b", " ", s)
    return " ".join(s.split())


def squad_f1(pred, golds):
    def f1(p, g):
        pt, gt = _norm(p).split(), _norm(g).split()
        common = {}
        for w in pt:
            if w in gt:
                common[w] = min(pt.count(w), gt.count(w))
        num_same = sum(common.values())
        if num_same == 0:
            return 0.0
        prec, rec = num_same / len(pt), num_same / len(gt)
        return 2 * prec * rec / (prec + rec)
    return max(f1(pred, g) for g in golds) if golds else 0.0


# ---------- 运行 ----------

def cmd_run(args):
    from vllm import LLM, SamplingParams
    model = MODELS[args.config]
    gsm_items, sq_items = load_tasks()

    llm = LLM(model=model, gpu_memory_utilization=0.90, max_model_len=4096)
    tok = llm.get_tokenizer()

    def chat(prompts, max_tokens):
        texts = [tok.apply_chat_template([{"role": "user", "content": p}],
                                         tokenize=False, add_generation_prompt=True)
                 for p in prompts]
        sp = SamplingParams(temperature=0.0, max_tokens=max_tokens)
        t0 = time.perf_counter()
        outs = llm.generate(texts, sp)
        dt = time.perf_counter() - t0
        outs = sorted(outs, key=lambda o: int(o.request_id))
        gen_tokens = sum(len(o.outputs[0].token_ids) for o in outs)
        return [o.outputs[0].text for o in outs], dt, gen_tokens

    results = {"config": args.config, "model": model, "n_items": N_ITEMS, "seed": SEED}

    preds, dt, ntok = chat([gsm_prompt(x["question"]) for x in gsm_items], 512)
    results["gsm8k"] = {
        "wall_s": dt, "gen_tokens": ntok, "tok_per_s": ntok / dt,
        "items": [{"id": x["id"], "gold": x["gold"], "pred_text": p,
                   "score": gsm_correct(p, x["gold"])}
                  for x, p in zip(gsm_items, preds)],
    }

    preds, dt, ntok = chat([squad_prompt(x["context"], x["question"]) for x in sq_items], 48)
    results["squad"] = {
        "wall_s": dt, "gen_tokens": ntok, "tok_per_s": ntok / dt,
        "items": [{"id": x["id"], "gold": x["gold_answers"], "pred_text": p,
                   "score": squad_f1(p.strip(), x["gold_answers"])}
                  for x, p in zip(sq_items, preds)],
    }

    for task in ("gsm8k", "squad"):
        scores = [it["score"] for it in results[task]["items"]]
        results[task]["mean_score"] = sum(scores) / len(scores)

    with open(args.out, "w") as f:
        json.dump(results, f, ensure_ascii=False)
    print(json.dumps({k: {"mean": results[k]["mean_score"],
                          "tok_per_s": round(results[k]["tok_per_s"], 1)}
                      for k in ("gsm8k", "squad")}, indent=2))


# ---------- 统计 ----------

def bootstrap_ci(deltas, n_boot=10000, seed=SEED):
    import random
    rng = random.Random(seed)
    n = len(deltas)
    means = []
    for _ in range(n_boot):
        s = [deltas[rng.randrange(n)] for _ in range(n)]
        means.append(sum(s) / n)
    means.sort()
    return means[int(0.025 * n_boot)], means[int(0.975 * n_boot)]


def cmd_stats(args):
    a = json.load(open(args.results[0]))
    b = json.load(open(args.results[1]))
    fp16, awq = (a, b) if a["config"] == "fp16" else (b, a)

    report = {}
    task_deltas = {}
    for task in ("gsm8k", "squad"):
        f_items = {it["id"]: it["score"] for it in fp16[task]["items"]}
        q_items = {it["id"]: it["score"] for it in awq[task]["items"]}
        ids = sorted(f_items)
        assert ids == sorted(q_items), "两配置题目集不一致"
        deltas = [q_items[i] - f_items[i] for i in ids]  # AWQ - FP16,负=量化受损
        lo, hi = bootstrap_ci(deltas)
        task_deltas[task] = deltas
        report[task] = {
            "fp16_score": fp16[task]["mean_score"],
            "awq_score": awq[task]["mean_score"],
            "quality_delta_awq_minus_fp16": sum(deltas) / len(deltas),
            "delta_ci95": [lo, hi],
            "fp16_tok_per_s": fp16[task]["tok_per_s"],
            "awq_tok_per_s": awq[task]["tok_per_s"],
            "throughput_speedup": awq[task]["tok_per_s"] / fp16[task]["tok_per_s"],
        }

    # 交互效应:量化在 gsm8k 上的损失是否显著大于 squad(独立组 bootstrap)
    import random
    rng = random.Random(SEED)
    g, s = task_deltas["gsm8k"], task_deltas["squad"]
    diffs = []
    for _ in range(10000):
        gm = sum(g[rng.randrange(len(g))] for _ in range(len(g))) / len(g)
        sm = sum(s[rng.randrange(len(s))] for _ in range(len(s))) / len(s)
        diffs.append(gm - sm)
    diffs.sort()
    report["interaction_gsm_minus_squad"] = {
        "point": (sum(g) / len(g)) - (sum(s) / len(s)),
        "ci95": [diffs[250], diffs[9750]],
        "task_dependent": diffs[9750] < 0 or diffs[250] > 0,
    }
    print(json.dumps(report, indent=2))
    with open("stats_report.json", "w") as f:
        json.dump(report, f, ensure_ascii=False, indent=2)


def main():
    ap = argparse.ArgumentParser()
    sub = ap.add_subparsers(dest="cmd", required=True)
    r = sub.add_parser("run")
    r.add_argument("--config", choices=list(MODELS), required=True)
    r.add_argument("--out", required=True)
    r.set_defaults(func=cmd_run)
    s = sub.add_parser("stats")
    s.add_argument("results", nargs=2)
    s.set_defaults(func=cmd_stats)
    args = ap.parse_args()
    args.func(args)


if __name__ == "__main__":
    main()
