| |
| """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) |
|
|
| |
| adata_train = adata_full[train_idx].copy() |
| |
| gamma_train = np.median(adata_train.layers["gamma"], axis=0) |
|
|
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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() |
|
|