#!/usr/bin/env python3
"""单细胞患者泄漏：细胞随机划分 vs 患者分组划分，性能被高估多少。

判定规则写在 ../claims.json，跑之前定死。零模型是标签打乱。
置换次数按实测单次拟合耗时自适应，卡住 40 分钟墙钟预算 —— 免费层只烧 CPU。
"""
import gzip, json, time
import numpy as np
from joblib import Parallel, delayed
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import KFold, GroupKFold
from sklearn.metrics import f1_score, accuracy_score

N_GENES, SEED = 2000, 20260827
NULL_BUDGET_S, NULL_MIN, NULL_MAX, N_JOBS = 2400, 40, 200, 3
TYPES = {1: "T", 2: "B", 3: "Macro", 4: "Endo", 5: "CAF", 6: "NK"}


def load():
    with gzip.open("gse72056.txt.gz", "rt") as fh:
        fh.readline()
        tumor = np.array(fh.readline().rstrip("\n").split("\t")[1:], dtype=int)
        malig = np.array(fh.readline().rstrip("\n").split("\t")[1:], dtype=int)
        ctype = np.array(fh.readline().rstrip("\n").split("\t")[1:], dtype=int)
        rows = [line.rstrip("\n").split("\t")[1:] for line in fh]
    X = np.array(rows, dtype=np.float32).T          # 细胞 × 基因
    keep = (malig == 1) & (ctype >= 1) & (ctype <= 6)   # 只用非恶性、类型明确的细胞
    return X[keep], ctype[keep], tumor[keep]


def fit_eval(X, y, splits, seed):
    accs, f1s = [], []
    for tr, te in splits:
        clf = LogisticRegression(max_iter=300, C=0.1, random_state=seed)
        clf.fit(X[tr], y[tr])
        p = clf.predict(X[te])
        accs.append(accuracy_score(y[te], p))
        f1s.append(f1_score(y[te], p, average="macro", zero_division=0))
    return float(np.mean(accs)), float(np.mean(f1s))


def one_perm(X, y, rnd2, grp2, k):
    """一次置换：同一套打乱标签，分别在随机划分和患者划分下评估，返回两者之差。"""
    ys = y.copy()
    np.random.default_rng(SEED + 1000 + k).shuffle(ys)
    return fit_eval(X, ys, rnd2, SEED)[1] - fit_eval(X, ys, grp2, SEED)[1]


def main():
    t0 = time.time()
    X, y, g = load()
    counts = {TYPES[t]: int((y == t).sum()) for t in sorted(set(y))}
    print(f"[data] 非恶性细胞 {len(y)}｜患者 {len(set(g))} 位｜类型 {counts}", flush=True)
    idx = np.argsort(X.var(axis=0))[-N_GENES:]
    X = X[:, idx]
    X = (X - X.mean(0)) / (X.std(0) + 1e-8)
    print(f"[data] 取方差最大的 {N_GENES} 个基因｜载入耗时 {time.time()-t0:.0f}s", flush=True)

    rnd = list(KFold(5, shuffle=True, random_state=SEED).split(X))
    grp = list(GroupKFold(5).split(X, y, groups=g))
    t1 = time.time()
    a_acc, a_f1 = fit_eval(X, y, rnd, SEED)
    b_acc, b_f1 = fit_eval(X, y, grp, SEED)
    per_fit = (time.time() - t1) / 10
    gap_f1, gap_acc = a_f1 - b_f1, a_acc - b_acc
    print(f"\n[A] 细胞随机划分  准确率 {a_acc:.4f}  macro-F1 {a_f1:.4f}", flush=True)
    print(f"[B] 患者分组划分  准确率 {b_acc:.4f}  macro-F1 {b_f1:.4f}", flush=True)
    print(f"[差] 随机划分高估 准确率 +{gap_acc:.4f}  macro-F1 +{gap_f1:.4f}"
          f"（相对高估 {100*gap_f1/max(b_f1,1e-9):.1f}%）｜单次拟合 {per_fit:.1f}s", flush=True)

    n_perm = int(np.clip(NULL_BUDGET_S * N_JOBS / max(per_fit * 4, 1e-6), NULL_MIN, NULL_MAX))
    print(f"[零模型] 预算 {NULL_BUDGET_S}s × {N_JOBS} 并发 → 置换 {n_perm} 次", flush=True)
    rnd2, grp2 = rnd[:2], grp[:2]
    null = np.array(Parallel(n_jobs=N_JOBS, verbose=1)(
        delayed(one_perm)(X, y, rnd2, grp2, k) for k in range(n_perm)))

    p95 = float(np.percentile(null, 95))
    emp_p = float((null >= gap_f1).mean())
    h1, h2 = bool(gap_f1 > 0), bool(gap_f1 > p95)
    print(f"\n[零模型] 标签打乱 {n_perm} 次的 F1 差：均值 {null.mean():+.4f}｜"
          f"95分位 {p95:+.4f}｜最大 {null.max():+.4f}", flush=True)
    print(f"[H1] 随机划分高于患者划分 → {'成立' if h1 else '被推翻'}", flush=True)
    print(f"[H2] 差距超出零模型 95 分位 → {'成立' if h2 else '被推翻'}（经验 p={emp_p:.4f}）", flush=True)

    json.dump({"seed": SEED, "n_cells": int(len(y)), "n_patients": int(len(set(g))),
               "n_genes": N_GENES, "n_perm": int(n_perm), "sec_per_fit": per_fit,
               "random_split": {"acc": a_acc, "macro_f1": a_f1},
               "patient_split": {"acc": b_acc, "macro_f1": b_f1},
               "gap_f1": gap_f1, "gap_acc": gap_acc,
               "relative_overestimate": gap_f1 / max(b_f1, 1e-9),
               "null_mean": float(null.mean()), "null_p95": p95, "empirical_p": emp_p,
               "H1_passes": h1, "H2_passes": h2, "class_counts": counts,
               "null_distribution": null.tolist(),
               "wall_clock_s": round(time.time() - t0)},
              open("results.json", "w"), ensure_ascii=False, indent=1)
    print(f"[done] 总耗时 {round(time.time()-t0)}s", flush=True)


if __name__ == "__main__":
    main()
