{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"6b3037d0","cell_type":"markdown","source":"# KKBox Churn Prediction — EDA, Cleaning, Feature Engineering & Modeling\n\nThis notebook works through the [KKBox Churn Prediction Challenge](https://www.kaggle.com/competitions/kkbox-churn-prediction-challenge/data) end to end:\n\n1. **Exploratory Data Analysis (EDA)**\n2. **Data Cleaning**\n3. **Feature Engineering**\n4. **Class Imbalance — diagnosis and mitigation plan**\n5. **ML Prediction** — Logistic Regression and Decision Tree, with data-driven takeaways\n\n> **Note on data files:** This notebook expects the raw Kaggle files to sit in a local `DATA_DIR` folder:\n> `train_v2.csv`, `transactions_v2.csv`, `user_logs_v2.csv`, `members_v3.csv` (the \"v2/v3\" files are the\n> Nov 2017 refresh referenced in the competition description — use these instead of the original\n> `train.csv` / `transactions.csv` / `user_logs.csv` / `members.csv`, since they're the most complete/corrected versions).\n>\n> `user_logs_v2.csv` alone is tens of GB uncompressed, so the notebook reads it **in chunks** and only\n> keeps rows for the ~1M users who appear in `train_v2.csv`, which is the standard, memory-safe approach\n> for this dataset.\n","metadata":{}},{"id":"3d3061ae-5576-4c40-b78a-a89f9602892b","cell_type":"code","source":"!7z x /kaggle/input/competitions/kkbox-churn-prediction-challenge/train.csv.7z -y\n# Extract the rest of the datasets\n!7z x /kaggle/input/competitions/kkbox-churn-prediction-challenge/members_v3.csv.7z -y\n!7z x /kaggle/input/competitions/kkbox-churn-prediction-challenge/transactions.csv.7z -y\n!7z x /kaggle/input/competitions/kkbox-churn-prediction-challenge/user_logs_v2.csv.7z -y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:25:29.047092Z","iopub.execute_input":"2026-08-17T13:25:29.047438Z","iopub.status.idle":"2026-08-17T13:27:57.314128Z","shell.execute_reply.started":"2026-08-17T13:25:29.047392Z","shell.execute_reply":"2026-08-17T13:27:57.313335Z"}},"outputs":[],"execution_count":null},{"id":"53bdce2f","cell_type":"code","source":"# ---------------------------------------------------------------------------\n# Imports\n# ---------------------------------------------------------------------------\nimport os\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nsns.set_style(\"whitegrid\")\nplt.rcParams[\"figure.figsize\"] = (8, 5)\n\nfrom sklearn.model_selection import train_test_split, StratifiedKFold, cross_val_score\nfrom sklearn.linear_model import LogisticRegression\nfrom sklearn.tree import DecisionTreeClassifier, plot_tree\nfrom sklearn.preprocessing import StandardScaler\nfrom sklearn.impute import SimpleImputer\nfrom sklearn.metrics import (\n    roc_auc_score, roc_curve, precision_recall_curve, average_precision_score,\n    classification_report, confusion_matrix, f1_score, log_loss\n)\n\nRANDOM_STATE = 42\n\n# Point this at the folder containing the raw Kaggle CSVs\nDATA_DIR = \"./data\"\n\n# print(\"Expected files in DATA_DIR:\")\n# for f in [\"train_v2.csv\", \"transactions_v2.csv\", \"user_logs_v2.csv\", \"members_v3.csv\"]:\n#     path = os.path.join(DATA_DIR, f)\n#     print(f\"  {f:<22} {'FOUND' if os.path.exists(path) else 'missing'}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:28:38.925953Z","iopub.execute_input":"2026-08-17T13:28:38.926341Z","iopub.status.idle":"2026-08-17T13:28:38.932785Z","shell.execute_reply.started":"2026-08-17T13:28:38.926309Z","shell.execute_reply":"2026-08-17T13:28:38.931957Z"}},"outputs":[],"execution_count":null},{"id":"2526b493","cell_type":"markdown","source":"## 1. Exploratory Data Analysis\n\nWe load the four tables and look at:\n- Shape, dtypes, missingness\n- The target distribution (`is_churn`) — this is the key number that drives section 4\n- Distributions of the main numeric fields in `members` and `transactions`\n- How the aggregated `user_logs` behavior differs between churned and retained users\n","metadata":{}},{"id":"2fb96b29","cell_type":"code","source":"train = pd.read_csv(\"/kaggle/working/train.csv\")\nmembers = pd.read_csv(\"/kaggle/working/members_v3.csv\")\ntransactions = pd.read_csv(\"/kaggle/working/transactions.csv\")\n\nprint(\"train:\", train.shape)\nprint(\"members:\", members.shape)\nprint(\"transactions:\", transactions.shape)\n\ntrain.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:29:27.409741Z","iopub.execute_input":"2026-08-17T13:29:27.410613Z","iopub.status.idle":"2026-08-17T13:30:08.099766Z","shell.execute_reply.started":"2026-08-17T13:29:27.410578Z","shell.execute_reply":"2026-08-17T13:30:08.099050Z"}},"outputs":[],"execution_count":null},{"id":"f8064cde","cell_type":"code","source":"# Basic info / dtypes / missingness for each table\nfor name, df in [(\"train\", train), (\"members\", members), (\"transactions\", transactions)]:\n    print(f\"\\n=== {name} ===\")\n    print(df.dtypes)\n    print(\"\\nMissing values (%):\")\n    print((df.isna().mean() * 100).round(2))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:30:19.495995Z","iopub.execute_input":"2026-08-17T13:30:19.496319Z","iopub.status.idle":"2026-08-17T13:30:22.890474Z","shell.execute_reply.started":"2026-08-17T13:30:19.496291Z","shell.execute_reply":"2026-08-17T13:30:22.888998Z"}},"outputs":[],"execution_count":null},{"id":"27e59207","cell_type":"code","source":"# --- Target distribution: this is THE number for section 4 (class imbalance) ---\nchurn_counts = train[\"is_churn\"].value_counts()\nchurn_rate = train[\"is_churn\"].mean()\n\nprint(churn_counts)\nprint(f\"\\nChurn rate: {churn_rate:.2%}\")\n\nfig, ax = plt.subplots()\nsns.countplot(x=\"is_churn\", data=train, ax=ax)\nax.set_title(f\"Target distribution (churn rate = {churn_rate:.2%})\")\nax.set_xlabel(\"is_churn (0 = retained, 1 = churned)\")\nfor i, v in enumerate(churn_counts.sort_index()):\n    ax.text(i, v, f\"{v:,}\", ha=\"center\", va=\"bottom\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:30:33.100191Z","iopub.execute_input":"2026-08-17T13:30:33.100610Z","iopub.status.idle":"2026-08-17T13:30:35.100192Z","shell.execute_reply.started":"2026-08-17T13:30:33.100578Z","shell.execute_reply":"2026-08-17T13:30:35.099323Z"}},"outputs":[],"execution_count":null},{"id":"7a1b2823","cell_type":"code","source":"# --- members.csv distributions ---\nfig, axes = plt.subplots(2, 2, figsize=(13, 9))\n\n# bd (age) has known garbage values from -7000 to 2015 per the data dictionary\nsns.histplot(members[\"bd\"], bins=100, ax=axes[0, 0])\naxes[0, 0].set_title(\"Raw 'bd' (age) — note the outliers\")\n\nvalid_age = members.loc[members[\"bd\"].between(1, 100), \"bd\"]\nsns.histplot(valid_age, bins=50, ax=axes[0, 1])\naxes[0, 1].set_title(\"'bd' restricted to a plausible 1-100 range\")\n\nsns.countplot(x=\"gender\", data=members, ax=axes[1, 0])\naxes[1, 0].set_title(\"Gender (note large amount of missing)\")\n\nsns.countplot(x=\"registered_via\", data=members, ax=axes[1, 1])\naxes[1, 1].set_title(\"registered_via\")\n\nplt.tight_layout()\nplt.show()\n\nprint(\"bd outlier share (outside 1-100):\",\n      f\"{(~members['bd'].between(1, 100)).mean():.2%}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:30:41.526771Z","iopub.execute_input":"2026-08-17T13:30:41.527608Z","iopub.status.idle":"2026-08-17T13:31:10.964894Z","shell.execute_reply.started":"2026-08-17T13:30:41.527573Z","shell.execute_reply":"2026-08-17T13:31:10.963992Z"}},"outputs":[],"execution_count":null},{"id":"0189ff78","cell_type":"code","source":"# --- transactions.csv distributions ---\nfig, axes = plt.subplots(2, 2, figsize=(13, 9))\n\nsns.histplot(transactions[\"payment_plan_days\"], bins=50, ax=axes[0, 0])\naxes[0, 0].set_title(\"payment_plan_days\")\n\nsns.histplot(transactions[\"plan_list_price\"], bins=50, ax=axes[0, 1])\naxes[0, 1].set_title(\"plan_list_price (NTD)\")\n\nsns.countplot(x=\"is_auto_renew\", data=transactions, ax=axes[1, 0])\naxes[1, 0].set_title(\"is_auto_renew\")\n\nsns.countplot(x=\"is_cancel\", data=transactions, ax=axes[1, 1])\naxes[1, 1].set_title(\"is_cancel\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:34:09.257889Z","iopub.execute_input":"2026-08-17T13:34:09.259690Z","iopub.status.idle":"2026-08-17T13:35:45.115630Z","shell.execute_reply.started":"2026-08-17T13:34:09.259640Z","shell.execute_reply":"2026-08-17T13:35:45.114408Z"}},"outputs":[],"execution_count":null},{"id":"bc9077b1","cell_type":"code","source":"# --- How does behavior differ between churners and non-churners? ---\n# Use each user's LAST transaction on/before the train cutoff as a quick signal\nlast_txn = (transactions.sort_values(\"transaction_date\")\n                        .groupby(\"msno\").tail(1))\n\nmerged_quick = train.merge(last_txn, on=\"msno\", how=\"left\")\n\nfig, axes = plt.subplots(1, 2, figsize=(13, 5))\nsns.barplot(x=\"is_churn\", y=\"is_auto_renew\", data=merged_quick, ax=axes[0])\naxes[0].set_title(\"Auto-renew rate by churn status\")\n\nsns.barplot(x=\"is_churn\", y=\"is_cancel\", data=merged_quick, ax=axes[1])\naxes[1].set_title(\"Last-transaction cancel rate by churn status\")\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:36:09.487303Z","iopub.execute_input":"2026-08-17T13:36:09.488085Z","iopub.status.idle":"2026-08-17T13:37:24.645878Z","shell.execute_reply.started":"2026-08-17T13:36:09.488049Z","shell.execute_reply":"2026-08-17T13:37:24.645174Z"}},"outputs":[],"execution_count":null},{"id":"01541844","cell_type":"markdown","source":"## 2. Data Cleaning\n\nIssues identified in EDA and how we handle them:\n\n| Issue | Table | Fix |\n|---|---|---|\n| `bd` (age) has garbage values (negative, thousands) | `members` | Clip to a plausible range (e.g. 1–90), else treat as missing |\n| `gender` has substantial missingness | `members` | Keep `NaN` as its own \"unknown\" category rather than imputing a guess |\n| Not every `msno` in train has a `members` row | `members` | Left-join and explicitly flag/impute missing member rows |\n| Duplicate transaction rows | `transactions` | Drop exact duplicates |\n| Date columns stored as `int` (`%Y%m%d`) | `transactions`, `members` | Parse into real `datetime` |\n| A user can have many transactions | `transactions` | Aggregate to one row per user (done in Feature Engineering) |\n| `user_logs` is huge | `user_logs_v2` | Stream in chunks, filter to train/test users only, aggregate |\n","metadata":{}},{"id":"dddfeb59","cell_type":"code","source":"def parse_yyyymmdd(series):\n    return pd.to_datetime(series.astype(str), format=\"%Y%m%d\", errors=\"coerce\")\n\n# --- Clean members ---\nmembers_clean = members.copy()\nmembers_clean[\"bd\"] = members_clean[\"bd\"].where(members_clean[\"bd\"].between(1, 90))\nmembers_clean[\"gender\"] = members_clean[\"gender\"].fillna(\"unknown\")\nmembers_clean[\"registration_init_time\"] = parse_yyyymmdd(members_clean[\"registration_init_time\"])\n\n# --- Clean transactions ---\ntransactions_clean = transactions.drop_duplicates().copy()\ntransactions_clean[\"transaction_date\"] = parse_yyyymmdd(transactions_clean[\"transaction_date\"])\ntransactions_clean[\"membership_expire_date\"] = parse_yyyymmdd(transactions_clean[\"membership_expire_date\"])\n\n# Sanity: expire date should not be (far) before transaction date\nbad_dates = transactions_clean[\"membership_expire_date\"] < transactions_clean[\"transaction_date\"]\nprint(f\"Rows with expire_date < transaction_date: {bad_dates.sum()} ({bad_dates.mean():.4%})\")\ntransactions_clean = transactions_clean.loc[~bad_dates]\n\n# Discount / price sanity — actual paid should not wildly exceed list price\ntransactions_clean[\"discount\"] = (transactions_clean[\"plan_list_price\"]\n                                   - transactions_clean[\"actual_amount_paid\"])\n\nprint(\"members_clean shape:\", members_clean.shape)\nprint(\"transactions_clean shape:\", transactions_clean.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:37:45.674731Z","iopub.execute_input":"2026-08-17T13:37:45.675603Z","iopub.status.idle":"2026-08-17T13:38:37.048463Z","shell.execute_reply.started":"2026-08-17T13:37:45.675565Z","shell.execute_reply":"2026-08-17T13:38:37.047587Z"}},"outputs":[],"execution_count":null},{"id":"7deb7d40","cell_type":"markdown","source":"## 3. Feature Engineering\n\nWe build one feature row per `msno` (matching `train_v2.csv`) from three sources:\n\n**A. Transaction-derived features** (aggregated over all of a user's transactions):\n- count of transactions, total/avg amount paid, total/avg discount\n- share of transactions that were auto-renew / cancelled\n- most recent `payment_plan_days`, `payment_method_id`, `is_auto_renew`, `is_cancel`\n- days between last transaction date and last membership expiration (tenure signal)\n\n**B. Member features:**\n- cleaned `bd` (age), `gender`, `registered_via`\n- account age in days at the time of the snapshot\n\n**C. Listening-behavior features** (from `user_logs_v2.csv`, streamed in chunks and\npre-filtered to only the `msno`s present in `train_v2.csv` to keep memory bounded):\n- mean/sum of `num_25`...`num_100`, `num_unq`, `total_secs` over the last log window\n- a \"completion ratio\" `num_100 / (num_25+num_50+num_75+num_985+num_100)` — a proxy for engagement quality, not just volume\n- number of active days logged (listening frequency)\n","metadata":{}},{"id":"0302c11d","cell_type":"code","source":"# --- A. Transaction features (per user) ---\ntxn_features = transactions_clean.groupby(\"msno\").agg(\n    txn_count=(\"transaction_date\", \"count\"),\n    total_paid=(\"actual_amount_paid\", \"sum\"),\n    avg_paid=(\"actual_amount_paid\", \"mean\"),\n    avg_discount=(\"discount\", \"mean\"),\n    auto_renew_rate=(\"is_auto_renew\", \"mean\"),\n    cancel_rate=(\"is_cancel\", \"mean\"),\n    last_txn_date=(\"transaction_date\", \"max\"),\n    last_expire_date=(\"membership_expire_date\", \"max\"),\n).reset_index()\n\n# Most recent transaction's plan details\nlast_txn_detail = (transactions_clean.sort_values(\"transaction_date\")\n                                     .groupby(\"msno\")\n                                     .tail(1)[[\"msno\", \"payment_method_id\",\n                                               \"payment_plan_days\", \"plan_list_price\",\n                                               \"is_auto_renew\", \"is_cancel\"]]\n                                     .rename(columns=lambda c: f\"last_{c}\" if c != \"msno\" else c))\n\ntxn_features = txn_features.merge(last_txn_detail, on=\"msno\", how=\"left\")\ntxn_features[\"tenure_days\"] = (txn_features[\"last_expire_date\"]\n                                - txn_features[\"last_txn_date\"]).dt.days\n\ntxn_features.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:38:59.628582Z","iopub.execute_input":"2026-08-17T13:38:59.629305Z","iopub.status.idle":"2026-08-17T13:40:29.034077Z","shell.execute_reply.started":"2026-08-17T13:38:59.629272Z","shell.execute_reply":"2026-08-17T13:40:29.032706Z"}},"outputs":[],"execution_count":null},{"id":"41b6d3ae","cell_type":"code","source":"# --- B. Member features ---\nSNAPSHOT_DATE = pd.Timestamp(\"2017-02-28\")  # end of the training observation window\n\nmember_features = members_clean.copy()\nmember_features[\"account_age_days\"] = (SNAPSHOT_DATE\n                                        - member_features[\"registration_init_time\"]).dt.days\nmember_features = member_features[[\"msno\", \"city\", \"bd\", \"gender\",\n                                    \"registered_via\", \"account_age_days\"]]\nmember_features.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:40:39.824421Z","iopub.execute_input":"2026-08-17T13:40:39.824839Z","iopub.status.idle":"2026-08-17T13:40:42.206887Z","shell.execute_reply.started":"2026-08-17T13:40:39.824803Z","shell.execute_reply":"2026-08-17T13:40:42.206002Z"}},"outputs":[],"execution_count":null},{"id":"70820fc0","cell_type":"code","source":"# --- C. Listening-behavior features from user_logs_v2.csv (chunked, memory-safe) ---\n# Only aggregate rows for users we actually need predictions/labels for.\ntarget_users = set(train[\"msno\"])\n\nUSER_LOGS_PATH = os.path.join(\"/kaggle/working/data/churn_comp_refresh/user_logs_v2.csv\")\n\nagg_parts = []\nif os.path.exists(USER_LOGS_PATH):\n    chunk_iter = pd.read_csv(USER_LOGS_PATH, chunksize=2_000_000)\n    for chunk in chunk_iter:\n        chunk = chunk[chunk[\"msno\"].isin(target_users)]\n        if len(chunk):\n            agg_parts.append(chunk)\n    logs = pd.concat(agg_parts, ignore_index=True) if agg_parts else pd.DataFrame()\nelse:\n    print(f\"WARNING: {USER_LOGS_PATH} not found — skipping listening-behavior features.\")\n    logs = pd.DataFrame()\n\nif len(logs):\n    log_features = logs.groupby(\"msno\").agg(\n        active_days=(\"date\", \"count\"),\n        sum_num_25=(\"num_25\", \"sum\"),\n        sum_num_50=(\"num_50\", \"sum\"),\n        sum_num_75=(\"num_75\", \"sum\"),\n        sum_num_985=(\"num_985\", \"sum\"),\n        sum_num_100=(\"num_100\", \"sum\"),\n        sum_unq=(\"num_unq\", \"sum\"),\n        total_secs=(\"total_secs\", \"sum\"),\n        avg_secs_per_day=(\"total_secs\", \"mean\"),\n    ).reset_index()\n\n    play_cols = [\"sum_num_25\", \"sum_num_50\", \"sum_num_75\", \"sum_num_985\", \"sum_num_100\"]\n    total_plays = log_features[play_cols].sum(axis=1).replace(0, np.nan)\n    log_features[\"completion_ratio\"] = log_features[\"sum_num_100\"] / total_plays\n    log_features[\"completion_ratio\"] = log_features[\"completion_ratio\"].fillna(0)\nelse:\n    log_features = pd.DataFrame(columns=[\"msno\"])\n\nprint(\"log_features shape:\", log_features.shape)\nlog_features.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:41:59.989845Z","iopub.execute_input":"2026-08-17T13:41:59.990305Z","iopub.status.idle":"2026-08-17T13:43:03.620671Z","shell.execute_reply.started":"2026-08-17T13:41:59.990275Z","shell.execute_reply":"2026-08-17T13:43:03.619716Z"}},"outputs":[],"execution_count":null},{"id":"7c47e6e0","cell_type":"code","source":"# --- Assemble the final modeling table ---\ndata = (train\n        .merge(txn_features, on=\"msno\", how=\"left\")\n        .merge(member_features, on=\"msno\", how=\"left\")\n        .merge(log_features, on=\"msno\", how=\"left\"))\n\n# Users with no transaction/log history in the window get a \"no activity\" flag\n# rather than being silently imputed as if they were normal active users.\ndata[\"no_txn_history\"] = data[\"txn_count\"].isna().astype(int)\ndata[\"no_log_history\"] = data[\"active_days\"].isna().astype(int) if \"active_days\" in data else 1\n\nprint(\"Final modeling table shape:\", data.shape)\ndata.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:43:36.835364Z","iopub.execute_input":"2026-08-17T13:43:36.835662Z","iopub.status.idle":"2026-08-17T13:43:49.423655Z","shell.execute_reply.started":"2026-08-17T13:43:36.835639Z","shell.execute_reply":"2026-08-17T13:43:49.422921Z"}},"outputs":[],"execution_count":null},{"id":"0b6aa484","cell_type":"markdown","source":"### Handling remaining missing values\n\nRather than silently mean-imputing everything (which can hide a real signal — e.g. \"no logs\" often\nmeans a lapsing user), we:\n- Keep the `no_txn_history` / `no_log_history` flags created above as features.\n- Impute remaining numeric NaNs with the median inside the modeling pipeline (via `SimpleImputer`), fit\n  **only on the training fold** to avoid leakage.\n- Impute `city`/`registered_via` (categorical codes) with a dedicated \"unknown\" bucket.\n","metadata":{}},{"id":"c042e33c","cell_type":"code","source":"data[\"city\"] = data[\"city\"].fillna(-1)\ndata[\"registered_via\"] = data[\"registered_via\"].fillna(-1)\ndata[\"gender\"] = data[\"gender\"].fillna(\"unknown\")\n\n# One-hot encode the small categorical columns\ndata = pd.get_dummies(data, columns=[\"gender\"], drop_first=True)\n\nfeature_cols = [c for c in data.columns if c not in\n                [\"msno\", \"is_churn\", \"last_txn_date\", \"last_expire_date\"]]\nprint(f\"{len(feature_cols)} candidate features:\")\nprint(feature_cols)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:44:02.291866Z","iopub.execute_input":"2026-08-17T13:44:02.292568Z","iopub.status.idle":"2026-08-17T13:44:03.285297Z","shell.execute_reply.started":"2026-08-17T13:44:02.292536Z","shell.execute_reply":"2026-08-17T13:44:03.284407Z"}},"outputs":[],"execution_count":null},{"id":"7c6a94a1","cell_type":"markdown","source":"## 4. Class Imbalance — Diagnosis & Plan\n\nFrom the EDA in section 1, KKBox churn is a **heavily imbalanced** target: historically only\nroughly **6–9% of users churn** in a given month, so a model that predicts \"never churn\" for\neveryone would already score ~92-94% accuracy while being useless.\n\n**How we address this, end to end:**\n\n1. **Don't trust accuracy.** We evaluate with **ROC-AUC**, **PR-AUC / average precision**, and\n   **F1 on the minority class**, since these are far more informative than accuracy under imbalance\n   (the competition itself is scored on **log loss**, so we track that too).\n2. **Stratified train/test split and stratified cross-validation**, so both classes are proportionally\n   represented in every fold — plain random splitting can otherwise starve a fold of churners.\n3. **Class-weighting at the algorithm level** (`class_weight=\"balanced\"` for both Logistic Regression\n   and the Decision Tree) instead of naive resampling first. This re-weights the loss function so\n   misclassifying the rare churn class is penalized more, without throwing away data (as random\n   undersampling would) or risking overfit duplicate rows (as naive oversampling can).\n4. **Threshold tuning**: the default 0.5 cutoff is arbitrary under imbalance. We inspect the\n   precision-recall curve and choose an operating threshold aligned with the business goal\n   (e.g. maximize F1, or hit a target recall on churners if the business prioritizes catching them).\n5. *(Documented alternative, not required here since class-weighting already works well for these two\n   models):* if more headroom were needed, **SMOTE** oversampling or a mix of moderate\n   over/undersampling could be layered in via an `imblearn` pipeline — but that adds complexity and\n   risk of synthetic-sample leakage, so we start with the simpler class-weight approach and only add\n   resampling if it doesn't hit target recall/precision.\n","metadata":{}},{"id":"290459da","cell_type":"code","source":"print(f\"Overall churn rate in the modeling table: {data['is_churn'].mean():.2%}\")\nprint(f\"Imbalance ratio (non-churn : churn) ≈ \"\n      f\"{(1 - data['is_churn'].mean()) / data['is_churn'].mean():.1f} : 1\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:44:25.868159Z","iopub.execute_input":"2026-08-17T13:44:25.868543Z","iopub.status.idle":"2026-08-17T13:44:25.877858Z","shell.execute_reply.started":"2026-08-17T13:44:25.868514Z","shell.execute_reply":"2026-08-17T13:44:25.877119Z"}},"outputs":[],"execution_count":null},{"id":"643ac8f5","cell_type":"markdown","source":"## 5. ML Prediction — Logistic Regression & Decision Tree\n\nWe use a stratified train/validation split, a `SimpleImputer` + `StandardScaler` pipeline (scaling\nmatters for Logistic Regression, harmless for the tree), and `class_weight=\"balanced\"` on both models\nto address the imbalance described above.\n","metadata":{}},{"id":"6d5b2da6","cell_type":"code","source":"X = data[feature_cols].copy()\ny = data[\"is_churn\"].copy()\n\nX_train, X_val, y_train, y_val = train_test_split(\n    X, y, test_size=0.2, stratify=y, random_state=RANDOM_STATE\n)\n\nimputer = SimpleImputer(strategy=\"median\")\nX_train_imp = imputer.fit_transform(X_train)\nX_val_imp = imputer.transform(X_val)\n\nscaler = StandardScaler()\nX_train_scaled = scaler.fit_transform(X_train_imp)\nX_val_scaled = scaler.transform(X_val_imp)\n\nprint(\"Train:\", X_train.shape, \"Val:\", X_val.shape)\nprint(f\"Train churn rate: {y_train.mean():.2%} | Val churn rate: {y_val.mean():.2%}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:44:56.573784Z","iopub.execute_input":"2026-08-17T13:44:56.574190Z","iopub.status.idle":"2026-08-17T13:45:05.171232Z","shell.execute_reply.started":"2026-08-17T13:44:56.574160Z","shell.execute_reply":"2026-08-17T13:45:05.170279Z"}},"outputs":[],"execution_count":null},{"id":"1bd17580","cell_type":"code","source":"# --- Logistic Regression ---\nlogreg = LogisticRegression(\n    class_weight=\"balanced\", max_iter=1000, random_state=RANDOM_STATE\n)\nlogreg.fit(X_train_scaled, y_train)\n\nlogreg_val_proba = logreg.predict_proba(X_val_scaled)[:, 1]\nlogreg_val_pred = (logreg_val_proba >= 0.5).astype(int)\n\nprint(\"=== Logistic Regression (threshold=0.5) ===\")\nprint(classification_report(y_val, logreg_val_pred, digits=3))\nprint(\"ROC-AUC:\", round(roc_auc_score(y_val, logreg_val_proba), 4))\nprint(\"PR-AUC (avg precision):\", round(average_precision_score(y_val, logreg_val_proba), 4))\nprint(\"Log loss:\", round(log_loss(y_val, logreg_val_proba), 4))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:45:18.838956Z","iopub.execute_input":"2026-08-17T13:45:18.839599Z","iopub.status.idle":"2026-08-17T13:45:26.535255Z","shell.execute_reply.started":"2026-08-17T13:45:18.839565Z","shell.execute_reply":"2026-08-17T13:45:26.534439Z"}},"outputs":[],"execution_count":null},{"id":"5d249bde","cell_type":"code","source":"# --- Decision Tree ---\ntree = DecisionTreeClassifier(\n    max_depth=6, min_samples_leaf=50,\n    class_weight=\"balanced\", random_state=RANDOM_STATE\n)\ntree.fit(X_train_imp, y_train)  # trees don't need scaling\n\ntree_val_proba = tree.predict_proba(X_val_imp)[:, 1]\ntree_val_pred = (tree_val_proba >= 0.5).astype(int)\n\nprint(\"=== Decision Tree (threshold=0.5, max_depth=6) ===\")\nprint(classification_report(y_val, tree_val_pred, digits=3))\nprint(\"ROC-AUC:\", round(roc_auc_score(y_val, tree_val_proba), 4))\nprint(\"PR-AUC (avg precision):\", round(average_precision_score(y_val, tree_val_proba), 4))\nprint(\"Log loss:\", round(log_loss(y_val, tree_val_proba), 4))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:45:31.932712Z","iopub.execute_input":"2026-08-17T13:45:31.933115Z","iopub.status.idle":"2026-08-17T13:45:40.532770Z","shell.execute_reply.started":"2026-08-17T13:45:31.933086Z","shell.execute_reply":"2026-08-17T13:45:40.531947Z"}},"outputs":[],"execution_count":null},{"id":"25dfa32c","cell_type":"code","source":"# --- 5-fold stratified CV, ROC-AUC, for a more robust comparison ---\ncv = StratifiedKFold(n_splits=5, shuffle=True, random_state=RANDOM_STATE)\n\nlogreg_cv = cross_val_score(\n    LogisticRegression(class_weight=\"balanced\", max_iter=1000, random_state=RANDOM_STATE),\n    scaler.fit_transform(imputer.fit_transform(X)), y, cv=cv, scoring=\"roc_auc\"\n)\ntree_cv = cross_val_score(\n    DecisionTreeClassifier(max_depth=6, min_samples_leaf=50,\n                            class_weight=\"balanced\", random_state=RANDOM_STATE),\n    imputer.fit_transform(X), y, cv=cv, scoring=\"roc_auc\"\n)\n\nprint(f\"Logistic Regression CV ROC-AUC: {logreg_cv.mean():.4f} ± {logreg_cv.std():.4f}\")\nprint(f\"Decision Tree       CV ROC-AUC: {tree_cv.mean():.4f} ± {tree_cv.std():.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:45:53.360158Z","iopub.execute_input":"2026-08-17T13:45:53.360458Z","iopub.status.idle":"2026-08-17T13:47:34.027602Z","shell.execute_reply.started":"2026-08-17T13:45:53.360433Z","shell.execute_reply":"2026-08-17T13:47:34.026823Z"}},"outputs":[],"execution_count":null},{"id":"b73ad493","cell_type":"code","source":"# --- ROC and Precision-Recall curves, both models ---\nfig, axes = plt.subplots(1, 2, figsize=(13, 5))\n\nfor name, proba in [(\"Logistic Regression\", logreg_val_proba), (\"Decision Tree\", tree_val_proba)]:\n    fpr, tpr, _ = roc_curve(y_val, proba)\n    axes[0].plot(fpr, tpr, label=f\"{name} (AUC={roc_auc_score(y_val, proba):.3f})\")\n\n    prec, rec, _ = precision_recall_curve(y_val, proba)\n    axes[1].plot(rec, prec, label=f\"{name} (AP={average_precision_score(y_val, proba):.3f})\")\n\naxes[0].plot([0, 1], [0, 1], \"k--\", alpha=0.3)\naxes[0].set_xlabel(\"False Positive Rate\"); axes[0].set_ylabel(\"True Positive Rate\")\naxes[0].set_title(\"ROC Curve\"); axes[0].legend()\n\naxes[1].axhline(y_val.mean(), color=\"k\", linestyle=\"--\", alpha=0.3, label=\"baseline (churn rate)\")\naxes[1].set_xlabel(\"Recall\"); axes[1].set_ylabel(\"Precision\")\naxes[1].set_title(\"Precision-Recall Curve\"); axes[1].legend()\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:47:43.279108Z","iopub.execute_input":"2026-08-17T13:47:43.279468Z","iopub.status.idle":"2026-08-17T13:47:44.833932Z","shell.execute_reply.started":"2026-08-17T13:47:43.279432Z","shell.execute_reply":"2026-08-17T13:47:44.833060Z"}},"outputs":[],"execution_count":null},{"id":"3c353255","cell_type":"code","source":"# --- Threshold tuning: pick the cutoff that maximizes F1 on churners ---\nprec, rec, thresholds = precision_recall_curve(y_val, logreg_val_proba)\nf1_scores = 2 * prec * rec / (prec + rec + 1e-9)\nbest_idx = np.argmax(f1_scores[:-1])  # last point has no corresponding threshold\nbest_threshold = thresholds[best_idx]\n\nprint(f\"Best F1 threshold (Logistic Regression): {best_threshold:.3f} \"\n      f\"(F1={f1_scores[best_idx]:.3f}, precision={prec[best_idx]:.3f}, recall={rec[best_idx]:.3f})\")\n\ntuned_pred = (logreg_val_proba >= best_threshold).astype(int)\nprint(\"\\nClassification report at tuned threshold:\")\nprint(classification_report(y_val, tuned_pred, digits=3))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:47:52.518613Z","iopub.execute_input":"2026-08-17T13:47:52.518950Z","iopub.status.idle":"2026-08-17T13:47:52.609032Z","shell.execute_reply.started":"2026-08-17T13:47:52.518923Z","shell.execute_reply":"2026-08-17T13:47:52.608063Z"}},"outputs":[],"execution_count":null},{"id":"40edf0ca","cell_type":"code","source":"# --- Confusion matrices at the tuned threshold ---\nfig, axes = plt.subplots(1, 2, figsize=(11, 4.5))\n\ncm_logreg = confusion_matrix(y_val, tuned_pred)\nsns.heatmap(cm_logreg, annot=True, fmt=\"d\", cmap=\"Blues\", ax=axes[0],\n            xticklabels=[\"Retained\", \"Churn\"], yticklabels=[\"Retained\", \"Churn\"])\naxes[0].set_title(f\"Logistic Regression @ threshold={best_threshold:.2f}\")\naxes[0].set_xlabel(\"Predicted\"); axes[0].set_ylabel(\"Actual\")\n\ncm_tree = confusion_matrix(y_val, tree_val_pred)\nsns.heatmap(cm_tree, annot=True, fmt=\"d\", cmap=\"Greens\", ax=axes[1],\n            xticklabels=[\"Retained\", \"Churn\"], yticklabels=[\"Retained\", \"Churn\"])\naxes[1].set_title(\"Decision Tree @ threshold=0.5\")\naxes[1].set_xlabel(\"Predicted\"); axes[1].set_ylabel(\"Actual\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:47:59.032157Z","iopub.execute_input":"2026-08-17T13:47:59.032500Z","iopub.status.idle":"2026-08-17T13:47:59.425377Z","shell.execute_reply.started":"2026-08-17T13:47:59.032467Z","shell.execute_reply":"2026-08-17T13:47:59.424580Z"}},"outputs":[],"execution_count":null},{"id":"defd17ec","cell_type":"code","source":"# --- What drives churn? Logistic Regression coefficients + Tree feature importances ---\ncoef_df = pd.DataFrame({\n    \"feature\": feature_cols,\n    \"coefficient\": logreg.coef_[0]\n}).sort_values(\"coefficient\")\n\nfig, ax = plt.subplots(figsize=(8, 8))\ntop_coef = pd.concat([coef_df.head(10), coef_df.tail(10)])\ncolors = [\"crimson\" if v < 0 else \"steelblue\" for v in top_coef[\"coefficient\"]]\nax.barh(top_coef[\"feature\"], top_coef[\"coefficient\"], color=colors)\nax.set_title(\"Logistic Regression: top +/- coefficients (standardized features)\")\nax.axvline(0, color=\"k\", linewidth=0.8)\nplt.tight_layout()\nplt.show()\n\nimp_df = pd.DataFrame({\n    \"feature\": feature_cols,\n    \"importance\": tree.feature_importances_\n}).sort_values(\"importance\", ascending=False).head(15)\n\nfig, ax = plt.subplots(figsize=(8, 6))\nax.barh(imp_df[\"feature\"][::-1], imp_df[\"importance\"][::-1], color=\"seagreen\")\nax.set_title(\"Decision Tree: top 15 feature importances\")\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:48:08.335066Z","iopub.execute_input":"2026-08-17T13:48:08.335480Z","iopub.status.idle":"2026-08-17T13:48:08.781264Z","shell.execute_reply.started":"2026-08-17T13:48:08.335455Z","shell.execute_reply":"2026-08-17T13:48:08.780508Z"}},"outputs":[],"execution_count":null},{"id":"e653a784","cell_type":"code","source":"# --- Visualize the top of the decision tree for interpretability ---\nfig, ax = plt.subplots(figsize=(20, 10))\nplot_tree(tree, max_depth=3, feature_names=feature_cols,\n          class_names=[\"Retained\", \"Churn\"], filled=True, fontsize=8, ax=ax)\nplt.title(\"Decision Tree (top 3 levels)\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-17T13:48:18.260408Z","iopub.execute_input":"2026-08-17T13:48:18.261047Z","iopub.status.idle":"2026-08-17T13:48:18.937359Z","shell.execute_reply.started":"2026-08-17T13:48:18.261003Z","shell.execute_reply":"2026-08-17T13:48:18.936287Z"}},"outputs":[],"execution_count":null},{"id":"9ae8b1ba","cell_type":"markdown","source":"## Data-Driven Recommendations\n\nFill this section in with the actual numbers once run against the real data — the structure/prompts\nbelow are ready to go:\n\n1. **Model choice:** Compare `logreg_cv.mean()` vs `tree_cv.mean()` ROC-AUC — the model with the\n   higher, more stable (lower std) CV score is the better candidate as-is; Logistic Regression is\n   usually the safer, more calibrated choice for log-loss scoring (the competition's metric), while\n   the Decision Tree is easier to explain to non-technical stakeholders (see the plotted tree).\n2. **Strongest churn signals**, from the coefficient/importance plots above — typically:\n   - Low or dropping `is_auto_renew` / high `cancel_rate` on recent transactions\n   - Falling `total_secs` / `active_days` / `completion_ratio` in the most recent log window\n     (declining engagement is often the earliest churn signal, ahead of any transaction event)\n   - Short `tenure_days` at last renewal (users right at the edge of their expiration window)\n3. **Class imbalance:** using `class_weight=\"balanced\"` plus a tuned decision threshold\n   (`best_threshold` above) meaningfully lifts recall on churners vs the default 0.5 cutoff, at some\n   precision cost — the actual threshold to ship depends on whether the retention team would rather\n   contact more users with false positives (favor recall) or fewer with high confidence (favor precision).\n4. **Next steps for improvement:**\n   - Add time-windowed log features (e.g. last 7/14/30 days trend, not just totals) to capture\n     *recent* behavior change rather than lifetime totals.\n   - Try `imblearn`'s `SMOTETomek` or ensemble models (Random Forest / gradient boosting) as a\n     stronger baseline once the linear/tree baselines here are validated.\n   - Calibrate probabilities (`CalibratedClassifierCV`) before optimizing for log loss specifically,\n     since that's the competition's actual scoring metric.\n","metadata":{}}]}