#!/usr/bin/env python3
"""核心复核 1 —— 这把尺子到底测的是不是"患者身份"。

主实验只给出 +0.0237 的差，太小，必须先证明测量本身有效：
  a) 假患者（打乱患者标签）：分组划分退化成随机划分，差应该塌到 ~0
  b) 注入患者特异偏移：差应该明显变大
  c) 折大小匹配的随机划分：排除"KFold 与 GroupKFold 折几何不同"这个混杂
a) 或 b) 不成立，说明主实验那个差不能解释成泄漏。
"""
import gzip, json
import numpy as np
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import KFold, GroupKFold
from sklearn.metrics import f1_score

SEED, N_GENES = 20260827, 2000


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)
    X, y, g = X[keep], ctype[keep], tumor[keep]
    idx = np.argsort(X.var(axis=0))[-N_GENES:]
    X = X[:, idx]
    return (X - X.mean(0)) / (X.std(0) + 1e-8), y, g


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


def gap(X, y, g, seed=SEED):
    rnd = list(KFold(5, shuffle=True, random_state=seed).split(X))
    grp = list(GroupKFold(5).split(X, y, groups=g))
    return f1_of(X, y, rnd) - f1_of(X, y, grp)


def matched_random(X, y, g, seed=SEED):
    """把患者折的大小照搬给随机划分 —— 折几何一样，只差"同患者是否跨侧"。"""
    grp = list(GroupKFold(5).split(X, y, groups=g))
    sizes = [len(te) for _, te in grp]
    order = np.random.default_rng(seed).permutation(len(y))
    splits, start = [], 0
    for size in sizes:
        te = order[start:start + size]
        splits.append((np.setdiff1d(order, te), te))
        start += size
    return f1_of(X, y, splits) - f1_of(X, y, grp)


def main():
    X, y, g = load()
    out = {"observed_gap": gap(X, y, g)}
    print(f"[主实验复现] 差 = {out['observed_gap']:+.4f}", flush=True)

    rng = np.random.default_rng(SEED)
    fake = [gap(X, y, rng.permutation(g)) for _ in range(3)]
    out["fake_patient_gaps"] = fake
    out["fake_patient_mean"] = float(np.mean(fake))
    print(f"[a 假患者] 差 = {np.mean(fake):+.4f}（3 次 {['%+.4f' % v for v in fake]}）"
          f" → 应≈0，实际{'塌了 ✅' if abs(np.mean(fake)) < abs(out['observed_gap']) / 2 else '没塌 ❌'}",
          flush=True)

    spiked = X.copy()
    for pid in set(g):
        m = g == pid
        spiked[m] += rng.normal(0, 1.0, size=(1, X.shape[1])).astype(np.float32)
    out["spiked_gap"] = gap(spiked, y, g)
    print(f"[b 注入患者偏移] 差 = {out['spiked_gap']:+.4f} → 应明显变大，"
          f"实际{'变大 ✅' if out['spiked_gap'] > out['observed_gap'] else '没变大 ❌'}", flush=True)

    out["matched_geometry_gap"] = matched_random(X, y, g)
    print(f"[c 折几何匹配] 差 = {out['matched_geometry_gap']:+.4f} → "
          f"与主实验同号同量级即说明不是折几何造成的", flush=True)

    out["verdict_instrument_valid"] = bool(
        abs(out["fake_patient_mean"]) < abs(out["observed_gap"]) / 2
        and out["spiked_gap"] > out["observed_gap"])
    json.dump(out, open("check1_instrument.json", "w"), ensure_ascii=False, indent=1)
    print(f"\n[结论] 尺子有效 = {out['verdict_instrument_valid']}", flush=True)


if __name__ == "__main__":
    main()
