#!/usr/bin/env python3
"""核心复核 2 —— 2.6% 这个数字是不是被"任务太容易"压住了。

主实验准确率 98.8%，离 1.0 只剩 1.2 个点，泄漏再大也没地方涨（天花板效应）。
如果差随任务变难而变大，那 2.6% 就不是普适结论，而是这个易任务的下界 ——
这是最需要知道的一条，也是把这套做法搬到别的数据上时的判断依据。
把任务逐级变难：减基因数、减每类训练细胞数。
"""
import gzip, json
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

SEED = 20260827


def load(n_genes):
    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 score(X, y, splits, cap, seed):
    """cap = 每类最多用多少训练细胞，用来把任务调难。"""
    rng = np.random.default_rng(seed)
    accs = []
    for tr, te in splits:
        if cap:
            sub = np.concatenate([rng.permutation(tr[y[tr] == c])[:cap] for c in np.unique(y[tr])])
            tr = sub if len(np.unique(y[sub])) > 1 else tr
        clf = LogisticRegression(max_iter=300, C=0.1, random_state=seed).fit(X[tr], y[tr])
        accs.append(f1_score(y[te], clf.predict(X[te]), average="macro", zero_division=0))
    return float(np.mean(accs))


def one(n_genes, cap):
    X, y, g = load(n_genes)
    rnd = list(KFold(5, shuffle=True, random_state=SEED).split(X))
    grp = list(GroupKFold(5).split(X, y, groups=g))
    a, b = score(X, y, rnd, cap, SEED), score(X, y, grp, cap, SEED)
    return {"n_genes": n_genes, "cap_per_class": cap, "random_f1": a, "patient_f1": b,
            "gap": a - b, "relative": (a - b) / max(b, 1e-9)}


def main():
    grid = [(2000, None), (2000, 40), (2000, 15), (200, None), (200, 15), (50, 15)]
    rows = Parallel(n_jobs=3, verbose=1)(delayed(one)(n, c) for n, c in grid)
    rows.sort(key=lambda r: r["patient_f1"], reverse=True)
    print(f"\n{'基因数':>6} {'每类训练细胞':>12} {'患者划分F1':>11} {'随机划分F1':>11} "
          f"{'绝对高估':>9} {'相对高估':>9}")
    for r in rows:
        print(f"{r['n_genes']:>6} {str(r['cap_per_class'] or '全部'):>12} "
              f"{r['patient_f1']:>11.4f} {r['random_f1']:>11.4f} "
              f"{r['gap']:>+9.4f} {100*r['relative']:>8.1f}%")
    easiest, hardest = rows[0], rows[-1]
    grows = hardest["gap"] > easiest["gap"]
    print(f"\n[结论] 任务从 F1 {easiest['patient_f1']:.3f} 变难到 {hardest['patient_f1']:.3f} 时，"
          f"高估幅度 {easiest['gap']:+.4f} → {hardest['gap']:+.4f}"
          f"（{'随难度放大 ✅ 2.6% 只是易任务下界' if grows else '未随难度放大'}）")
    json.dump({"grid": rows, "gap_grows_with_difficulty": bool(grows)},
              open("check2_difficulty.json", "w"), ensure_ascii=False, indent=1)


if __name__ == "__main__":
    main()
