# src/cox_train_once.py from __future__ import annotations import os import re from pathlib import Path import numpy as np import pandas as pd from lifelines.exceptions import ConvergenceError from src.survival_utils import prepare_cox_df, fit_cox from src.cox_persist import save_cox_artifacts from src.calibration_utils import bootstrap_c_index_ci def _default_cox_dir() -> Path: """ Hugging Face Spaces persistent volume (if enabled) mounts at /data. Fall back to /tmp otherwise. """ data_root = Path("/data") if data_root.exists() and os.access(str(data_root), os.W_OK): return data_root / "saved_models" return Path("/tmp") / "saved_models" def _sanitize_name(value: str) -> str: """ Convert free-text label into filesystem-safe token. """ value = str(value).strip().lower() value = re.sub(r"[^a-z0-9]+", "_", value) value = re.sub(r"_+", "_", value).strip("_") return value def _normalize_target_name(target_name: str) -> str: """ Map UI / dataset target labels to stable internal names. """ t = str(target_name).strip().lower() if "acute" in t: return "acute_gvhd" if "chronic" in t: return "chronic_gvhd" return _sanitize_name(target_name) def _build_risk_col_name(target_name: str) -> str: """ Build explicit Cox predictor column name for predicted GVHD risk. """ normalized_target = _normalize_target_name(target_name) return f"Predicted_{normalized_target}_risk" def train_and_save_cox( df_surv_ok: pd.DataFrame, preds_surv: np.ndarray, covariates: list[str], target_name: str, prefix: str | None = None, save_dir: Path | str | None = None, ): """ Train and save a Cox proportional hazards model for overall survival. Parameters ---------- df_surv_ok : pd.DataFrame Survival-ready dataframe containing at least: - OS_time_days - Event_clean preds_surv : np.ndarray Predicted probabilities/risk scores from the selected GVHD classifier. covariates : list[str] Clinical covariates to include in the Cox model. target_name : str Example: - "Acute GVHD(<100 days)" - "Chronic GVHD>100 days" prefix : str | None Optional custom prefix for saved artifacts. If None, a target-specific prefix will be generated automatically. save_dir : Path | str | None Output directory for saved artifacts. Returns ------- cph : lifelines.CoxPHFitter df_cox : pd.DataFrame model_path : Path meta_path : Path """ if df_surv_ok is None or df_surv_ok.empty: raise ValueError("df_surv_ok is empty. Cannot train Cox model.") if preds_surv is None: raise ValueError("preds_surv is None. Cannot train Cox model.") if not isinstance(covariates, list) or len(covariates) == 0: raise ValueError("covariates must be a non-empty list.") if not target_name or not str(target_name).strip(): raise ValueError("target_name must be a non-empty string.") preds_surv = np.asarray(preds_surv).ravel() if len(preds_surv) != len(df_surv_ok): raise ValueError( f"Length mismatch: preds_surv has {len(preds_surv)} rows but " f"df_surv_ok has {len(df_surv_ok)} rows." ) normalized_target = _normalize_target_name(target_name) risk_col_name = _build_risk_col_name(target_name) if prefix is None: prefix = f"cox_os_{normalized_target}" # Build Cox dataset (one-hot encoding etc.) # NOTE: prepare_cox_df must accept risk_col_name= df_cox, design_cols, cat_cols = prepare_cox_df( df_surv=df_surv_ok, preds_surv=preds_surv, covariates=covariates, risk_col_name=risk_col_name, ) if df_cox is None or df_cox.empty: raise RuntimeError("prepare_cox_df returned an empty dataframe.") if risk_col_name not in df_cox.columns: raise RuntimeError( f"Expected risk column '{risk_col_name}' not found in prepared Cox dataframe." ) # Drop zero-variance columns to reduce convergence problems zero_var = [c for c in design_cols if c in df_cox.columns and df_cox[c].nunique(dropna=False) <= 1] if zero_var: df_cox = df_cox.drop(columns=zero_var) design_cols = [c for c in design_cols if c not in zero_var] if risk_col_name not in design_cols: raise RuntimeError( f"Expected risk column '{risk_col_name}' not found in design_cols." ) # Fit Cox model try: cph = fit_cox(df_cox) except ConvergenceError as e: raise RuntimeError(f"Cox did not converge: {e}") except Exception as e: raise RuntimeError(f"Cox fit error: {e}") # ------------------------------------------------------------------ # Bootstrap 95% CI for the training-cohort C-index. # Risk scores come from the fitted Cox model's partial hazard prediction. # ------------------------------------------------------------------ try: partial_hazard = cph.predict_partial_hazard(df_cox).values.ravel() c_point, c_lo, c_hi = bootstrap_c_index_ci( durations=df_cox["OS_time_days"].values, events=df_cox["Event_clean"].values, risk_scores=partial_hazard, n_bootstraps=1000, seed=42, ) except Exception: c_point, c_lo, c_hi = (float(cph.concordance_index_), float("nan"), float("nan")) # Resolve output directory out_dir = Path(save_dir) if save_dir is not None else _default_cox_dir() out_dir.mkdir(parents=True, exist_ok=True) # Save artifacts # NOTE: save_cox_artifacts should accept extra_meta= model_path, meta_path = save_cox_artifacts( cph=cph, design_cols=design_cols, cat_cols=cat_cols, covariates=covariates, df_cox=df_cox, out_dir=out_dir, prefix=prefix, extra_meta={ "target_name": target_name, "normalized_target": normalized_target, "risk_col_name": risk_col_name, "model_type": "cox_os", "n_rows": int(len(df_cox)), "n_events": int(df_cox["Event_clean"].sum()) if "Event_clean" in df_cox.columns else None, "duration_col": "OS_time_days", "event_col": "Event_clean", "c_index": float(c_point) if c_point is not None else None, "c_index_ci_low": float(c_lo) if c_lo is not None else None, "c_index_ci_high": float(c_hi) if c_hi is not None else None, }, ) return cph, df_cox, model_path, meta_path