#!/usr/bin/env python3
"""
Round 2:量化损失是"格式"还是"内容"、能否被提示词救回、吞吐收益的拐点在哪。

预注册假设(跑之前写死,负结果照实报):
  H3(机制判定):把"内容正确率"(gold 片段是否出现在输出里)与 F1 分开测,
     FP16 与 AWQ 的内容正确率无显著差异,而 F1 有显著差异
     → 量化损伤的是输出格式保真,不是抽取能力。
  H4(命题生死线):提示词缓解(严格提示 / 2-shot / 2-shot+长度上限)能把 AWQ 与 FP16
     的 F1 差距缩到不显著 → 若成立,"运行时按任务选配置"的必要性下降(永远开缓解即可);
     若不成立,该命题增强。
  H5(吞吐拐点):speedup 随并发数单调下降,存在 batch* 使显著加速只出现在 batch < batch*。

产出:每个条件的 F1 / 内容正确率 / 格式合规率(逐题落盘),以及 5 个并发点的吞吐曲线。

用法(A100,两条腿分别跑):
  python run_round2.py run --config fp16 --out r2_fp16.json
  python run_round2.py run --config awq  --out r2_awq.json
  python run_round2.py stats r2_fp16.json r2_awq.json
"""
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          # 与 round 1 同种子,但抽签顺序不同(round 1 先抽 GSM8K)
                         # → 实际题目子集与 round 1 只重合 4 道,两轮绝对数字不可直接比。
                         # 本轮内部四个条件 × 两个模型用同一批题,逐题配对,内部对比有效。
BATCH_POINTS = [1, 4, 16, 64, 256]
THROUGHPUT_TOKENS = 128  # 固定输出长度,ignore_eos,使 tok/s 可比

SHOTS = [
    ("The Amazon rainforest covers most of the Amazon basin of South America, "
     "an area of 7,000,000 square kilometres.",
     "How large is the Amazon basin?", "7,000,000 square kilometres"),
    ("Chopin was born in Zelazowa Wola in 1810 and left Poland at the age of 20.",
     "In what year was Chopin born?", "1810"),
]


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

def load_squad():
    from datasets import load_dataset
    import random
    rng = random.Random(SEED)
    sq = load_dataset("rajpurkar/squad", split="validation")
    idx = rng.sample(range(len(sq)), N_ITEMS)
    return [{"id": f"squad-{i}", "context": sq[i]["context"],
             "question": sq[i]["question"],
             "gold_answers": sq[i]["answers"]["text"]} for i in idx]


# ---------- 四个条件的 prompt ----------

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


def p_strict(c, q):
    return ("Extract the answer span from the context. Rules: output ONLY the exact "
            "words copied from the context; no full sentence; no explanation; "
            "no more than 10 words.\n\n"
            f"Context: {c}\nQuestion: {q}\nAnswer:")


def p_fewshot(c, q):
    parts = ["Extract the answer span from the context. Output only the span."]
    for sc, sq_, sa in SHOTS:
        parts.append(f"\nContext: {sc}\nQuestion: {sq_}\nAnswer: {sa}")
    parts.append(f"\nContext: {c}\nQuestion: {q}\nAnswer:")
    return "".join(parts)


CONDITIONS = {
    "baseline": (p_baseline, 48),
    "strict": (p_strict, 48),
    "fewshot": (p_fewshot, 48),
    "fewshot_cap": (p_fewshot, 12),   # 2-shot + 硬性长度上限(工程上最常见的兜底)
}


# ---------- 评分:F1 / 内容正确率 / 格式合规率 三者分开 ----------

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()
        if not pt or not gt:
            return 0.0
        common = {w: min(pt.count(w), gt.count(w)) for w in set(pt) if w in gt}
        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 containment(pred, golds):
    """内容正确率:gold 片段(规范化后)是否作为子串出现在输出里 —— 不惩罚啰嗦。"""
    p = _norm(pred)
    return float(any(_norm(g) and _norm(g) in p for g in golds))


def compliance(pred, context):
    """格式合规:输出是上下文的精确子串(规范化后)且不超过 15 词。"""
    p = _norm(pred)
    return float(bool(p) and len(p.split()) <= 15 and p in _norm(context))


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

