{"metadata":{"kernelspec":{"language":"python","name":"python3","display_name":"Python 3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":7163,"databundleVersionId":44582}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install lifelines","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:53:11.426222Z","iopub.execute_input":"2025-12-16T22:53:11.427043Z","iopub.status.idle":"2025-12-16T22:53:20.819491Z","shell.execute_reply.started":"2025-12-16T22:53:11.427014Z","shell.execute_reply":"2025-12-16T22:53:20.818542Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install py7zr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:53:20.821103Z","iopub.execute_input":"2025-12-16T22:53:20.821365Z","iopub.status.idle":"2025-12-16T22:53:25.602091Z","shell.execute_reply.started":"2025-12-16T22:53:20.821341Z","shell.execute_reply":"2025-12-16T22:53:25.601190Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install pycox torchtuples","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:53:25.603237Z","iopub.execute_input":"2025-12-16T22:53:25.603533Z","iopub.status.idle":"2025-12-16T22:53:31.283834Z","shell.execute_reply.started":"2025-12-16T22:53:25.603498Z","shell.execute_reply":"2025-12-16T22:53:31.283137Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load in \n\nimport pandas as pd\nimport numpy as np\nimport torch\nfrom sklearn.model_selection import train_test_split\nfrom sklearn_pandas import DataFrameMapper\nimport torchtuples as tt\nfrom pycox.models import CoxPH\nfrom pycox.evaluation import EvalSurv\n\n# Any results you write to the current directory are saved as output.\nimport os\nprint(os.listdir(\"../input\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:53:31.285555Z","iopub.execute_input":"2025-12-16T22:53:31.285789Z","iopub.status.idle":"2025-12-16T22:53:39.948785Z","shell.execute_reply.started":"2025-12-16T22:53:31.285765Z","shell.execute_reply":"2025-12-16T22:53:39.948136Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls -d /kaggle/input/kkbox-churn-prediction-challenge/*","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:53:39.949747Z","iopub.execute_input":"2025-12-16T22:53:39.950324Z","iopub.status.idle":"2025-12-16T22:53:40.096330Z","shell.execute_reply.started":"2025-12-16T22:53:39.950303Z","shell.execute_reply":"2025-12-16T22:53:40.095608Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport py7zr\nimport pandas as pd\nimport json\n\nINPUT_PATH = \"/kaggle/input/kkbox-churn-prediction-challenge\"\nOUTPUT_PATH = \"/kaggle/working/extracted\"\nos.makedirs(OUTPUT_PATH, exist_ok=True)\n\n\ndef safe_preview_csv(path, n=10):\n    try:\n        df = pd.read_csv(path, nrows=n)\n        display(df)\n    except Exception as e:\n        print(f\"Failed to read CSV: {e}\")\n\n\ndef preview_file(path, preview_rows=10):\n    fname = os.path.basename(path)\n\n    # CSV\n    if fname.endswith(\".csv\"):\n        print(f\"\\n--- {fname} (CSV) ---\")\n        safe_preview_csv(path, preview_rows)\n        return\n\n    # TSV\n    if fname.endswith(\".tsv\"):\n        print(f\"\\n--- {fname} (TSV) ---\")\n        try:\n            df = pd.read_csv(path, sep=\"\\t\", nrows=preview_rows)\n            display(df)\n        except Exception as e:\n            print(f\"Failed to read TSV: {e}\")\n        return\n\n    # JSON\n    if fname.endswith(\".json\"):\n        print(f\"\\n--- {fname} (JSON) ---\")\n        try:\n            with open(path) as j:\n                data = json.load(j)\n            print(json.dumps(data, indent=2)[:2000])\n        except Exception as e:\n            print(f\"Failed to read JSON: {e}\")\n        return\n\n    # Plain text / code formats\n    if fname.endswith((\".txt\", \".md\", \".py\", \".log\", \".scala\")):\n        print(f\"\\n--- {fname} (text) ---\")\n        try:\n            with open(path, \"r\", errors=\"ignore\") as t:\n                print(\"\".join(t.readlines()[:preview_rows]))\n        except Exception as e:\n            print(f\"Failed to read text: {e}\")\n        return\n\n    # FALLBACK FOR BINARY / UNKNOWN\n    print(f\"\\n--- {fname} (unknown/binary) ---\")\n    try:\n        with open(path, \"rb\") as b:\n            print(b.read(200))\n    except Exception as e:\n        print(f\"(unreadable) {e}\")\n\n\ndef inspect_folder(folder, preview_rows=10):\n    for root, dirs, files in os.walk(folder):\n        for f in files:\n            preview_file(os.path.join(root, f), preview_rows)\n\n\ndef inspect_7z(filename, preview_rows=10):\n    src = f\"{INPUT_PATH}/{filename}\"\n    dst = f\"{OUTPUT_PATH}/{filename.replace('.7z', '')}\"\n\n    os.makedirs(dst, exist_ok=True)\n\n    print(f\"\\n=== Extracting {filename} ===\")\n\n    with py7zr.SevenZipFile(src, \"r\") as archive:\n        archive.extractall(path=dst)\n\n    print(\"\\nContents:\")\n    for item in os.listdir(dst):\n        print(\" -\", item)\n\n    print(\"\\n=== PREVIEW ===\")\n    inspect_folder(dst, preview_rows)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:53:40.097422Z","iopub.execute_input":"2025-12-16T22:53:40.097716Z","iopub.status.idle":"2025-12-16T22:53:40.513071Z","shell.execute_reply.started":"2025-12-16T22:53:40.097679Z","shell.execute_reply":"2025-12-16T22:53:40.512276Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"files = [\n    \"train_v2.csv.7z\",\n    \"members_v3.csv.7z\",\n    \"transactions_v2.csv.7z\",\n    \"user_logs_v2.csv.7z\",\n    \"sample_submission_v2.csv.7z\"    \n]\n\nfor f in files:\n    inspect_7z(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:53:40.513896Z","iopub.execute_input":"2025-12-16T22:53:40.514133Z","iopub.status.idle":"2025-12-16T22:55:17.597693Z","shell.execute_reply.started":"2025-12-16T22:53:40.514105Z","shell.execute_reply":"2025-12-16T22:55:17.596976Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls -d /kaggle/working/extracted/*","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:17.598480Z","iopub.execute_input":"2025-12-16T22:55:17.598672Z","iopub.status.idle":"2025-12-16T22:55:17.729169Z","shell.execute_reply.started":"2025-12-16T22:55:17.598655Z","shell.execute_reply":"2025-12-16T22:55:17.728448Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mem_df = pd.read_csv('/kaggle/working/extracted/members_v3.csv/members_v3.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:17.730165Z","iopub.execute_input":"2025-12-16T22:55:17.730439Z","iopub.status.idle":"2025-12-16T22:55:25.351758Z","shell.execute_reply.started":"2025-12-16T22:55:17.730413Z","shell.execute_reply":"2025-12-16T22:55:25.351182Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mem_df[mem_df['msno']==\"moRTKhKIDvb+C8ZHOgmaF4dXMLk0jOn65d7a8tQ2Eds=\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:25.354487Z","iopub.execute_input":"2025-12-16T22:55:25.354704Z","iopub.status.idle":"2025-12-16T22:55:25.822636Z","shell.execute_reply.started":"2025-12-16T22:55:25.354688Z","shell.execute_reply":"2025-12-16T22:55:25.821979Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mem_df['registration_init_time_dt']=pd.to_datetime(mem_df['registration_init_time'], format='%Y%m%d')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:25.823470Z","iopub.execute_input":"2025-12-16T22:55:25.823711Z","iopub.status.idle":"2025-12-16T22:55:26.002719Z","shell.execute_reply.started":"2025-12-16T22:55:25.823693Z","shell.execute_reply":"2025-12-16T22:55:26.001960Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mem_df = mem_df.drop(columns=['registration_init_time'])\nmem_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:26.003568Z","iopub.execute_input":"2025-12-16T22:55:26.003774Z","iopub.status.idle":"2025-12-16T22:55:26.316893Z","shell.execute_reply.started":"2025-12-16T22:55:26.003757Z","shell.execute_reply":"2025-12-16T22:55:26.316268Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys, subprocess\ntry:\n    import pycox, torch  # quick check\nexcept Exception:\n    print(\"Installing pycox + torchtuples + light torch (CPU)...\")\n    !pip install --no-cache-dir pycox torchtuples\n    # install a CPU-only torch if available (pytorch cpu wheels tend to be auto-selected)\n    !pip install --no-cache-dir torch --index-url https://download.pytorch.org/whl/cpu","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:26.317679Z","iopub.execute_input":"2025-12-16T22:55:26.317967Z","iopub.status.idle":"2025-12-16T22:55:26.322868Z","shell.execute_reply.started":"2025-12-16T22:55:26.317938Z","shell.execute_reply":"2025-12-16T22:55:26.322289Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport os\nfrom datetime import datetime\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import StandardScaler\nimport torch\nimport torchtuples as tt\nfrom pycox.models import CoxPH\nfrom pycox.evaluation import EvalSurv\nimport gc  # For garbage collection to free memory","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:26.323583Z","iopub.execute_input":"2025-12-16T22:55:26.323836Z","iopub.status.idle":"2025-12-16T22:55:26.344904Z","shell.execute_reply.started":"2025-12-16T22:55:26.323790Z","shell.execute_reply":"2025-12-16T22:55:26.344408Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Paths ---\nEXTRACTED = \"/kaggle/working/extracted\"\ntrain_v2_path = os.path.join(EXTRACTED, \"train_v2.csv\", \"data\", \"churn_comp_refresh\", \"train_v2.csv\")\ntransactions_v2_path = os.path.join(EXTRACTED, \"transactions_v2.csv\", \"data\", \"churn_comp_refresh\", \"transactions_v2.csv\")\nmembers_path = os.path.join(EXTRACTED, \"members_v3.csv\", \"members_v3.csv\")\n\n# --- Load datasets ---\nprint(\"Loading train...\")\ntrain = pd.read_csv(train_v2_path, usecols=[\"msno\",\"is_churn\"])\nprint(f\"Train shape: {train.shape}\")\n\nprint(\"Loading members...\")\nmembers = pd.read_csv(members_path)\nprint(f\"Members shape: {members.shape}\")\n\n# --- Load transactions (only necessary columns) ---\nprint(\"Loading transactions (this may take a while)...\")\ntransactions = pd.read_csv(\n    transactions_v2_path, \n    usecols=['msno', 'transaction_date', 'membership_expire_date']\n)\nprint(f\"Transactions shape: {transactions.shape}\")\nprint(transactions.head(10))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:26.345484Z","iopub.execute_input":"2025-12-16T22:55:26.345641Z","iopub.status.idle":"2025-12-16T22:55:35.994699Z","shell.execute_reply.started":"2025-12-16T22:55:26.345629Z","shell.execute_reply":"2025-12-16T22:55:35.994000Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Preprocess transactions ---\nprint(\"Processing transactions...\")\ntransactions['transaction_date'] = pd.to_datetime(transactions['transaction_date'], format='%Y%m%d', errors='coerce')\ntransactions['membership_expire_date'] = pd.to_datetime(transactions['membership_expire_date'], format='%Y%m%d', errors='coerce')\n\n# Remove invalid dates\ntransactions = transactions.dropna(subset=['transaction_date', 'membership_expire_date'])\n\n# Aggregate per user\nprint(\"Aggregating transactions per user...\")\ntrans_agg = transactions.groupby('msno', as_index=False).agg(\n    first_trans=('transaction_date', 'min'),\n    last_expire=('membership_expire_date', 'max'),\n    num_transactions=('transaction_date', 'count')  # Transaction count as feature\n)\n\n# Free up memory\ndel transactions\ngc.collect()\nprint(f\"Aggregated transactions shape: {trans_agg.shape}\")\nprint(trans_agg.head(10))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:35.995485Z","iopub.execute_input":"2025-12-16T22:55:35.995763Z","iopub.status.idle":"2025-12-16T22:55:38.079341Z","shell.execute_reply.started":"2025-12-16T22:55:35.995747Z","shell.execute_reply":"2025-12-16T22:55:38.078737Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Merge with train ---\nprint(\"Merging datasets...\")\ndf = train.merge(trans_agg, on='msno', how='left')\ndf = df.merge(mem_df, on='msno', how='left')\nprint(f\"Merged shape: {df.shape}\")\n\n# df table\nprint(df.head(10))\n\n# Free memory\ndel train, trans_agg, members\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:38.080035Z","iopub.execute_input":"2025-12-16T22:55:38.080217Z","iopub.status.idle":"2025-12-16T22:55:47.160332Z","shell.execute_reply.started":"2025-12-16T22:55:38.080203Z","shell.execute_reply":"2025-12-16T22:55:47.159615Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Calculate duration ---\nprint(\"Calculating duration...\")\ndf['duration'] = (df['last_expire'] - df['first_trans']).dt.days\n\n# Handle invalid durations (negative or NaN)\n# For users with no transaction history, use a default of 30 days\ndf['duration'] = df['duration'].apply(lambda x: max(1, x) if pd.notna(x) else 30)\n\n# Event observed\ndf['event_observed'] = df['is_churn'].astype(int)\n\nprint(f\"Duration range: [{df['duration'].min()}, {df['duration'].max()}]\")\nprint(f\"Event rate: {df['event_observed'].mean():.2%}\")\nprint(f\"Rows with NaN first_trans: {df['first_trans'].isna().sum()}\")\nprint(f\"Rows with NaN last_expire: {df['last_expire'].isna().sum()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:47.161151Z","iopub.execute_input":"2025-12-16T22:55:47.161410Z","iopub.status.idle":"2025-12-16T22:55:47.727468Z","shell.execute_reply.started":"2025-12-16T22:55:47.161388Z","shell.execute_reply":"2025-12-16T22:55:47.726839Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Clean member features ---\nprint(\"Cleaning features...\")\n\n# 1. Age (bd) - filter outliers and fill missing\ndf['bd'] = df['bd'].apply(lambda x: x if (x >= 0 and x <= 90) else np.nan)\n\n# 2. Gender - encode and handle missing\ndf['gender'] = df['gender'].map({'male': 1, 'female': 0})\ndf.fillna({'gender': -1}, inplace=True)\n# df['gender'].fillna(-1, inplace=True)  # -1 for unknown\n\n# 4. Registration via - categorical\ndf['registered_via'].fillna(-1, inplace=True)\n\nprint(\"Gender distribution:\")\nprint(df['gender'].value_counts())\nprint(df.head(10))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:47.728253Z","iopub.execute_input":"2025-12-16T22:55:47.728731Z","iopub.status.idle":"2025-12-16T22:55:48.050768Z","shell.execute_reply.started":"2025-12-16T22:55:47.728710Z","shell.execute_reply":"2025-12-16T22:55:48.050021Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Remove rows with invalid duration or missing key features ---\nprint(\"\\n=== Final Data Validation ===\")\ninitial_rows = len(df)\n\n# Remove rows with invalid duration\ndf = df[df['duration'] > 0].copy()\nprint(f\"Removed {initial_rows - len(df)} rows with invalid duration\")\n\n# Remove rows with missing event\ndf = df.dropna(subset=['event_observed'])\nprint(f\"Final dataset: {len(df)} rows\")\n\n# --- Define feature columns ---\nfeature_cols = [\n    'bd', \n    'gender', \n    'city', \n    'registered_via',\n    'num_transactions'\n]\n\n# Add registration features if they exist\nif 'registration_year' in df.columns:\n    feature_cols.extend(['registration_year', 'registration_month', 'account_age_days', 'registration_missing'])\nelif 'has_registration_data' in df.columns:\n    feature_cols.append('has_registration_data')\n\nprint(f\"\\nFeatures to use ({len(feature_cols)}): {feature_cols}\")\n\n# Create final dataframe with only needed columns\ndf_final = df[feature_cols + ['duration', 'event_observed']].copy()\n\n# Final check for any remaining NaNs\nprint(\"\\n=== Missing Values Check ===\")\nmissing_summary = df_final.isnull().sum()\nif missing_summary.sum() > 0:\n    print(\"⚠️ WARNING: Some columns still have missing values:\")\n    print(missing_summary[missing_summary > 0])\n    print(\"\\nFilling remaining NaNs with -1...\")\n    df_final = df_final.fillna(-1)\nelse:\n    print(\"✓ No missing values!\")\n\nprint(\"\\n=== Final Dataset Summary ===\")\nprint(f\"Shape: {df_final.shape}\")\nprint(f\"\\nDuration stats:\")\nprint(df_final['duration'].describe())\nprint(f\"\\nEvent distribution:\")\nprint(df_final['event_observed'].value_counts())\n\n# Free memory\ndel df\ngc.collect()\n\nprint(\"\\n✓ Data cleaning complete!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:48.051506Z","iopub.execute_input":"2025-12-16T22:55:48.051743Z","iopub.status.idle":"2025-12-16T22:55:48.506417Z","shell.execute_reply.started":"2025-12-16T22:55:48.051718Z","shell.execute_reply":"2025-12-16T22:55:48.505692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Sample for faster training (OPTIONAL) ---\n# Comment this out to use full dataset once you verify it works\nSAMPLE_SIZE = 50000  # Adjust based on your memory\n\nif len(df_final) > SAMPLE_SIZE:\n    df_sample = df_final.sample(SAMPLE_SIZE, random_state=42)\n    print(f\"Sampled {SAMPLE_SIZE} rows from {len(df_final)} total rows\")\nelse:\n    df_sample = df_final.copy()\n    print(f\"Using all {len(df_final)} rows\")\n\nprint(f\"\\nSample event rate: {df_sample['event_observed'].mean():.2%}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:48.507226Z","iopub.execute_input":"2025-12-16T22:55:48.507487Z","iopub.status.idle":"2025-12-16T22:55:48.543308Z","shell.execute_reply.started":"2025-12-16T22:55:48.507465Z","shell.execute_reply":"2025-12-16T22:55:48.542735Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Prepare features and target ---\nprint(\"=== Preparing Data for Cox Model ===\")\n\n# Get feature columns (everything except duration and event_observed)\nfeature_cols = [col for col in df_sample.columns if col not in ['duration', 'event_observed']]\n\nprint(f\"Using {len(feature_cols)} features: {feature_cols}\")\n\nX = df_sample[feature_cols].values.astype('float32')\ndurations = df_sample['duration'].values.astype('float32')\nevents = df_sample['event_observed'].values.astype('float32')\n\nprint(f\"\\nData shapes:\")\nprint(f\"  X: {X.shape}\")\nprint(f\"  Durations: {durations.shape}\")\nprint(f\"  Events: {events.shape}\")\n\nprint(f\"\\nData ranges:\")\nprint(f\"  Duration: [{durations.min():.1f}, {durations.max():.1f}] days\")\nprint(f\"  Event rate: {events.mean():.2%}\")\n\n# Check for any remaining NaNs or infs\nif np.isnan(X).any():\n    print(f\"⚠️ WARNING: X contains {np.isnan(X).sum()} NaN values\")\nif np.isinf(X).any():\n    print(f\"⚠️ WARNING: X contains {np.isinf(X).sum()} Inf values\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:48.544002Z","iopub.execute_input":"2025-12-16T22:55:48.544200Z","iopub.status.idle":"2025-12-16T22:55:48.553428Z","shell.execute_reply.started":"2025-12-16T22:55:48.544184Z","shell.execute_reply":"2025-12-16T22:55:48.552833Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Standardize features ---\nprint(\"\\n=== Standardizing Features ===\")\n\nscaler = StandardScaler()\nX_std = scaler.fit_transform(X)\n\nprint(f\"Standardized X shape: {X_std.shape}\")\nprint(f\"Mean of first feature: {X_std[:, 0].mean():.6f} (should be ~0)\")\nprint(f\"Std of first feature: {X_std[:, 0].std():.6f} (should be ~1)\")\n\n# Verify no NaNs after scaling\nif np.isnan(X_std).any():\n    print(f\"⚠️ ERROR: Standardization created {np.isnan(X_std).sum()} NaN values!\")\n    print(\"This usually means a feature had zero variance (all same value)\")\nelse:\n    print(\"✓ No NaN values after standardization\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:48.554166Z","iopub.execute_input":"2025-12-16T22:55:48.554619Z","iopub.status.idle":"2025-12-16T22:55:48.608578Z","shell.execute_reply.started":"2025-12-16T22:55:48.554596Z","shell.execute_reply":"2025-12-16T22:55:48.607863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Train/test split with stratification ---\nprint(\"\\n=== Splitting Train/Test ===\")\n\nX_train, X_test, durations_train, durations_test, events_train, events_test = train_test_split(\n    X_std, durations, events, \n    test_size=0.2, \n    random_state=42, \n    stratify=events\n)\n\nprint(f\"Train size: {X_train.shape[0]:,} samples\")\nprint(f\"Test size: {X_test.shape[0]:,} samples\")\nprint(f\"Train event rate: {events_train.mean():.2%}\")\nprint(f\"Test event rate: {events_test.mean():.2%}\")\n\n# Check duration distributions\nprint(f\"\\nTrain duration: [{durations_train.min():.1f}, {durations_train.max():.1f}]\")\nprint(f\"Test duration: [{durations_test.min():.1f}, {durations_test.max():.1f}]\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:48.609537Z","iopub.execute_input":"2025-12-16T22:55:48.609985Z","iopub.status.idle":"2025-12-16T22:55:48.639913Z","shell.execute_reply.started":"2025-12-16T22:55:48.609945Z","shell.execute_reply":"2025-12-16T22:55:48.639095Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n=== Converting to PyTorch Tensors ===\")\n\nX_train_tensor = torch.tensor(X_train, dtype=torch.float32)\nX_test_tensor = torch.tensor(X_test, dtype=torch.float32)\ndurations_train_tensor = torch.tensor(durations_train, dtype=torch.float32)\nevents_train_tensor = torch.tensor(events_train, dtype=torch.float32)\ndurations_test_tensor = torch.tensor(durations_test, dtype=torch.float32)\nevents_test_tensor = torch.tensor(events_test, dtype=torch.float32)\n\nprint(\"✓ Tensors created successfully!\")\nprint(f\"  X_train: {X_train_tensor.shape}\")\nprint(f\"  X_test: {X_test_tensor.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:48.640723Z","iopub.execute_input":"2025-12-16T22:55:48.641314Z","iopub.status.idle":"2025-12-16T22:55:48.706426Z","shell.execute_reply.started":"2025-12-16T22:55:48.641295Z","shell.execute_reply":"2025-12-16T22:55:48.705757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Build and train Cox model ---\nprint(\"\\n=== Building and Training Cox Model ===\")\n\nin_features = X_train_tensor.shape[1]\n\n# Create neural network for Cox model\nnet = tt.practical.MLPVanilla(\n    in_features, \n    [32, 32],  # Two hidden layers with 32 nodes each\n    out_features=1, \n    batch_norm=True, \n    dropout=0.1\n)\n\n# Create Cox model\nmodel = CoxPH(net, tt.optim.Adam(lr=0.01))\n\nprint(f\"Model architecture: {in_features} -> 32 -> 32 -> 1\")\n\n# Train the model\nprint(\"\\nTraining...\")\nlog = model.fit(\n    X_train_tensor,\n    (durations_train_tensor, events_train_tensor),\n    batch_size=256,\n    epochs=100,\n    verbose=True,\n    val_data=(X_test_tensor, (durations_test_tensor, events_test_tensor)),\n    val_batch_size=256\n)\n\nprint(\"\\n✓ Training complete!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:55:48.707191Z","iopub.execute_input":"2025-12-16T22:55:48.707464Z","iopub.status.idle":"2025-12-16T22:56:43.483022Z","shell.execute_reply.started":"2025-12-16T22:55:48.707442Z","shell.execute_reply":"2025-12-16T22:56:43.482263Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import brier_score_loss\n\ndef compute_brier_score_at_time(surv_df, durations, events, time_point):\n\n    # Get survival probability at time_point\n    if time_point in surv_df.index:\n        surv_prob = surv_df.loc[time_point].values\n    else:\n        # Find closest time point\n        idx = (surv_df.index - time_point).abs().argmin()\n        surv_prob = surv_df.iloc[idx].values\n    \n    # Binary outcome: did event occur before time_point?\n    # For Brier score: 1 if event occurred before time_point, 0 otherwise\n    y_true = np.zeros(len(durations))\n    \n    for i in range(len(durations)):\n        if events[i] == 1 and durations[i] <= time_point:\n            # Event occurred before time_point\n            y_true[i] = 1\n        elif events[i] == 0 and durations[i] <= time_point:\n            # Censored before time_point - exclude from calculation\n            y_true[i] = np.nan\n        else:\n            # Event occurred after time_point or not yet observed\n            y_true[i] = 0\n    \n    # Remove censored observations\n    mask = ~np.isnan(y_true)\n    y_true_clean = y_true[mask]\n    surv_prob_clean = surv_prob[mask]\n    \n    if len(y_true_clean) == 0:\n        return np.nan\n    \n    # Brier score uses (1 - survival_prob) as predicted event probability\n    y_pred = 1 - surv_prob_clean\n    \n    # Calculate Brier score (MSE between predicted and actual)\n    brier = np.mean((y_pred - y_true_clean) ** 2)\n    \n    return brier\n\n# Calculate Brier scores at multiple time points\ntime_points = [30, 60, 90, 120, 180, 365]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:56:43.483926Z","iopub.execute_input":"2025-12-16T22:56:43.484586Z","iopub.status.idle":"2025-12-16T22:56:43.491050Z","shell.execute_reply.started":"2025-12-16T22:56:43.484566Z","shell.execute_reply":"2025-12-16T22:56:43.490247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Evaluate model on TRAINING data ---\nprint(\"\\n\" + \"=\"*60)\nprint(\"EVALUATING ON TRAINING DATA (to check for overfitting)\")\nprint(\"=\"*60)\n\n# Compute baseline hazards\n_ = model.compute_baseline_hazards()\n\n# Predict survival curves on TRAINING data\nsurv_train = model.predict_surv_df(X_train_tensor)\n\n# Create EvalSurv object for training data\nev_train = EvalSurv(surv_train, durations_train, events_train, censor_surv='km')\n\n# 1. C-INDEX on training data\nc_index_train = ev_train.concordance_td()\nprint(f\"\\nC-INDEX (Train): {c_index_train:.4f}\")\n\n# --- Same for TRAINING data ---\nprint(\"\\n=== BRIER SCORES (TRAINING DATA) ===\")\nbrier_scores_train = []\nvalid_times_train = []\n\nfor t in time_points:\n    if t <= durations_train.max() and t >= durations_train.min():\n        try:\n            bs = compute_brier_score_at_time(surv_train, durations_train, events_train, t)\n            if not np.isnan(bs):\n                brier_scores_train.append(bs)\n                valid_times_train.append(t)\n                print(f\"Day {t:3d}: {bs:.4f}\")\n        except Exception as e:\n            print(f\"Day {t:3d}: Could not compute ({e})\")\n\nif len(brier_scores_train) > 0:\n    ibs_train = np.mean(brier_scores_train)\n    print(f\"\\nIntegrated Brier Score (Train): {ibs_train:.4f}\")\nelse:\n    print(\"\\nCould not compute Integrated Brier Score\")\n    \n# 3. Time-dependent AUC at specific time points on training data\nprint(\"\\n=== Time-Dependent AUC at Specific Time Points (Train) ===\")\ntry:\n    from sklearn.metrics import roc_auc_score\n    \n    time_points = [30, 60, 90]\n    \n    for t in time_points:\n        if t <= durations_train.max() and t >= durations_train.min():\n            # Get survival probability at time t\n            surv_prob_at_t = surv_train.loc[t] if t in surv_train.index else surv_train.iloc[(surv_train.index - t).abs().argmin()]\n            \n            # Create binary outcome\n            y_true_at_t = ((durations_train <= t) & (events_train == 1)).astype(int)\n            \n            if len(np.unique(y_true_at_t)) > 1:\n                risk_scores = 1 - surv_prob_at_t.values\n                auc_at_t = roc_auc_score(y_true_at_t, risk_scores)\n                print(f\"AUC at day {t}: {auc_at_t:.4f}\")\n            else:\n                print(f\"AUC at day {t}: Cannot compute (only one class present)\")\n                \nexcept Exception as e:\n    print(f\"Could not compute time-dependent AUC: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:56:43.494138Z","iopub.execute_input":"2025-12-16T22:56:43.494406Z","iopub.status.idle":"2025-12-16T22:56:53.715127Z","shell.execute_reply.started":"2025-12-16T22:56:43.494390Z","shell.execute_reply":"2025-12-16T22:56:53.714454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Evaluate model ---\nprint(\"\\n=== Evaluating Model ===\")\n\n# Compute baseline hazards\n_ = model.compute_baseline_hazards()\n\n# Predict survival curves\nsurv = model.predict_surv_df(X_test_tensor)\n\n# Create EvalSurv object\nev = EvalSurv(surv, durations_test, events_test, censor_surv='km')\n\n# 1. C-INDEX (Concordance Index)\nc_index = ev.concordance_td()\nprint(f\"\\n{'='*50}\")\nprint(f\"C-INDEX: {c_index:.4f}\")\nprint(f\"{'='*50}\")\n\n# 2. INTEGRATED BRIER SCORE\nprint(\"\\n=== BRIER SCORES (TEST DATA) ===\")\nbrier_scores_test = []\nvalid_times_test = []\n\nfor t in time_points:\n    if t <= durations_test.max() and t >= durations_test.min():\n        try:\n            bs = compute_brier_score_at_time(surv, durations_test, events_test, t)\n            if not np.isnan(bs):\n                brier_scores_test.append(bs)\n                valid_times_test.append(t)\n                print(f\"Day {t:3d}: {bs:.4f}\")\n        except Exception as e:\n            print(f\"Day {t:3d}: Could not compute ({e})\")\n\n# Calculate Integrated Brier Score (average across time points)\nif len(brier_scores_test) > 0:\n    ibs_test = np.mean(brier_scores_test)\n    print(f\"\\nIntegrated Brier Score (Test): {ibs_test:.4f}\")\nelse:\n    print(\"\\nCould not compute Integrated Brier Score\")\n\n# Alternative: Time-dependent AUC at specific time points\nprint(\"\\n=== Time-Dependent AUC at Specific Time Points ===\")\ntry:\n    from sklearn.metrics import roc_auc_score\n    \n    # Evaluate AUC at 30, 60, and 90 days\n    time_points = [30, 60, 90]\n    \n    for t in time_points:\n        if t <= durations_test.max() and t >= durations_test.min():\n            # Get survival probability at time t\n            surv_prob_at_t = surv.loc[t] if t in surv.index else surv.iloc[(surv.index - t).abs().argmin()]\n            \n            # Create binary outcome: did event occur before time t?\n            y_true_at_t = ((durations_test <= t) & (events_test == 1)).astype(int)\n            \n            # Only compute AUC if we have both classes\n            if len(np.unique(y_true_at_t)) > 1:\n                # Use 1 - survival probability as risk score\n                risk_scores = 1 - surv_prob_at_t.values\n                auc_at_t = roc_auc_score(y_true_at_t, risk_scores)\n                print(f\"AUC at day {t}: {auc_at_t:.4f}\")\n            else:\n                print(f\"AUC at day {t}: Cannot compute (only one class present)\")\n                \nexcept Exception as e:\n    print(f\"Could not compute time-dependent AUC: {e}\")\n\nprint(\"\\n\" + \"=\"*50)\nprint(\"EVALUATION COMPLETE\")\nprint(\"=\"*50)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:56:53.715930Z","iopub.execute_input":"2025-12-16T22:56:53.716208Z","iopub.status.idle":"2025-12-16T22:56:54.468500Z","shell.execute_reply.started":"2025-12-16T22:56:53.716186Z","shell.execute_reply":"2025-12-16T22:56:54.467708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"COMPARISON: TRAINING vs TEST PERFORMANCE\")\nprint(\"=\"*60)\n\n# Compare C-index\nprint(f\"\\n=== C-INDEX ===\")\nprint(f\"  Train: {c_index_train:.4f}\")\nprint(f\"  Test:  {c_index:.4f}\")\nprint(f\"  Difference: {c_index_train - c_index:.4f}\")\n\nif c_index_train - c_index > 0.05:\n    print(\"  ⚠️  WARNING: Significant overfitting detected!\")\n    print(\"     (Train C-index is >0.05 higher than test)\")\nelif c_index_train - c_index > 0.02:\n    print(\"  ⚠️  Mild overfitting detected\")\n    print(\"     (Train C-index is 0.02-0.05 higher than test)\")\nelse:\n    print(\"  ✅ Good generalization! Model is not overfitting significantly.\")\n\nprint(f\"\\n=== Brier Score ===\")\nprint(f\"  Train: {ibs_train:.4f}\")\nprint(f\"  Test:  {ibs_test:.4f}\")\nprint(f\"  Difference: {ibs_train - ibs_test:.4f}\")\n\n# Compare AUC at day 30 if available\nprint(f\"\\n=== Time-Dependent AUC Comparison ===\")\ntry:\n    for t in [30, 60, 90]:\n        if t <= durations_test.max() and t <= durations_train.max():\n            # Test AUC\n            surv_prob_test = surv.loc[t] if t in surv.index else surv.iloc[(surv.index - t).abs().argmin()]\n            y_true_test = ((durations_test <= t) & (events_test == 1)).astype(int)\n            \n            # Train AUC\n            surv_prob_train = surv_train.loc[t] if t in surv_train.index else surv_train.iloc[(surv_train.index - t).abs().argmin()]\n            y_true_train = ((durations_train <= t) & (events_train == 1)).astype(int)\n            \n            if len(np.unique(y_true_test)) > 1 and len(np.unique(y_true_train)) > 1:\n                auc_test = roc_auc_score(y_true_test, 1 - surv_prob_test.values)\n                auc_train = roc_auc_score(y_true_train, 1 - surv_prob_train.values)\n                \n                print(f\"\\nDay {t}:\")\n                print(f\"  Train AUC: {auc_train:.4f}\")\n                print(f\"  Test AUC:  {auc_test:.4f}\")\n                print(f\"  Difference: {auc_train - auc_test:.4f}\")\nexcept Exception as e:\n    print(f\"Could not compare AUCs: {e}\")\n\nprint(\"\\n\" + \"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T22:56:54.469406Z","iopub.execute_input":"2025-12-16T22:56:54.469691Z","iopub.status.idle":"2025-12-16T22:56:54.533695Z","shell.execute_reply.started":"2025-12-16T22:56:54.469666Z","shell.execute_reply":"2025-12-16T22:56:54.533059Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}