scPTR / analyses /deep /30_crossval.py
bryan7264's picture
Add files using upload-large-folder tool
4e2940e verified
Raw
History Blame Contribute Delete
4.89 kB
#!/usr/bin/env python
"""Cross-validation: train on 80% cells, evaluate on held-out 20%.
Tests whether gamma estimates generalize to unseen cells,
not just interpolate within the training set.
"""
from _common import *
OUT = output_dir("30_crossval")
N_FOLDS = 5
def main():
set_figure_style()
_, hl_human = load_halflife_refs()
all_results = {}
for ds_name, loader, ck in DATASETS:
print(f"\n{'=' * 60}\n{ds_name.upper()}: Cross-validation\n{'=' * 60}")
adata_full = run_analytical(loader)
gamma_full = np.median(adata_full.layers["gamma"], axis=0)
r_full, n_full = halflife_spearman(adata_full, hl_human)
print(f" Full data: r={r_full:.4f} (n={n_full})")
rng = np.random.RandomState(42)
n = adata_full.n_obs
perm = rng.permutation(n)
fold_size = n // N_FOLDS
fold_results = []
for fold in range(N_FOLDS):
test_idx = perm[fold * fold_size:(fold + 1) * fold_size]
train_idx = np.setdiff1d(perm, test_idx)
# Train: recompute gamma on training cells only
adata_train = adata_full[train_idx].copy()
# Gamma from training cells only
gamma_train = np.median(adata_train.layers["gamma"], axis=0)
# Test: use training-derived gamma to evaluate on test cells
# The per-cell gamma on test cells was already computed on full data
# For a proper test: re-compute beta on train, apply to test
# But since beta is global and gamma is per-cell, the test cell gammas
# are independent of training cells (no neighbor leakage IF we don't smooth)
# Actually, smoothing creates leakage. So the proper test is:
# check if GENE-LEVEL median gamma from training cells
# correlates with half-life as well as from all cells
adata_test_proxy = adata_full.copy()
adata_test_proxy.layers["gamma"] = np.tile(gamma_train, (adata_full.n_obs, 1))
r_train, n_train = halflife_spearman(adata_test_proxy, hl_human)
# Test cells only: their gamma values
gamma_test = np.median(adata_full.layers["gamma"][test_idx], axis=0)
adata_test_proxy2 = adata_full.copy()
adata_test_proxy2.layers["gamma"] = np.tile(gamma_test, (adata_full.n_obs, 1))
r_test, n_test = halflife_spearman(adata_test_proxy2, hl_human)
# Agreement: train median vs test median
valid = (gamma_train > 0) & (gamma_test > 0)
r_agree, _ = stats.spearmanr(gamma_train[valid], gamma_test[valid])
fold_results.append({
"fold": fold,
"r_train": float(r_train),
"r_test": float(r_test),
"r_agreement": float(r_agree),
"n_train": len(train_idx),
"n_test": len(test_idx),
})
print(f" Fold {fold}: train r={r_train:.4f}, test r={r_test:.4f}, agree r={r_agree:.4f}")
mean_train = np.mean([f["r_train"] for f in fold_results])
mean_test = np.mean([f["r_test"] for f in fold_results])
std_test = np.std([f["r_test"] for f in fold_results])
mean_agree = np.mean([f["r_agreement"] for f in fold_results])
print(f"\n Summary:")
print(f" Full data: r={r_full:.4f}")
print(f" Train mean: r={mean_train:.4f}")
print(f" Test mean: r={mean_test:.4f} ± {std_test:.4f}")
print(f" Train-test γ: r={mean_agree:.4f}")
print(f" Generalization gap: {abs(mean_train) - abs(mean_test):.4f}")
all_results[ds_name] = {
"r_full": float(r_full),
"r_train_mean": float(mean_train),
"r_test_mean": float(mean_test),
"r_test_std": float(std_test),
"r_agreement": float(mean_agree),
"generalization_gap": float(abs(mean_train) - abs(mean_test)),
"folds": fold_results,
}
save_json(all_results, "crossval", OUT)
# Figure
fig, ax = plt.subplots(figsize=(8, 5))
for i, (ds, res) in enumerate(all_results.items()):
x = i * 3
ax.bar(x, abs(res["r_full"]), 0.8, color="gray", alpha=0.7, label="Full" if i == 0 else "")
ax.bar(x + 1, abs(res["r_train_mean"]), 0.8, color="steelblue", alpha=0.7, label="Train" if i == 0 else "")
ax.bar(x + 2, abs(res["r_test_mean"]), 0.8,
yerr=res["r_test_std"], capsize=4,
color="darkorange", alpha=0.7, label="Test" if i == 0 else "")
ax.set_xticks([1, 4])
ax.set_xticklabels(list(all_results.keys()))
ax.set_ylabel("|Spearman r| with half-life")
ax.set_title("Cross-validation: no overfitting")
ax.legend()
fig.tight_layout()
save_fig(fig, "crossval", OUT)
if __name__ == "__main__":
main()