{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# \ud83d\ude80 Day 5: Guided Lab \u2014 Churn + Value Modeling with Explainable ML\n",
    "\n",
    "## Learning Objectives\n",
    "\n",
    "By the end of this lab, you will be able to:\n",
    "- Build a reusable preprocessing + modeling **pipeline** for tabular data\n",
    "- Train and compare **logistic regression \u2192 random forest \u2192 gradient boosting** for churn\n",
    "- Evaluate models using **ROC-AUC, PR-AUC, and lift by deciles**\n",
    "- Predict **MonthlyCharges** (regression) as a value proxy\n",
    "- Explain predictions using **SHAP** (global + local)\n",
    "- Build a **Revenue-at-Risk** ranked call list: $p(\\text{churn}) \\times \\widehat{\\text{MonthlyCharges}}$"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## \ud83e\udd1d GenAI Copilot (how to use it in this lab)\n",
    "\n",
    "Use GenAI to **speed up** your work, especially for:\n",
    "- creating a robust scikit-learn pipeline (imputation + one-hot encoding + model)\n",
    "- debugging errors (sklearn API changes, SHAP issues)\n",
    "- translating evaluation metrics into a short business interpretation\n",
    "\n",
    "**Do not** use GenAI to fabricate results. Any claim in your write-up must be backed by a computed metric/plot.\n",
    "\n",
    "### Prompt starters\n",
    "- \u201cDraft a scikit-learn ColumnTransformer pipeline for mixed numeric/categorical churn data.\u201d\n",
    "- \u201cExplain why PR-AUC is better than accuracy for imbalanced churn.\u201d\n",
    "- \u201cHelp me interpret this SHAP summary plot in business language.\u201d"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "## Part 0: Setup & Data Load (10 min)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "!pip install -q -U google-genai pandas numpy scikit-learn shap\n",
    "# Optional: if you want AutoML later\n",
    "# !pip install -q -U flaml"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os, json, warnings\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "import matplotlib.pyplot as plt\n",
    "from datetime import datetime, timezone\n",
    "\n",
    "from sklearn.model_selection import train_test_split\n",
    "from sklearn.compose import ColumnTransformer\n",
    "from sklearn.pipeline import Pipeline\n",
    "from sklearn.preprocessing import OneHotEncoder, StandardScaler\n",
    "from sklearn.impute import SimpleImputer\n",
    "\n",
    "from sklearn.linear_model import LogisticRegression, Ridge\n",
    "from sklearn.ensemble import RandomForestClassifier, GradientBoostingClassifier\n",
    "from sklearn.ensemble import RandomForestRegressor, GradientBoostingRegressor\n",
    "\n",
    "from sklearn.metrics import (\n",
    "    roc_auc_score, average_precision_score,\n",
    "    precision_score, recall_score, accuracy_score,\n",
    "    confusion_matrix, ConfusionMatrixDisplay,\n",
    "    mean_absolute_error, mean_squared_error, r2_score,\n",
    "    roc_curve, precision_recall_curve\n",
    ")\n",
    "\n",
    "import shap\n",
    "shap.initjs()\n",
    "\n",
    "warnings.filterwarnings(\"ignore\", category=FutureWarning)\n",
    "RANDOM_STATE = 42"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# \u2500\u2500 GenAI Setup \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\n",
    "from google import genai\n",
    "from google.genai import types\n",
    "\n",
    "try:\n",
    "    from google.colab import userdata\n",
    "    os.environ[\"GEMINI_API_KEY\"] = userdata.get(\"GEMINI_API_KEY\")\n",
    "except Exception:\n",
    "    pass\n",
    "\n",
    "if not os.environ.get(\"GEMINI_API_KEY\"):\n",
    "    import getpass\n",
    "    os.environ[\"GEMINI_API_KEY\"] = getpass.getpass(\"Paste your GEMINI_API_KEY: \")\n",
    "\n",
    "client = genai.Client(api_key=os.environ[\"GEMINI_API_KEY\"])\n",
    "MODEL_ID = \"gemini-2.5-flash-lite\"\n",
    "\n",
    "# \u2500\u2500 Logging Infrastructure \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\n",
    "PROMPT_LOG = []\n",
    "\n",
    "def _now():\n",
    "    return datetime.now(timezone.utc).isoformat(timespec=\"seconds\").replace(\"+00:00\", \"Z\")\n",
    "\n",
    "def log_interaction(role, content, label=None):\n",
    "    \"\"\"Log a GenAI interaction.\"\"\"\n",
    "    entry = {\n",
    "        \"ts\": _now(),\n",
    "        \"role\": role,\n",
    "        \"content\": content if isinstance(content, str) else json.dumps(content),\n",
    "        \"label\": label or \"\",\n",
    "    }\n",
    "    PROMPT_LOG.append(entry)\n",
    "    return entry\n",
    "\n",
    "print(\"\u2705 GenAI + logging ready.\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Load the Telco dataset\n",
    "\n",
    "We\u2019ll load the dataset from a public GitHub mirror. If the download fails (no internet), upload the CSV manually and replace the `url`."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "url = \"https://raw.githubusercontent.com/IBM/telco-customer-churn-on-icp4d/master/data/Telco-Customer-Churn.csv\"\n",
    "df = pd.read_csv(url)\n",
    "print(f\"\u2705 Loaded: {df.shape[0]:,} rows \u00d7 {df.shape[1]} columns\")\n",
    "df.head()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Quick cleaning\n",
    "\n",
    "- Convert `TotalCharges` to numeric (it sometimes contains blanks)\n",
    "- Keep `customerID` for final outputs, but do not use it as a feature"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "df[\"TotalCharges\"] = pd.to_numeric(df[\"TotalCharges\"], errors=\"coerce\")\n",
    "customer_ids = df[\"customerID\"].copy()\n",
    "\n",
    "# Targets\n",
    "y_clf = (df[\"Churn\"] == \"Yes\").astype(int)\n",
    "y_reg = df[\"MonthlyCharges\"].astype(float)\n",
    "\n",
    "# Features (drop customerID and Churn; keep MonthlyCharges for classification)\n",
    "X = df.drop(columns=[\"customerID\", \"Churn\"])\n",
    "\n",
    "print(f\"Missing values:\\n{X.isna().sum().sort_values(ascending=False).head(5)}\")\n",
    "print(f\"\\nChurn rate: {y_clf.mean():.1%}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Train/test split\n",
    "\n",
    "- For churn we use a **stratified** split (keeps churn rate similar in train and test).\n",
    "- We will reuse the same split indices for the regression task."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "X_train, X_test, y_train, y_test = train_test_split(\n",
    "    X, y_clf,\n",
    "    test_size=0.25,\n",
    "    random_state=RANDOM_STATE,\n",
    "    stratify=y_clf\n",
    ")\n",
    "\n",
    "# Use the same row selection for regression\n",
    "yreg_train = y_reg.loc[X_train.index]\n",
    "yreg_test = y_reg.loc[X_test.index]\n",
    "\n",
    "print(f\"Train: {X_train.shape[0]:,} rows | Test: {X_test.shape[0]:,} rows\")\n",
    "print(f\"Churn rate \u2014 train: {y_train.mean():.3f} | test: {y_test.mean():.3f}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "## Part 1: Preprocessing Pipeline (20 min)\n",
    "\n",
    "We\u2019ll build a reusable preprocessing component:\n",
    "\n",
    "- **Numeric:** median impute \u2192 standardize\n",
    "- **Categorical:** most-frequent impute \u2192 one-hot encode\n",
    "\n",
    "Then we can plug in different models (logit, RF, boosting) without rewriting preprocessing."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "numeric_cols = X_train.select_dtypes(include=[\"number\"]).columns.tolist()\n",
    "categorical_cols = X_train.select_dtypes(exclude=[\"number\"]).columns.tolist()\n",
    "\n",
    "print(f\"Numeric columns ({len(numeric_cols)}): {numeric_cols}\")\n",
    "print(f\"Categorical columns ({len(categorical_cols)}): {categorical_cols[:8]}...\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# OneHotEncoder API changed across sklearn versions; handle both.\n",
    "try:\n",
    "    ohe = OneHotEncoder(handle_unknown=\"ignore\", sparse_output=False)\n",
    "except TypeError:\n",
    "    ohe = OneHotEncoder(handle_unknown=\"ignore\", sparse=False)\n",
    "\n",
    "numeric_transformer = Pipeline(steps=[\n",
    "    (\"impute\", SimpleImputer(strategy=\"median\")),\n",
    "    (\"scale\", StandardScaler())\n",
    "])\n",
    "\n",
    "categorical_transformer = Pipeline(steps=[\n",
    "    (\"impute\", SimpleImputer(strategy=\"most_frequent\")),\n",
    "    (\"onehot\", ohe)\n",
    "])\n",
    "\n",
    "preprocess = ColumnTransformer(\n",
    "    transformers=[\n",
    "        (\"num\", numeric_transformer, numeric_cols),\n",
    "        (\"cat\", categorical_transformer, categorical_cols)\n",
    "    ],\n",
    "    remainder=\"drop\"\n",
    ")\n",
    "\n",
    "print(\"\u2705 Preprocessor defined.\")\n",
    "preprocess"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "## Part 2: Churn Models + Evaluation (20 min)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Helper functions\n",
    "\n",
    "We\u2019ll compute ROC-AUC, PR-AUC, precision/recall at a chosen threshold, and lift by deciles (a manager-friendly ranking metric)."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def evaluate_classifier(name, model, X_te, y_te, threshold=0.5):\n",
    "    \"\"\"Evaluate a classifier and return metrics + probabilities.\"\"\"\n",
    "    proba = model.predict_proba(X_te)[:, 1]\n",
    "    pred = (proba >= threshold).astype(int)\n",
    "    out = {\n",
    "        \"model\": name,\n",
    "        \"roc_auc\": roc_auc_score(y_te, proba),\n",
    "        \"pr_auc\": average_precision_score(y_te, proba),\n",
    "        \"accuracy\": accuracy_score(y_te, pred),\n",
    "        \"precision\": precision_score(y_te, pred, zero_division=0),\n",
    "        \"recall\": recall_score(y_te, pred, zero_division=0),\n",
    "    }\n",
    "    return out, proba, pred\n",
    "\n",
    "\n",
    "def lift_by_decile(y_true, y_score, n_bins=10):\n",
    "    \"\"\"Compute lift by decile (manager-friendly ranking metric).\"\"\"\n",
    "    tmp = pd.DataFrame({\"y\": y_true, \"score\": y_score}).copy()\n",
    "    tmp[\"decile\"] = pd.qcut(tmp[\"score\"].rank(method=\"first\"), q=n_bins, labels=False) + 1\n",
    "    overall = tmp[\"y\"].mean()\n",
    "    table = (\n",
    "        tmp.groupby(\"decile\")\n",
    "           .agg(n=(\"y\", \"size\"), churn_rate=(\"y\", \"mean\"), avg_score=(\"score\", \"mean\"))\n",
    "           .sort_index(ascending=False)\n",
    "           .reset_index()\n",
    "    )\n",
    "    table[\"lift\"] = table[\"churn_rate\"] / overall\n",
    "    return table, overall\n",
    "\n",
    "\n",
    "def plot_lift(table, overall_rate, title=\"Lift by decile\"):\n",
    "    \"\"\"Plot observed churn rate by decile.\"\"\"\n",
    "    plt.figure(figsize=(7, 4))\n",
    "    plt.plot(table[\"decile\"], table[\"churn_rate\"], marker=\"o\")\n",
    "    plt.axhline(overall_rate, linestyle=\"--\", color=\"grey\", label=f\"Overall ({overall_rate:.1%})\")\n",
    "    plt.gca().invert_xaxis()\n",
    "    plt.xlabel(\"Decile (10 = highest risk)\")\n",
    "    plt.ylabel(\"Observed churn rate\")\n",
    "    plt.title(title)\n",
    "    plt.legend()\n",
    "    plt.show()\n",
    "\n",
    "\n",
    "def eval_regression(name, model, X_te, y_te):\n",
    "    \"\"\"Evaluate a regression model.\"\"\"\n",
    "    pred = model.predict(X_te)\n",
    "    out = {\n",
    "        \"model\": name,\n",
    "        \"mae\": mean_absolute_error(y_te, pred),\n",
    "        \"rmse\": mean_squared_error(y_te, pred, squared=False),\n",
    "        \"r2\": r2_score(y_te, pred)\n",
    "    }\n",
    "    return out, pred\n",
    "\n",
    "print(\"\u2705 Helper functions defined.\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Model 1 \u2014 Logistic Regression (baseline)\n",
    "\n",
    "Logistic regression is a strong baseline because it is fast, stable, and interpretable. We use `class_weight=\"balanced\"` to handle the churn class imbalance."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "logit = Pipeline(steps=[\n",
    "    (\"prep\", preprocess),\n",
    "    (\"clf\", LogisticRegression(max_iter=2000, class_weight=\"balanced\", random_state=RANDOM_STATE))\n",
    "])\n",
    "logit.fit(X_train, y_train)\n",
    "\n",
    "m_logit, p_logit, pred_logit = evaluate_classifier(\"LogisticRegression\", logit, X_test, y_test)\n",
    "m_logit"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Confusion matrix (threshold = 0.5)\n",
    "cm = confusion_matrix(y_test, pred_logit)\n",
    "disp = ConfusionMatrixDisplay(cm)\n",
    "disp.plot(values_format=\"d\")\n",
    "plt.title(\"Logistic Regression \u2014 Confusion Matrix (0.5 threshold)\")\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Model 2 \u2014 Random Forest\n",
    "\n",
    "Random forests often improve performance by capturing nonlinear patterns and interactions."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "rf = Pipeline(steps=[\n",
    "    (\"prep\", preprocess),\n",
    "    (\"clf\", RandomForestClassifier(\n",
    "        n_estimators=400,\n",
    "        random_state=RANDOM_STATE,\n",
    "        n_jobs=-1,\n",
    "        class_weight=\"balanced_subsample\"\n",
    "    ))\n",
    "])\n",
    "rf.fit(X_train, y_train)\n",
    "\n",
    "m_rf, p_rf, pred_rf = evaluate_classifier(\"RandomForest\", rf, X_test, y_test)\n",
    "m_rf"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Model 3 \u2014 Gradient Boosting\n",
    "\n",
    "Boosting often performs very well on tabular data."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "gb = Pipeline(steps=[\n",
    "    (\"prep\", preprocess),\n",
    "    (\"clf\", GradientBoostingClassifier(random_state=RANDOM_STATE))\n",
    "])\n",
    "gb.fit(X_train, y_train)\n",
    "\n",
    "m_gb, p_gb, pred_gb = evaluate_classifier(\"GradientBoosting\", gb, X_test, y_test)\n",
    "m_gb"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Compare models"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "clf_results = pd.DataFrame([m_logit, m_rf, m_gb]).sort_values([\"pr_auc\", \"roc_auc\"], ascending=False)\n",
    "clf_results"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# \u2500\u2500 ROC + PR Curves \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\n",
    "model_probas = {\"LogisticRegression\": p_logit, \"RandomForest\": p_rf, \"GradientBoosting\": p_gb}\n",
    "model_metrics = {\"LogisticRegression\": m_logit, \"RandomForest\": m_rf, \"GradientBoosting\": m_gb}\n",
    "\n",
    "fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5))\n",
    "\n",
    "for name, proba in model_probas.items():\n",
    "    fpr, tpr, _ = roc_curve(y_test, proba)\n",
    "    ax1.plot(fpr, tpr, label=f\"{name} (AUC={model_metrics[name]['roc_auc']:.3f})\")\n",
    "\n",
    "ax1.plot([0, 1], [0, 1], \"k--\", alpha=0.3)\n",
    "ax1.set_xlabel(\"False Positive Rate\")\n",
    "ax1.set_ylabel(\"True Positive Rate\")\n",
    "ax1.set_title(\"ROC Curves\")\n",
    "ax1.legend()\n",
    "ax1.grid(True, alpha=0.3)\n",
    "\n",
    "for name, proba in model_probas.items():\n",
    "    prec, rec, _ = precision_recall_curve(y_test, proba)\n",
    "    ax2.plot(rec, prec, label=f\"{name} (AUC={model_metrics[name]['pr_auc']:.3f})\")\n",
    "\n",
    "baseline = y_test.mean()\n",
    "ax2.axhline(y=baseline, color=\"k\", linestyle=\"--\", alpha=0.3, label=f\"Baseline ({baseline:.2f})\")\n",
    "ax2.set_xlabel(\"Recall\")\n",
    "ax2.set_ylabel(\"Precision\")\n",
    "ax2.set_title(\"Precision-Recall Curves\")\n",
    "ax2.legend()\n",
    "ax2.grid(True, alpha=0.3)\n",
    "\n",
    "plt.tight_layout()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Lift by deciles (manager-friendly ranking view)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Choose the best churn model for the rest of the lab\n",
    "best_name = clf_results.iloc[0][\"model\"]\n",
    "best_model = {\"LogisticRegression\": logit, \"RandomForest\": rf, \"GradientBoosting\": gb}[best_name]\n",
    "best_proba = {\"LogisticRegression\": p_logit, \"RandomForest\": p_rf, \"GradientBoosting\": p_gb}[best_name]\n",
    "\n",
    "lift_tbl, overall = lift_by_decile(y_test, best_proba, n_bins=10)\n",
    "print(f\"Best model: {best_name}\")\n",
    "lift_tbl"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plot_lift(lift_tbl, overall, title=f\"{best_name} \u2014 Lift by decile\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### GenAI: Interpret the model comparison\n",
    "\n",
    "Let\u2019s ask Gemini to help interpret these results in business language."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# \u2500\u2500 GenAI: Interpret Model Comparison \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\n",
    "metrics_summary = clf_results.to_string(index=False)\n",
    "lift_top = lift_tbl.head(3).to_string(index=False)\n",
    "\n",
    "interpret_prompt = f\"\"\"Here are churn model results on a holdout test set:\n",
    "\n",
    "{metrics_summary}\n",
    "\n",
    "Lift by decile (top 3 deciles of the best model):\n",
    "{lift_top}\n",
    "\n",
    "Context: We are building a churn retention campaign. Missing a churner (false negative)\n",
    "costs roughly 5\u00d7 more than contacting a non-churner (false positive).\n",
    "\n",
    "In 3\u20134 sentences, recommend which model to use and why. Mention the precision-recall\n",
    "tradeoff and what the lift chart tells a manager about targeting.\"\"\"\n",
    "\n",
    "log_interaction(\"user\", interpret_prompt, label=\"model_interpretation\")\n",
    "\n",
    "response = client.models.generate_content(model=MODEL_ID, contents=interpret_prompt)\n",
    "interpretation = response.text\n",
    "log_interaction(\"assistant\", interpretation, label=\"model_interpretation\")\n",
    "\n",
    "print(\"\ud83e\udd16 Gemini's Model Recommendation:\")\n",
    "print(\"=\" * 60)\n",
    "print(interpretation)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "## Part 3: Regression \u2014 Predict MonthlyCharges (15 min)\n",
    "\n",
    "We\u2019ll predict **MonthlyCharges** using the same feature set.\n",
    "\n",
    "Important: Do **not** include MonthlyCharges itself as a predictor when forecasting it."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Regression features (drop MonthlyCharges from X)\n",
    "Xr_train = X_train.drop(columns=[\"MonthlyCharges\"])\n",
    "Xr_test = X_test.drop(columns=[\"MonthlyCharges\"])\n",
    "\n",
    "numeric_cols_r = Xr_train.select_dtypes(include=[\"number\"]).columns.tolist()\n",
    "categorical_cols_r = Xr_train.select_dtypes(exclude=[\"number\"]).columns.tolist()\n",
    "\n",
    "preprocess_r = ColumnTransformer(\n",
    "    transformers=[\n",
    "        (\"num\", Pipeline([(\"impute\", SimpleImputer(strategy=\"median\")), (\"scale\", StandardScaler())]), numeric_cols_r),\n",
    "        (\"cat\", Pipeline([(\"impute\", SimpleImputer(strategy=\"most_frequent\")), (\"onehot\", ohe)]), categorical_cols_r),\n",
    "    ]\n",
    ")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Regression models: Ridge \u2192 Random Forest \u2192 Gradient Boosting"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "ridge = Pipeline(steps=[(\"prep\", preprocess_r), (\"reg\", Ridge(alpha=1.0))])\n",
    "ridge.fit(Xr_train, yreg_train)\n",
    "m_ridge, pred_ridge = eval_regression(\"Ridge\", ridge, Xr_test, yreg_test)\n",
    "\n",
    "rfr = Pipeline(steps=[\n",
    "    (\"prep\", preprocess_r),\n",
    "    (\"reg\", RandomForestRegressor(n_estimators=400, random_state=RANDOM_STATE, n_jobs=-1))\n",
    "])\n",
    "rfr.fit(Xr_train, yreg_train)\n",
    "m_rfr, pred_rfr = eval_regression(\"RandomForestRegressor\", rfr, Xr_test, yreg_test)\n",
    "\n",
    "gbr = Pipeline(steps=[\n",
    "    (\"prep\", preprocess_r),\n",
    "    (\"reg\", GradientBoostingRegressor(random_state=RANDOM_STATE))\n",
    "])\n",
    "gbr.fit(Xr_train, yreg_train)\n",
    "m_gbr, pred_gbr = eval_regression(\"GradientBoostingRegressor\", gbr, Xr_test, yreg_test)\n",
    "\n",
    "reg_results = pd.DataFrame([m_ridge, m_rfr, m_gbr]).sort_values([\"mae\", \"rmse\"], ascending=True)\n",
    "reg_results"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "## Part 4: SHAP + Revenue-at-Risk Call List (10 min)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### SHAP (global + local)\n",
    "\n",
    "We\u2019ll explain the **best churn model**. To keep SHAP fast, we transform features using the pipeline\u2019s preprocessing and sample a subset of rows."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Extract transformed matrices + feature names\n",
    "prep_fitted = best_model.named_steps[\"prep\"]\n",
    "clf_fitted = best_model.named_steps[\"clf\"]\n",
    "\n",
    "X_train_enc = prep_fitted.transform(X_train)\n",
    "feature_names = prep_fitted.get_feature_names_out()\n",
    "\n",
    "# Sample for speed\n",
    "sample_idx = np.random.RandomState(RANDOM_STATE).choice(\n",
    "    X_train_enc.shape[0], size=min(800, X_train_enc.shape[0]), replace=False\n",
    ")\n",
    "X_shap = X_train_enc[sample_idx]\n",
    "\n",
    "print(f\"Encoded shape: {X_train_enc.shape}\")\n",
    "print(f\"SHAP sample shape: {X_shap.shape}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# \u2500\u2500 SHAP: Global Summary Plot \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\n",
    "explainer = None\n",
    "try:\n",
    "    explainer = shap.TreeExplainer(clf_fitted)\n",
    "    shap_values = explainer.shap_values(X_shap)\n",
    "except Exception as e:\n",
    "    print(f\"TreeExplainer failed, falling back to shap.Explainer: {e}\")\n",
    "    explainer = shap.Explainer(clf_fitted, X_shap)\n",
    "    shap_values = explainer(X_shap)\n",
    "\n",
    "plt.figure(figsize=(10, 6))\n",
    "try:\n",
    "    # For some classifiers, shap_values is a list; class 1 is churn\n",
    "    if isinstance(shap_values, list):\n",
    "        shap.summary_plot(shap_values[1], X_shap, feature_names=feature_names, max_display=15, show=False)\n",
    "    else:\n",
    "        shap.summary_plot(shap_values, X_shap, feature_names=feature_names, max_display=15, show=False)\n",
    "    plt.title(f\"SHAP Summary \u2014 {best_name}\")\n",
    "    plt.tight_layout()\n",
    "    plt.show()\n",
    "except Exception as e:\n",
    "    print(f\"Could not render summary plot: {e}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Local explanation for one customer"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "row_i = 0\n",
    "x_one = X_shap[row_i:row_i+1]\n",
    "\n",
    "try:\n",
    "    if hasattr(explainer, \"__call__\") and not isinstance(shap_values, list):\n",
    "        exp = explainer(x_one)\n",
    "        shap.plots.waterfall(exp[0])\n",
    "    else:\n",
    "        # TreeExplainer older API\n",
    "        sv = shap_values[1][row_i] if isinstance(shap_values, list) else shap_values[row_i]\n",
    "        base = explainer.expected_value[1] if isinstance(explainer.expected_value, (list, np.ndarray)) else explainer.expected_value\n",
    "        shap.plots._waterfall.waterfall_legacy(base, sv, feature_names=feature_names)\n",
    "except Exception as e:\n",
    "    print(f\"Local explanation failed: {e}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### GenAI: Turn SHAP into a narrative\n",
    "\n",
    "Let\u2019s ask Gemini to interpret the SHAP values for a specific customer in plain business language."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# \u2500\u2500 GenAI: SHAP Narrative \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\n",
    "# Get top SHAP contributors for this customer\n",
    "if isinstance(shap_values, list):\n",
    "    sv_row = shap_values[1][row_i]\n",
    "else:\n",
    "    sv_row = shap_values[row_i]\n",
    "\n",
    "shap_df = pd.DataFrame({\"feature\": feature_names, \"shap_value\": sv_row})\n",
    "shap_df[\"abs_shap\"] = shap_df[\"shap_value\"].abs()\n",
    "top_features = shap_df.nlargest(6, \"abs_shap\")\n",
    "\n",
    "shap_summary = \"\\n\".join(\n",
    "    f\"  - {row['feature']}: SHAP={row['shap_value']:+.3f} ({'increases' if row['shap_value'] > 0 else 'decreases'} churn risk)\"\n",
    "    for _, row in top_features.iterrows()\n",
    ")\n",
    "\n",
    "narrative_prompt = f\"\"\"A customer has been flagged by our churn model.\n",
    "The overall churn rate is {y_test.mean():.0%}.\n",
    "\n",
    "Top factors driving this prediction (SHAP values):\n",
    "{shap_summary}\n",
    "\n",
    "Write a 3-sentence explanation for a retention manager who needs to decide\n",
    "whether to call this customer. Use plain business language, not technical jargon.\"\"\"\n",
    "\n",
    "log_interaction(\"user\", narrative_prompt, label=\"shap_narrative\")\n",
    "\n",
    "response = client.models.generate_content(model=MODEL_ID, contents=narrative_prompt)\n",
    "narrative = response.text\n",
    "log_interaction(\"assistant\", narrative, label=\"shap_narrative\")\n",
    "\n",
    "print(\"\ud83e\udd16 GenAI Narrative for Retention Manager:\")\n",
    "print(\"=\" * 60)\n",
    "print(narrative)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Revenue-at-Risk ranking\n",
    "\n",
    "$$\\text{Revenue at Risk} = p(\\text{churn}) \\times \\widehat{\\text{MonthlyCharges}}$$\n",
    "\n",
    "Then create a ranked call list."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Choose best regression model\n",
    "best_reg_name = reg_results.iloc[0][\"model\"]\n",
    "best_reg_model = {\"Ridge\": ridge, \"RandomForestRegressor\": rfr, \"GradientBoostingRegressor\": gbr}[best_reg_name]\n",
    "\n",
    "# Predictions on the test set\n",
    "p_churn = best_proba\n",
    "pred_value = best_reg_model.predict(Xr_test)\n",
    "pred_value = np.clip(pred_value, 0, None)  # avoid negative predictions\n",
    "\n",
    "call_list = pd.DataFrame({\n",
    "    \"customerID\": customer_ids.loc[X_test.index].values,\n",
    "    \"p_churn\": p_churn,\n",
    "    \"pred_monthly_charges\": pred_value,\n",
    "})\n",
    "call_list[\"revenue_at_risk\"] = call_list[\"p_churn\"] * call_list[\"pred_monthly_charges\"]\n",
    "\n",
    "call_list = call_list.sort_values(\"revenue_at_risk\", ascending=False)\n",
    "print(f\"Top 10 customers by Revenue-at-Risk:\")\n",
    "call_list.head(10)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Save top-N call list\n",
    "TOP_N = 200\n",
    "out_path = \"day5_guided_call_list_top200.csv\"\n",
    "call_list.head(TOP_N).to_csv(out_path, index=False)\n",
    "print(f\"\u2705 Saved top {TOP_N} call list to {out_path}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "## Wrap-up\n",
    "\n",
    "**What you built:**\n",
    "- A reproducible ML pipeline (imputation + encoding + model)\n",
    "- Model progression: logistic regression \u2192 random forest \u2192 gradient boosting\n",
    "- Business evaluation: ROC/PR curves + lift by deciles\n",
    "- Regression for MonthlyCharges as a value proxy\n",
    "- SHAP explanations (global + local) with GenAI narrative\n",
    "- Revenue-at-Risk ranked call list\n",
    "\n",
    "**Next:** In the **Independent Lab**, you\u2019ll choose an extension track (cost targeting, calibration, AutoML, or segment stress test) and turn this into a manager-ready recommendation."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# \u2500\u2500 Export Prompt Log \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\n",
    "if PROMPT_LOG:\n",
    "    log_df = pd.DataFrame(PROMPT_LOG)\n",
    "    log_df.to_csv(\"day5_lab1_prompt_log.csv\", index=False)\n",
    "    print(f\"\u2705 Exported {len(log_df)} log entries to day5_lab1_prompt_log.csv\")\n",
    "else:\n",
    "    print(\"No GenAI interactions logged.\")"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python",
   "version": "3.11.0"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 4
}