def cmd_run(args):
    from vllm import LLM, SamplingParams
    model = MODELS[args.config]
    items = load_squad()
    llm = LLM(model=model, gpu_memory_utilization=0.90, max_model_len=4096)
    tok = llm.get_tokenizer()

    def chat(prompts, max_tokens, ignore_eos=False):
        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, ignore_eos=ignore_eos)
        t0 = time.perf_counter()
        outs = llm.generate(texts, sp)
        dt = time.perf_counter() - t0
        outs = sorted(outs, key=lambda o: int(o.request_id))
        ntok = sum(len(o.outputs[0].token_ids) for o in outs)
        return [o.outputs[0].text for o in outs], dt, ntok

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

    # --- 质量:四个条件 ---
    for name, (pfn, mt) in CONDITIONS.items():
        preds, dt, ntok = chat([pfn(x["context"], x["question"]) for x in items], mt)
        rows = []
        for x, p in zip(items, preds):
            t = p.strip()
            rows.append({"id": x["id"], "gold": x["gold_answers"], "pred_text": t,
                         "f1": squad_f1(t, x["gold_answers"]),
                         "contain": containment(t, x["gold_answers"]),
                         "comply": compliance(t, x["context"]),
                         "n_words": len(t.split())})
        results["conditions"][name] = {
            "wall_s": dt, "gen_tokens": ntok,
            "mean_f1": sum(r["f1"] for r in rows) / len(rows),
            "mean_contain": sum(r["contain"] for r in rows) / len(rows),
            "mean_comply": sum(r["comply"] for r in rows) / len(rows),
            "items": rows,
        }
        print(f"[{args.config}/{name}] F1={results['conditions'][name]['mean_f1']:.3f} "
              f"contain={results['conditions'][name]['mean_contain']:.3f} "
              f"comply={results['conditions'][name]['mean_comply']:.3f}")

    # --- 吞吐:并发扫描(固定输出长度,ignore_eos) ---
    base_prompts = [p_baseline(x["context"], x["question"]) for x in items]
    for b in BATCH_POINTS:
        prompts = [base_prompts[i % len(base_prompts)] for i in range(b)]
        chat(prompts[:min(b, 4)], 8)                       # warmup
        _, dt, ntok = chat(prompts, THROUGHPUT_TOKENS, ignore_eos=True)
        results["throughput"][str(b)] = {"wall_s": dt, "gen_tokens": ntok,
                                         "tok_per_s": ntok / dt}
        print(f"[{args.config}/batch={b}] {ntok/dt:.0f} tok/s")

    with open(args.out, "w") as f:
        json.dump(results, f, ensure_ascii=False)
    print("saved", args.out)


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

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


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

    for cond in CONDITIONS:
        F = {r["id"]: r for r in fp16["conditions"][cond]["items"]}
        A = {r["id"]: r for r in awq["conditions"][cond]["items"]}
        ids = sorted(F)
        block = {}
        for metric in ("f1", "contain", "comply"):
            d = [A[i][metric] - F[i][metric] for i in ids]
            lo, hi = bootstrap_ci(d)
            block[metric] = {
                "fp16": sum(F[i][metric] for i in ids) / len(ids),
                "awq": sum(A[i][metric] for i in ids) / len(ids),
                "delta": sum(d) / len(d), "ci95": [lo, hi],
                "significant": (hi < 0) or (lo > 0),
            }
        block["mean_words"] = {
            "fp16": sum(F[i]["n_words"] for i in ids) / len(ids),
            "awq": sum(A[i]["n_words"] for i in ids) / len(ids),
        }
        report["quality"][cond] = block

    for b_ in fp16["throughput"]:
        f_, a_ = fp16["throughput"][b_]["tok_per_s"], awq["throughput"][b_]["tok_per_s"]
        report["throughput"][b_] = {"fp16_tok_per_s": f_, "awq_tok_per_s": a_,
                                    "speedup": a_ / f_}

    # H3:内容无损而 F1 有损?
    base = report["quality"]["baseline"]
    report["H3_damage_is_format_not_content"] = (
        base["f1"]["significant"] and not base["contain"]["significant"])
    # H4:某个缓解条件把 F1 差距做到不显著?
    rescued = [c for c in CONDITIONS if c != "baseline"
               and not report["quality"][c]["f1"]["significant"]]
    report["H4_prompting_rescues"] = {"rescued_by": rescued, "any": bool(rescued)}
    # H5:speedup 是否随并发下降 + 拐点
    sp = [(int(k), v["speedup"]) for k, v in sorted(report["throughput"].items(), key=lambda kv: int(kv[0]))]
    report["H5_throughput"] = {
        "curve": sp,
        "monotone_decreasing": all(sp[i][1] >= sp[i + 1][1] - 0.02 for i in range(len(sp) - 1)),
        "crossover_batch": next((b for b, s in sp if s < 1.2), None),
    }

    print(json.dumps(report, indent=2, ensure_ascii=False))
    with open("stats_round2.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()
