{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":7163,"databundleVersionId":44582}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\nos.system(\"7z e /kaggle/input/competitions/kkbox-churn-prediction-challenge/transactions_v2.csv.7z -o/kaggle/working/ -y\")\nos.system(\"7z e /kaggle/input/competitions/kkbox-churn-prediction-challenge/members_v3.csv.7z -o/kaggle/working/ -y\")\nos.system(\"7z e /kaggle/input/competitions/kkbox-churn-prediction-challenge/train_v2.csv.7z -o/kaggle/working/ -y\")\nos.system(\"7z e /kaggle/input/competitions/kkbox-churn-prediction-challenge/user_logs_v2.csv.7z -o/kaggle/working/ -y\")\n\nprint(\"Done!\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-04T16:59:05.227553Z","iopub.execute_input":"2026-05-04T16:59:05.227826Z","iopub.status.idle":"2026-05-04T17:00:24.707710Z","shell.execute_reply.started":"2026-05-04T16:59:05.227803Z","shell.execute_reply":"2026-05-04T17:00:24.706807Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip -q install torch==2.4.0 torch-geometric==2.4.0 scikit-learn pandas numpy tqdm\nprint(\"Done!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T17:12:48.361075Z","iopub.execute_input":"2026-05-04T17:12:48.361631Z","iopub.status.idle":"2026-05-04T17:15:22.344873Z","shell.execute_reply.started":"2026-05-04T17:12:48.361601Z","shell.execute_reply":"2026-05-04T17:15:22.343804Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install torch-scatter torch-sparse torch-cluster torch-spline-conv -f https://data.pyg.org/whl/torch-2.4.0+cu121.html\nprint(\"Done!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T17:24:03.062683Z","iopub.execute_input":"2026-05-04T17:24:03.063224Z","iopub.status.idle":"2026-05-04T17:24:08.813823Z","shell.execute_reply.started":"2026-05-04T17:24:03.063183Z","shell.execute_reply":"2026-05-04T17:24:08.812903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# hybrid_gat_mlp_churn_FINAL.py\n# Makamal code — sab graphs aur evaluations included!\n\nimport os\nimport gc\nimport random\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom sklearn.preprocessing import LabelEncoder, StandardScaler\nfrom sklearn.metrics import (roc_auc_score, f1_score, accuracy_score,\n                              confusion_matrix, classification_report,\n                              roc_curve, precision_recall_curve, average_precision_score)\nfrom sklearn.model_selection import train_test_split\n\nfrom torch_geometric.data import Data\nfrom torch_geometric.loader import NeighborLoader\nfrom torch_geometric.nn import GATv2Conv, BatchNorm\n\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nimport matplotlib.gridspec as gridspec\n\n# -------------------------\n# CONFIG\n# -------------------------\nCFG = {\n    \"PATHS\": {\n        \"train\": \"/kaggle/working/train_v2.csv\",\n        \"members\": \"/kaggle/working/members_v3.csv\",\n        \"transactions\": \"/kaggle/working/transactions_v2.csv\",\n        \"user_logs\": \"/kaggle/working/user_logs_v2.csv\",\n    },\n    \"RANDOM_SEED\": 42,\n    \"LOG_CHUNK_ROWS\": 1_000_000,\n    \"MAX_CANCEL_CAP\": 4,\n    \"GRAPH_MAX_NEIGHBORS_PER_GROUP\": 10,\n    \"BATCH_SIZE\": 512,\n    \"LR\": 5e-3,\n    \"WEIGHT_DECAY\": 1e-4,\n    \"MAX_EPOCHS\": 50,\n    \"EARLY_STOP_PATIENCE\": 5,\n    \"GAT_HIDDEN\": 32,\n    \"GAT_HEADS\": 4,\n    \"GAT_DROPOUT\": 0.3,\n    \"MLP_HIDDEN\": 64,\n    \"VAL_RATIO\": 0.15,\n    \"TEST_RATIO\": 0.15,\n    \"THRESHOLD\": 0.89,\n    \"DEVICE\": \"cuda\" if torch.cuda.is_available() else \"cpu\",\n}\n\n# ===== SEEDS FIX =====\nrandom.seed(CFG[\"RANDOM_SEED\"])\nnp.random.seed(CFG[\"RANDOM_SEED\"])\ntorch.manual_seed(CFG[\"RANDOM_SEED\"])\ntorch.cuda.manual_seed_all(CFG[\"RANDOM_SEED\"])\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False\n\n# -------------------------\n# UTILITIES\n# -------------------------\ndef parse_date(series):\n    return pd.to_datetime(series, errors=\"coerce\")\n\n# -------------------------\n# DATA PREPARATION\n# -------------------------\ndef aggregate_transactions(path_transactions):\n    print(\"Aggregating transactions ...\")\n    tx = pd.read_csv(path_transactions)\n    for col in [\"transaction_date\", \"membership_expire_date\"]:\n        if col in tx.columns:\n            tx[col] = parse_date(tx[col])\n    pay_col = \"actual_amount_paid\" if \"actual_amount_paid\" in tx.columns else \"payment_plan_price\"\n    agg = tx.groupby(\"msno\").agg(\n        total_transactions=(\"msno\", \"size\"),\n        total_payment=(pay_col, \"sum\"),\n        is_cancel_sum=(\"is_cancel\", \"sum\"),\n        last_transaction_date=(\"transaction_date\", \"max\")\n    ).reset_index()\n    agg[\"is_cancel_sum\"] = agg[\"is_cancel_sum\"].clip(upper=CFG[\"MAX_CANCEL_CAP\"]).fillna(0).astype(int)\n    agg[\"total_payment\"] = agg[\"total_payment\"].fillna(0.0)\n    return agg\n\ndef aggregate_user_logs(path_logs, chunk_rows):\n    print(\"Aggregating user logs (chunked) ...\")\n    agg_dict = {}\n    chunks = pd.read_csv(path_logs, chunksize=chunk_rows, iterator=True)\n    for chunk in tqdm(chunks, desc=\"user_logs chunks\"):\n        if \"date\" in chunk.columns:\n            chunk[\"date\"] = parse_date(chunk[\"date\"])\n        if \"total_secs\" in chunk.columns:\n            chunk[\"total_secs\"] = pd.to_numeric(chunk[\"total_secs\"], errors=\"coerce\").fillna(0)\n        if \"num_unq\" in chunk.columns:\n            chunk[\"num_unq\"] = pd.to_numeric(chunk[\"num_unq\"], errors=\"coerce\").fillna(0)\n        g = chunk.groupby(\"msno\").agg(\n            log_days=(\"date\", \"nunique\") if \"date\" in chunk.columns else (\"msno\", \"size\"),\n            total_secs_sum=(\"total_secs\", \"sum\") if \"total_secs\" in chunk.columns else (\"msno\", \"size\"),\n            total_songs_played=(\"num_unq\", \"sum\") if \"num_unq\" in chunk.columns else (\"msno\", \"size\"),\n        )\n        for msno, row in g.iterrows():\n            d = agg_dict.get(msno)\n            if d is None:\n                agg_dict[msno] = {\"log_days\": int(row[\"log_days\"]),\n                                   \"total_secs_sum\": float(row[\"total_secs_sum\"]),\n                                   \"total_songs_played\": float(row[\"total_songs_played\"])}\n            else:\n                d[\"log_days\"] += int(row[\"log_days\"])\n                d[\"total_secs_sum\"] += float(row[\"total_secs_sum\"])\n                d[\"total_songs_played\"] += float(row[\"total_songs_played\"])\n        del chunk, g\n        gc.collect()\n    df = pd.DataFrame.from_dict(agg_dict, orient=\"index\").reset_index().rename(columns={\"index\": \"msno\"})\n    df[\"log_days\"] = df[\"log_days\"].astype(int)\n    return df\n\ndef load_members(path_members):\n    print(\"Loading members ...\")\n    mem = pd.read_csv(path_members)\n    if \"registration_init_time\" in mem.columns:\n        mem[\"registration_init_time\"] = pd.to_datetime(\n            mem[\"registration_init_time\"], format=\"%Y%m%d\", errors=\"coerce\")\n    return mem[[\"msno\", \"city\", \"gender\", \"registered_via\", \"registration_init_time\"]]\n\ndef load_train_labels(path_train):\n    tr = pd.read_csv(path_train)[[\"msno\", \"is_churn\"]]\n    tr[\"is_churn\"] = tr[\"is_churn\"].fillna(0).astype(int)\n    return tr\n\ndef build_user_level_table(paths):\n    tx_agg   = aggregate_transactions(paths[\"transactions\"])\n    logs_agg = aggregate_user_logs(paths[\"user_logs\"], CFG[\"LOG_CHUNK_ROWS\"])\n    members  = load_members(paths[\"members\"])\n    labels   = load_train_labels(paths[\"train\"])\n\n    print(\"Merging all sources ...\")\n    df = labels.merge(members, on=\"msno\", how=\"left\") \\\n               .merge(tx_agg,  on=\"msno\", how=\"left\") \\\n               .merge(logs_agg,on=\"msno\", how=\"left\")\n\n    for c in [\"total_transactions\",\"total_payment\",\"is_cancel_sum\",\n              \"log_days\",\"total_secs_sum\",\"total_songs_played\"]:\n        df[c] = pd.to_numeric(df[c], errors=\"coerce\").fillna(0)\n\n    print(\"Feature engineering ...\")\n    df[\"last_transaction_date\"]  = parse_date(df[\"last_transaction_date\"])\n    df[\"registration_init_time\"] = parse_date(df[\"registration_init_time\"])\n    df[\"membership_days\"] = (df[\"last_transaction_date\"] - df[\"registration_init_time\"]).dt.days.fillna(0).clip(lower=0)\n    df[\"registration_year\"]  = df[\"registration_init_time\"].dt.year.fillna(0).astype(int)\n    df[\"registration_month\"] = df[\"registration_init_time\"].dt.month.fillna(0).astype(int)\n\n    for col in [\"city\", \"gender\", \"registered_via\"]:\n        le = LabelEncoder()\n        df[col] = le.fit_transform(df[col].fillna(\"unknown\").astype(str))\n\n    feature_cols = [\"city\",\"gender\",\"registered_via\",\n                    \"total_transactions\",\"total_payment\",\"is_cancel_sum\",\n                    \"log_days\",\"total_secs_sum\",\"total_songs_played\",\n                    \"membership_days\",\"registration_year\",\"registration_month\"]\n\n    df[\"avg_secs_per_day\"]  = (df[\"total_secs_sum\"]    / (df[\"log_days\"] + 1e-6)).fillna(0)\n    df[\"avg_songs_per_day\"] = (df[\"total_songs_played\"] / (df[\"log_days\"] + 1e-6)).fillna(0)\n    df[\"pay_per_tx\"]        = (df[\"total_payment\"]      / (df[\"total_transactions\"] + 1e-6)).fillna(0)\n    feature_cols.extend([\"avg_secs_per_day\",\"avg_songs_per_day\",\"pay_per_tx\"])\n\n    print(\"Scaling numeric features ...\")\n    scaler = StandardScaler()\n    df[feature_cols] = scaler.fit_transform(df[feature_cols].astype(float))\n    return df, feature_cols, scaler\n\n# -------------------------\n# GRAPH CONSTRUCTION\n# -------------------------\ndef build_sparse_similarity_edges(df, max_neighbors=10):\n    print(\"Building sparse similarity graph ...\")\n    df = df.reset_index(drop=True)\n    df[\"node_id\"] = np.arange(len(df))\n    edges_src, edges_dst = [], []\n    for _, g in tqdm(df.groupby([\"city\",\"registered_via\"]),\n                     total=df.groupby([\"city\",\"registered_via\"]).ngroups):\n        ids = g[\"node_id\"].to_numpy()\n        if len(ids) <= 1:\n            continue\n        K = min(max_neighbors, len(ids) - 1)\n        for i in range(len(ids)):\n            for k in range(1, K + 1):\n                j = (i + k) % len(ids)\n                edges_src.append(ids[i]); edges_dst.append(ids[j])\n                edges_src.append(ids[j]); edges_dst.append(ids[i])\n    edge_index = torch.tensor([edges_src, edges_dst], dtype=torch.long)\n    print(f\"Edges built: {edge_index.size(1):,}\")\n    return edge_index, df[\"node_id\"].values\n\n# -------------------------\n# MODEL\n# -------------------------\nclass HybridGATMLP(nn.Module):\n    def __init__(self, in_dim_tabular, gat_in_dim, gat_hidden=32, gat_heads=4,\n                 gat_dropout=0.3, mlp_hidden=64):\n        super().__init__()\n        self.gat1 = GATv2Conv(gat_in_dim, gat_hidden, heads=gat_heads, dropout=gat_dropout, concat=True)\n        self.bn1  = BatchNorm(gat_hidden * gat_heads)\n        self.gat2 = GATv2Conv(gat_hidden * gat_heads, gat_hidden, heads=1, dropout=gat_dropout, concat=True)\n        self.bn2  = BatchNorm(gat_hidden)\n        self.fc1      = nn.Linear(in_dim_tabular, mlp_hidden)\n        self.bn_tab1  = nn.BatchNorm1d(mlp_hidden)\n        self.fc2      = nn.Linear(mlp_hidden, in_dim_tabular)\n        self.bn_tab2  = nn.BatchNorm1d(in_dim_tabular)\n        self.dropout  = nn.Dropout(gat_dropout)\n        fused_dim = gat_hidden + in_dim_tabular\n        self.classifier = nn.Sequential(\n            nn.Linear(fused_dim, fused_dim), nn.ReLU(), nn.Dropout(0.5), nn.Linear(fused_dim, 1))\n\n    def forward(self, x_tab, x_gat, edge_index):\n        z = self.dropout(F.elu(self.bn1(self.gat1(x_gat, edge_index))))\n        z = self.dropout(F.elu(self.bn2(self.gat2(z, edge_index))))\n        t = self.dropout(F.relu(self.bn_tab1(self.fc1(x_tab))))\n        t = self.dropout(F.relu(self.bn_tab2(self.fc2(t))))\n        return self.classifier(torch.cat([z, t], dim=1)).squeeze(1)\n\n# -------------------------\n# TRAIN / EVAL\n# -------------------------\ndef train_one_epoch(model, loader, optimizer, pos_weight, device):\n    model.train()\n    total_loss = 0.0; total = 0\n    for batch in loader:\n        batch = batch.to(device); optimizer.zero_grad()\n        loss = F.binary_cross_entropy_with_logits(\n            model(batch.x_tab, batch.x_gat, batch.edge_index),\n            batch.y.float(), pos_weight=pos_weight)\n        loss.backward(); optimizer.step()\n        total_loss += float(loss.item()) * batch.num_nodes\n        total += batch.num_nodes\n    return total_loss / max(1, total)\n\n@torch.no_grad()\ndef evaluate(model, loader, device):\n    model.eval()\n    ys, ps = [], []\n    for batch in loader:\n        batch = batch.to(device)\n        ps.append(torch.sigmoid(model(batch.x_tab, batch.x_gat, batch.edge_index)).cpu().numpy())\n        ys.append(batch.y.cpu().numpy())\n    y_true = np.concatenate(ys); y_prob = np.concatenate(ps)\n    y_pred = (y_prob >= CFG[\"THRESHOLD\"]).astype(int)\n    return (roc_auc_score(y_true, y_prob),\n            f1_score(y_true, y_pred),\n            accuracy_score(y_true, y_pred),\n            y_true, y_prob)\n\ndef build_pyg_data(df, feature_cols, edge_index):\n    X = torch.tensor(df[feature_cols].values, dtype=torch.float32)\n    y = torch.tensor(df[\"is_churn\"].values, dtype=torch.long)\n    data = Data()\n    data.x_tab = X.clone(); data.x_gat = X.clone()\n    data.y = y; data.edge_index = edge_index; data.num_nodes = X.size(0)\n    return data\n\n# =========================================================\n# GRAPHS — YEH SAB SHOW HOGA!\n# =========================================================\n\ndef plot_confusion_matrix(y_true, y_prob, threshold=0.5, title=\"Test Set\"):\n    \"\"\"Graph 1: Confusion Matrix\"\"\"\n    y_pred = (y_prob >= threshold).astype(int)\n    cm = confusion_matrix(y_true, y_pred)\n    plt.figure(figsize=(6, 5))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n                xticklabels=['Not Churn', 'Churn'],\n                yticklabels=['Not Churn', 'Churn'],\n                annot_kws={\"size\": 14})\n    auc = roc_auc_score(y_true, y_prob)\n    f1  = f1_score(y_true, y_pred)\n    acc = accuracy_score(y_true, y_pred)\n    plt.title(f'Confusion Matrix — {title}\\nAUC={auc:.4f} | F1={f1:.4f} | Acc={acc:.4f}', fontsize=12)\n    plt.xlabel('Predicted Label'); plt.ylabel('True Label')\n    plt.tight_layout()\n    plt.savefig(\"/content/graph1_confusion_matrix.png\", dpi=150)\n    plt.show()\n    print(f\"\\nTN:{cm[0][0]} | FP:{cm[0][1]}\")\n    print(f\"FN:{cm[1][0]} | TP:{cm[1][1]}\")\n    print(f\"\\nAccuracy: {acc:.4f} | AUC: {auc:.4f} | F1: {f1:.4f}\")\n    print(\"\\n\" + classification_report(y_true, y_pred,\n          target_names=['Not Churn','Churn'], digits=4))\n\ndef plot_training_curves(train_losses, val_aucs):\n    \"\"\"Graph 2: Training Loss + Validation AUC\"\"\"\n    epochs = np.arange(1, len(train_losses) + 1)\n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))\n\n    # Loss curve\n    ax1.plot(epochs, train_losses, color='tomato', linewidth=2, label='Train Loss')\n    ax1.set_xlabel('Epoch'); ax1.set_ylabel('Loss')\n    ax1.set_title('Training Loss over Epochs')\n    ax1.legend(); ax1.grid(True, alpha=0.3)\n\n    # AUC curve\n    best_epoch = np.argmax(val_aucs) + 1\n    best_auc   = max(val_aucs)\n    ax2.plot(epochs, val_aucs, color='steelblue', linewidth=2, label='Val AUC')\n    ax2.axvline(x=best_epoch, color='green', linestyle='--', alpha=0.7,\n                label=f'Best Epoch={best_epoch} (AUC={best_auc:.4f})')\n    ax2.set_xlabel('Epoch'); ax2.set_ylabel('AUC')\n    ax2.set_title('Validation AUC over Epochs')\n    ax2.legend(); ax2.grid(True, alpha=0.3)\n\n    plt.tight_layout()\n    plt.savefig(\"/content/graph2_training_curves.png\", dpi=150)\n    plt.show()\n    print(f\"✅ Best Val AUC: {best_auc:.4f} at Epoch {best_epoch}\")\n\ndef plot_roc_curve(y_true, y_prob):\n    \"\"\"Graph 3: ROC Curve\"\"\"\n    fpr, tpr, _ = roc_curve(y_true, y_prob)\n    auc = roc_auc_score(y_true, y_prob)\n    plt.figure(figsize=(6, 5))\n    plt.plot(fpr, tpr, color='darkorange', lw=2, label=f'ROC curve (AUC = {auc:.4f})')\n    plt.plot([0,1], [0,1], color='navy', lw=1, linestyle='--', label='Random classifier')\n    plt.fill_between(fpr, tpr, alpha=0.1, color='darkorange')\n    plt.xlabel('False Positive Rate'); plt.ylabel('True Positive Rate')\n    plt.title('ROC Curve — Hybrid GAT+MLP')\n    plt.legend(loc='lower right'); plt.grid(True, alpha=0.3)\n    plt.tight_layout()\n    plt.savefig(\"/content/graph3_roc_curve.png\", dpi=150)\n    plt.show()\n\ndef plot_precision_recall_curve(y_true, y_prob):\n    \"\"\"Graph 4: Precision-Recall Curve\"\"\"\n    precision, recall, _ = precision_recall_curve(y_true, y_prob)\n    ap = average_precision_score(y_true, y_prob)\n    baseline = y_true.mean()\n    plt.figure(figsize=(6, 5))\n    plt.plot(recall, precision, color='purple', lw=2, label=f'PR curve (AP = {ap:.4f})')\n    plt.axhline(y=baseline, color='gray', linestyle='--', label=f'Baseline = {baseline:.2f}')\n    plt.fill_between(recall, precision, alpha=0.1, color='purple')\n    plt.xlabel('Recall'); plt.ylabel('Precision')\n    plt.title('Precision-Recall Curve')\n    plt.legend(); plt.grid(True, alpha=0.3)\n    plt.tight_layout()\n    plt.savefig(\"/content/graph4_precision_recall.png\", dpi=150)\n    plt.show()\n\ndef plot_probability_distribution(y_true, y_prob):\n    \"\"\"Graph 5: Churn probability distribution\"\"\"\n    plt.figure(figsize=(8, 4))\n    plt.hist(y_prob[y_true == 0], bins=50, alpha=0.6, color='steelblue',\n             label='Not Churn', density=True)\n    plt.hist(y_prob[y_true == 1], bins=50, alpha=0.6, color='tomato',\n             label='Churn', density=True)\n    plt.axvline(x=CFG[\"THRESHOLD\"], color='black', linestyle='--',\n                linewidth=2, label=f'Threshold = {CFG[\"THRESHOLD\"]}')\n    plt.xlabel('Predicted Probability'); plt.ylabel('Density')\n    plt.title('Churn Probability Distribution')\n    plt.legend(); plt.grid(True, alpha=0.3)\n    plt.tight_layout()\n    plt.savefig(\"/content/graph5_prob_distribution.png\", dpi=150)\n    plt.show()\n\n# =========================================================\n# MAIN\n# =========================================================\ndef main():\n    df, feature_cols, scaler = build_user_level_table(CFG[\"PATHS\"])\n    print(f\"Users after merge: {len(df):,}\")\n\n    edge_index, _ = build_sparse_similarity_edges(\n        df[[\"city\",\"registered_via\"]].copy(),\n        max_neighbors=CFG[\"GRAPH_MAX_NEIGHBORS_PER_GROUP\"])\n\n    data = build_pyg_data(df, feature_cols, edge_index)\n\n    idx = np.arange(data.num_nodes)\n    train_idx, test_idx = train_test_split(idx, test_size=CFG[\"TEST_RATIO\"],\n                                           stratify=df[\"is_churn\"],\n                                           random_state=CFG[\"RANDOM_SEED\"])\n    train_idx, val_idx  = train_test_split(train_idx,\n                                           test_size=CFG[\"VAL_RATIO\"]/(1-CFG[\"TEST_RATIO\"]),\n                                           stratify=df[\"is_churn\"].iloc[train_idx],\n                                           random_state=CFG[\"RANDOM_SEED\"])\n\n    data.train_mask = torch.zeros(data.num_nodes, dtype=torch.bool)\n    data.val_mask   = torch.zeros(data.num_nodes, dtype=torch.bool)\n    data.test_mask  = torch.zeros(data.num_nodes, dtype=torch.bool)\n    data.train_mask[torch.tensor(train_idx)] = True\n    data.val_mask[torch.tensor(val_idx)]     = True\n    data.test_mask[torch.tensor(test_idx)]   = True\n\n    train_loader = NeighborLoader(data, num_neighbors=[15,10],\n                                  batch_size=CFG[\"BATCH_SIZE\"], input_nodes=data.train_mask)\n    val_loader   = NeighborLoader(data, num_neighbors=[15,10],\n                                  batch_size=CFG[\"BATCH_SIZE\"], input_nodes=data.val_mask)\n    test_loader  = NeighborLoader(data, num_neighbors=[15,10],\n                                  batch_size=CFG[\"BATCH_SIZE\"], input_nodes=data.test_mask)\n\n    device = CFG[\"DEVICE\"]; print(\"Device:\", device)\n\n    # Model banane se pehle seed lagao\n    torch.manual_seed(CFG[\"RANDOM_SEED\"])\n    torch.cuda.manual_seed_all(CFG[\"RANDOM_SEED\"])\n\n    model = HybridGATMLP(\n        in_dim_tabular=len(feature_cols), gat_in_dim=len(feature_cols),\n        gat_hidden=CFG[\"GAT_HIDDEN\"], gat_heads=CFG[\"GAT_HEADS\"],\n        gat_dropout=CFG[\"GAT_DROPOUT\"], mlp_hidden=CFG[\"MLP_HIDDEN\"],\n    ).to(device)\n\n    optimizer  = torch.optim.AdamW(model.parameters(),\n                                   lr=CFG[\"LR\"], weight_decay=CFG[\"WEIGHT_DECAY\"])\n    scheduler  = torch.optim.lr_scheduler.ReduceLROnPlateau(\n                    optimizer, mode='max', factor=0.5, patience=3, verbose=True)\n    pos_weight = torch.tensor([10.0], device=device)\n\n    # ===== TRAINING LOOP =====\n    best_auc   = -1.0\n    best_state = None\n    bad        = 0\n    train_losses = []   # loss store karo\n    val_aucs     = []   # AUC store karo\n\n    for epoch in range(1, CFG[\"MAX_EPOCHS\"] + 1):\n        tr_loss = train_one_epoch(model, train_loader, optimizer, pos_weight, device)\n        val_auc, val_f1, val_acc, _, _ = evaluate(model, val_loader, device)\n\n        # Lists mein save karo (graph ke liye)\n        train_losses.append(tr_loss)\n        val_aucs.append(val_auc)\n\n        scheduler.step(val_auc)\n        print(f\"Epoch {epoch:02d} | loss {tr_loss:.4f} | val AUC {val_auc:.4f} \"\n              f\"| F1 {val_f1:.4f} | Acc {val_acc:.4f}\")\n\n        if val_auc > best_auc + 1e-4:\n            best_auc   = val_auc\n            best_state = {k: v.cpu() for k, v in model.state_dict().items()}\n            bad = 0\n        else:\n            bad += 1\n            if bad >= CFG[\"EARLY_STOP_PATIENCE\"]:\n                print(\"Early stopping.\"); break\n\n    if best_state:\n        model.load_state_dict({k: v.to(device) for k, v in best_state.items()})\n\n    # ===== FINAL TEST =====\n    test_auc, test_f1, test_acc, y_true, y_prob = evaluate(model, test_loader, device)\n    print(f\"\\nTEST | AUC {test_auc:.4f} | F1 {test_f1:.4f} | Acc {test_acc:.4f}\")\n\n    # =========================================================\n    # SARE GRAPHS YAHAN CALL HO RAHE HAIN\n    # =========================================================\n    print(\"\\n\" + \"=\"*50)\n    print(\"GRAPHS BAN RAHE HAIN...\")\n    print(\"=\"*50)\n\n    print(\"\\n📊 Graph 1: Confusion Matrix\")\n    plot_confusion_matrix(y_true, y_prob,\n                          threshold=CFG[\"THRESHOLD\"], title=\"KKBox Test Set\")\n\n    print(\"\\n📈 Graph 2: Training Curves\")\n    plot_training_curves(train_losses, val_aucs)\n\n    print(\"\\n📉 Graph 3: ROC Curve\")\n    plot_roc_curve(y_true, y_prob)\n\n    print(\"\\n🎯 Graph 4: Precision-Recall Curve\")\n    plot_precision_recall_curve(y_true, y_prob)\n\n    print(\"\\n🔔 Graph 5: Probability Distribution\")\n    plot_probability_distribution(y_true, y_prob)\n\n    print(\"\\n✅ Sare graphs ban gaye aur /content/ mein save ho gaye!\")\n    print(\"Files: graph1_confusion_matrix.png\")\n    print(\"       graph2_training_curves.png\")\n    print(\"       graph3_roc_curve.png\")\n    print(\"       graph4_precision_recall.png\")\n    print(\"       graph5_prob_distribution.png\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-04T17:31:41.387754Z","iopub.execute_input":"2026-05-04T17:31:41.388241Z","iopub.status.idle":"2026-05-04T17:54:24.358702Z","shell.execute_reply.started":"2026-05-04T17:31:41.388205Z","shell.execute_reply":"2026-05-04T17:54:24.358011Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"ok\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-03T03:00:34.384904Z","iopub.execute_input":"2026-07-03T03:00:34.385194Z","iopub.status.idle":"2026-07-03T03:00:34.396937Z","shell.execute_reply.started":"2026-07-03T03:00:34.385166Z","shell.execute_reply":"2026-07-03T03:00:34.395641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}