{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Cell 1 — Environment, imports and reproducibility\nimport os, sys, random, warnings, gc, copy, subprocess, math\nwarnings.filterwarnings(\"ignore\")\n\ndef ensure_import(module, pip_name):\n    try:\n        __import__(module)\n    except Exception:\n        subprocess.check_call([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", pip_name])\n\nensure_import(\"torch_geometric\", \"torch-geometric\")\nensure_import(\"pynndescent\", \"pynndescent\")\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom pathlib import Path\nfrom scipy.sparse import csr_matrix\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import StandardScaler\nfrom sklearn.metrics import (\n    accuracy_score, balanced_accuracy_score,\n    precision_recall_fscore_support, roc_auc_score,\n    average_precision_score, confusion_matrix\n)\nfrom sklearn.pipeline import make_pipeline\nfrom sklearn.linear_model import LogisticRegression\nfrom sklearn.ensemble import RandomForestClassifier\nfrom sklearn.neural_network import MLPClassifier\n\nfrom torch_geometric.data import HeteroData\nfrom torch_geometric.nn import HeteroConv, GATConv, SAGEConv\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(\"Device:\", device)\nprint(\"GPU:\", torch.cuda.get_device_name(0) if torch.cuda.is_available() else \"CPU\")\nprint(\"PyTorch:\", torch.__version__)\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-10-02T12:53:42.260738Z","iopub.execute_input":"2026-10-02T12:53:42.261079Z","iopub.status.idle":"2026-10-02T12:54:33.117425Z","shell.execute_reply.started":"2026-10-02T12:53:42.261046Z","shell.execute_reply":"2026-10-02T12:54:33.116494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 2 — H&M paths and ID normalization\n\ndef normalize_article_id(s):\n    return (\n        s.astype(\"string\")\n         .str.strip()\n         .str.replace(r\"\\.0$\", \"\", regex=True)\n         .str.zfill(10)\n    )\n\ndef normalize_customer_id(s):\n    return s.astype(\"string\").str.strip()\n\nDATA_ROOT = Path(\"/kaggle/input/h-and-m-personalized-fashion-recommendations\")\n\nif not DATA_ROOT.exists():\n    DATA_ROOT = None\n    for p in Path(\"/kaggle/input\").glob(\"**/transactions_train.csv\"):\n        DATA_ROOT = p.parent\n        break\n\nif DATA_ROOT is None:\n    raise FileNotFoundError(\n        \"Attach the official H&M Personalized Fashion Recommendations \"\n        \"competition dataset with Kaggle Add Input.\"\n    )\n\nARTICLES = DATA_ROOT / \"articles.csv\"\nCUSTOMERS = DATA_ROOT / \"customers.csv\"\nTRANSACTIONS = DATA_ROOT / \"transactions_train.csv\"\n\nfor p in [ARTICLES, CUSTOMERS, TRANSACTIONS]:\n    if not p.exists():\n        raise FileNotFoundError(f\"Missing official H&M file: {p}\")\n\nprint(\"Using:\", DATA_ROOT)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T12:57:59.128478Z","iopub.execute_input":"2026-10-02T12:57:59.129378Z","iopub.status.idle":"2026-10-02T12:57:59.143123Z","shell.execute_reply.started":"2026-10-02T12:57:59.12934Z","shell.execute_reply":"2026-10-02T12:57:59.142291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 3 — Load immutable metadata and determine temporal cutoff\n\narticles_raw = pd.read_csv(ARTICLES)\ncustomers_raw = pd.read_csv(CUSTOMERS)\n\narticles_raw[\"article_id\"] = normalize_article_id(articles_raw[\"article_id\"])\ncustomers_raw[\"customer_id\"] = normalize_customer_id(customers_raw[\"customer_id\"])\n\nmax_date = None\nfor ch in pd.read_csv(\n    TRANSACTIONS,\n    usecols=[\"t_dat\"],\n    dtype={\"t_dat\":\"string\"},\n    chunksize=1_000_000\n):\n    d = pd.to_datetime(ch[\"t_dat\"], errors=\"coerce\").max()\n    if max_date is None or d > max_date:\n        max_date = d\n\nif pd.isna(max_date):\n    raise RuntimeError(\"Could not determine maximum transaction date.\")\n\ncutoff = max_date - pd.Timedelta(days=7)\n\nprint(\"Customers:\", customers_raw.shape)\nprint(\"Articles:\", articles_raw.shape)\nprint(\"Max transaction date:\", max_date)\nprint(\"Temporal hold-out starts:\", cutoff)\nprint(\"Article ID sample:\", articles_raw[\"article_id\"].head(5).tolist())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T12:58:03.232322Z","iopub.execute_input":"2026-10-02T12:58:03.232871Z","iopub.status.idle":"2026-10-02T12:59:05.579594Z","shell.execute_reply.started":"2026-10-02T12:58:03.232838Z","shell.execute_reply":"2026-10-02T12:59:05.578864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 4 — Publication-clean customer cohort and temporal targets\n#\n# IMPORTANT: cohort selection is based ONLY on information available at/before\n# the temporal cutoff. Future transactions are used only AFTER the cohort is\n# frozen to construct the retention target.\n\nMAX_CUSTOMERS = 100_000\nMIN_HIST_INTERACTIONS = 1\n\nhist_counts = {}\n\nfor ch in pd.read_csv(\n    TRANSACTIONS,\n    usecols=[\"t_dat\", \"customer_id\"],\n    dtype={\"t_dat\": \"string\", \"customer_id\": \"string\"},\n    chunksize=1_000_000\n):\n    ch[\"t_dat\"] = pd.to_datetime(ch[\"t_dat\"], errors=\"coerce\")\n    ch[\"customer_id\"] = normalize_customer_id(ch[\"customer_id\"])\n\n    hist = ch.loc[\n        (ch[\"t_dat\"] <= cutoff) & ch[\"customer_id\"].notna(),\n        \"customer_id\"\n    ]\n\n    for cid, n in hist.value_counts().items():\n        hist_counts[cid] = hist_counts.get(cid, 0) + int(n)\n\ncustomers_df = customers_raw.copy()\ncustomers_df[\"hist_count\"] = (\n    customers_df[\"customer_id\"]\n    .map(hist_counts)\n    .fillna(0)\n    .astype(int)\n)\n\n# Freeze the cohort using historical activity only.\ncustomers_df = customers_df[\n    customers_df[\"hist_count\"] >= MIN_HIST_INTERACTIONS\n].copy()\n\nif len(customers_df) > MAX_CUSTOMERS:\n    customers_df = (\n        customers_df\n        .sort_values(\n            [\"hist_count\", \"customer_id\"],\n            ascending=[False, True]\n        )\n        .head(MAX_CUSTOMERS)\n        .copy()\n    )\n\ncustomers_df = customers_df.reset_index(drop=True)\n\n# ONLY NOW derive the future retention target for the frozen cohort.\nfuture_customer_ids = set()\n\nfor ch in pd.read_csv(\n    TRANSACTIONS,\n    usecols=[\"t_dat\", \"customer_id\"],\n    dtype={\"t_dat\": \"string\", \"customer_id\": \"string\"},\n    chunksize=1_000_000\n):\n    ch[\"t_dat\"] = pd.to_datetime(ch[\"t_dat\"], errors=\"coerce\")\n    ch[\"customer_id\"] = normalize_customer_id(ch[\"customer_id\"])\n\n    future_customer_ids.update(\n        ch.loc[\n            (ch[\"t_dat\"] > cutoff)\n            & ch[\"customer_id\"].isin(customers_df[\"customer_id\"]),\n            \"customer_id\"\n        ].dropna().tolist()\n    )\n\ncustomers_df[\"retention_y\"] = customers_df[\"customer_id\"].isin(\n    future_customer_ids\n).astype(np.int64)\n\n# Historical-only VIP proxy.\nvip_threshold = customers_df[\"hist_count\"].quantile(0.80)\ncustomers_df[\"vip_y\"] = (\n    customers_df[\"hist_count\"] >= vip_threshold\n).astype(np.int64)\n\ncohort_ids = set(customers_df[\"customer_id\"])\n\nprint(\"Publication-clean cohort:\", len(customers_df))\nprint(\"Retention positive rate:\", round(float(customers_df[\"retention_y\"].mean()), 6))\nprint(\"VIP positive rate:\", round(float(customers_df[\"vip_y\"].mean()), 6))\nprint(\"Historical interaction median:\", customers_df[\"hist_count\"].median())\nprint(\"Cohort selected using historical data only: YES\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T12:59:05.580903Z","iopub.execute_input":"2026-10-02T12:59:05.581172Z","iopub.status.idle":"2026-10-02T13:00:54.825339Z","shell.execute_reply.started":"2026-10-02T12:59:05.581145Z","shell.execute_reply":"2026-10-02T13:00:54.824703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 5 — Publication-clean temporal interactions and candidate catalog\n#\n# Candidate products are determined from PRE-CUTOFF popularity only.\n# Future products are never used to construct the graph/catalog.\n# Future transactions are used only as held-out evaluation labels.\n\nitem_counts = {}\nfuture_cohort_rows = 0\nfuture_kept_rows = 0\n\nfor ch in pd.read_csv(\n    TRANSACTIONS,\n    usecols=[\"t_dat\", \"customer_id\", \"article_id\"],\n    dtype={\n        \"t_dat\": \"string\",\n        \"customer_id\": \"string\",\n        \"article_id\": \"string\"\n    },\n    chunksize=1_000_000\n):\n    ch[\"t_dat\"] = pd.to_datetime(ch[\"t_dat\"], errors=\"coerce\")\n    ch[\"customer_id\"] = normalize_customer_id(ch[\"customer_id\"])\n    ch[\"article_id\"] = normalize_article_id(ch[\"article_id\"])\n\n    cohort_mask = ch[\"customer_id\"].isin(cohort_ids)\n\n    h = ch[\n        (ch[\"t_dat\"] <= cutoff) & cohort_mask\n    ]\n    f = ch[\n        (ch[\"t_dat\"] > cutoff) & cohort_mask\n    ]\n\n    if not h.empty:\n        for aid, n in h[\"article_id\"].value_counts().items():\n            item_counts[aid] = item_counts.get(aid, 0) + int(n)\n\n    future_cohort_rows += int(len(f))\n\n# Candidate catalog: historical-only top products.\nMAX_HIST_ITEMS = 10_000\n\ntop_hist_items = set(\n    pd.Series(item_counts)\n    .sort_values(ascending=False)\n    .head(MAX_HIST_ITEMS)\n    .index\n    .astype(str)\n)\n\nkeep_items = top_hist_items\n\narticles_graph = (\n    articles_raw[\n        articles_raw[\"article_id\"].isin(keep_items)\n    ]\n    .drop_duplicates(\"article_id\")\n    .reset_index(drop=True)\n    .copy()\n)\n\nif articles_graph.empty:\n    raise RuntimeError(\"The historical candidate product universe is empty.\")\n\nhist_parts = []\nfuture_parts = []\n\nfor ch in pd.read_csv(\n    TRANSACTIONS,\n    usecols=[\"t_dat\", \"customer_id\", \"article_id\", \"price\"],\n    dtype={\n        \"t_dat\": \"string\",\n        \"customer_id\": \"string\",\n        \"article_id\": \"string\",\n        \"price\": \"float32\"\n    },\n    chunksize=1_000_000\n):\n    ch[\"t_dat\"] = pd.to_datetime(ch[\"t_dat\"], errors=\"coerce\")\n    ch[\"customer_id\"] = normalize_customer_id(ch[\"customer_id\"])\n    ch[\"article_id\"] = normalize_article_id(ch[\"article_id\"])\n\n    h = ch[\n        (ch[\"t_dat\"] <= cutoff)\n        & ch[\"customer_id\"].isin(cohort_ids)\n        & ch[\"article_id\"].isin(keep_items)\n    ]\n\n    f = ch[\n        (ch[\"t_dat\"] > cutoff)\n        & ch[\"customer_id\"].isin(cohort_ids)\n        & ch[\"article_id\"].isin(keep_items)\n    ]\n\n    if not h.empty:\n        hist_parts.append(h)\n\n    if not f.empty:\n        future_parts.append(f)\n        future_kept_rows += int(len(f))\n\nhist_df = (\n    pd.concat(hist_parts, ignore_index=True)\n    if hist_parts\n    else pd.DataFrame(columns=[\"t_dat\", \"customer_id\", \"article_id\", \"price\"])\n)\n\nfuture_df = (\n    pd.concat(future_parts, ignore_index=True)\n    if future_parts\n    else pd.DataFrame(columns=[\"t_dat\", \"customer_id\", \"article_id\", \"price\"])\n)\n\nMAX_HIST_PER_CUSTOMER = 30\n\nhist_df = hist_df.sort_values([\"customer_id\", \"t_dat\"])\nhist_df = hist_df.drop_duplicates(\n    [\"customer_id\", \"article_id\"],\n    keep=\"last\"\n)\nhist_df = (\n    hist_df\n    .groupby(\"customer_id\", group_keys=False)\n    .tail(MAX_HIST_PER_CUSTOMER)\n    .reset_index(drop=True)\n)\n\nproduct_id_to_idx = {\n    aid: i\n    for i, aid in enumerate(articles_graph[\"article_id\"].tolist())\n}\n\nfuture_catalog_coverage = (\n    future_kept_rows / future_cohort_rows\n    if future_cohort_rows > 0 else 0.0\n)\n\nprint(\"Historical graph interactions:\", len(hist_df))\nprint(\"Max historical unique products retained per customer:\", MAX_HIST_PER_CUSTOMER)\nprint(\"Future held-out interactions inside candidate catalog:\", len(future_df))\nprint(\"Historical candidate products:\", len(top_hist_items))\nprint(\"Graph product universe:\", len(articles_graph))\nprint(\"Future interaction catalog coverage:\", round(float(future_catalog_coverage), 6))\nprint(\"Future products used to build catalog: NO\")\nprint(\"Candidate catalog is pre-cutoff historical-only: YES\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:00:54.826406Z","iopub.execute_input":"2026-10-02T13:00:54.826742Z","iopub.status.idle":"2026-10-02T13:03:46.024163Z","shell.execute_reply.started":"2026-10-02T13:00:54.826705Z","shell.execute_reply":"2026-10-02T13:03:46.023291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 6 — Leakage-safe customer split and train-only feature statistics\n\nindices = np.arange(len(customers_df))\ny_ret = customers_df[\"retention_y\"].to_numpy()\n\ntrain_idx, test_idx = train_test_split(\n    indices,\n    test_size=0.15,\n    random_state=SEED,\n    stratify=y_ret\n)\n\ntrain_idx, val_idx = train_test_split(\n    train_idx,\n    test_size=0.1764705882,   # gives 70/15/15 overall\n    random_state=SEED,\n    stratify=y_ret[train_idx]\n)\n\ntrain_mask_np = np.zeros(len(customers_df), dtype=bool)\nval_mask_np = np.zeros(len(customers_df), dtype=bool)\ntest_mask_np = np.zeros(len(customers_df), dtype=bool)\n\ntrain_mask_np[train_idx] = True\nval_mask_np[val_idx] = True\ntest_mask_np[test_idx] = True\n\n# Leakage-safe historical behavior features.\n# Every feature below is computed from transactions at or before the temporal cutoff.\n# Scaling statistics are fit on TRAIN customers only.\n\nhist_behavior = (\n    hist_df.groupby(\"customer_id\")\n    .agg(\n        hist_unique_products=(\"article_id\", \"nunique\"),\n        hist_avg_price=(\"price\", \"mean\"),\n        hist_first_date=(\"t_dat\", \"min\"),\n        hist_last_date=(\"t_dat\", \"max\"),\n    )\n    .reset_index()\n)\n\nhist_behavior[\"hist_recency_days\"] = (\n    pd.Timestamp(cutoff) - pd.to_datetime(hist_behavior[\"hist_last_date\"])\n).dt.days.clip(lower=0)\nhist_behavior[\"hist_active_span_days\"] = (\n    pd.to_datetime(hist_behavior[\"hist_last_date\"])\n    - pd.to_datetime(hist_behavior[\"hist_first_date\"])\n).dt.days.clip(lower=0)\n\nbehavior_cols = [\n    \"hist_unique_products\",\n    \"hist_avg_price\",\n    \"hist_recency_days\",\n    \"hist_active_span_days\",\n]\n\ncustomers_df = customers_df.merge(\n    hist_behavior[[\"customer_id\"] + behavior_cols],\n    on=\"customer_id\", how=\"left\"\n)\n\nfor c in behavior_cols:\n    train_median = customers_df.loc[train_idx, c].median()\n    customers_df[c] = pd.to_numeric(customers_df[c], errors=\"coerce\").fillna(train_median if pd.notna(train_median) else 0.0)\n\n# Fit preprocessing ONLY on train customers.\nnum_cols = [\"age\", \"hist_count\"] + behavior_cols\ncat_cols = [\"club_member_status\", \"fashion_news_frequency\"]\n\nfor c in num_cols:\n    customers_df[c] = pd.to_numeric(\n        customers_df[c], errors=\"coerce\"\n    )\n    train_median = customers_df.loc[train_idx, c].median()\n    customers_df[c] = customers_df[c].fillna(train_median)\n\nfor c in cat_cols:\n    customers_df[c] = (\n        customers_df[c].fillna(\"Unknown\").astype(str)\n    )\n\nscaler = StandardScaler().fit(\n    customers_df.loc[train_idx, num_cols]\n)\ncustomers_df.loc[:, num_cols] = scaler.transform(\n    customers_df[num_cols]\n)\n\nprint(\n    \"Split sizes:\",\n    len(train_idx), len(val_idx), len(test_idx)\n)\nprint(\n    \"Train/Val/Test percentages:\",\n    round(len(train_idx)/len(indices)*100,2),\n    round(len(val_idx)/len(indices)*100,2),\n    round(len(test_idx)/len(indices)*100,2)\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:03:46.026437Z","iopub.execute_input":"2026-10-02T13:03:46.026819Z","iopub.status.idle":"2026-10-02T13:03:48.371711Z","shell.execute_reply.started":"2026-10-02T13:03:46.02679Z","shell.execute_reply":"2026-10-02T13:03:48.370919Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 7 — Build the real heterogeneous graph\n# Important: customer-peer edges are constructed from TRAIN customers only.\n\nfrom pynndescent import NNDescent\n\ncustomer_ids = customers_df[\"customer_id\"].tolist()\narticle_ids = articles_graph[\"article_id\"].tolist()\n\ncustomer_id_to_idx = {\n    v: i for i, v in enumerate(customer_ids)\n}\nproduct_id_to_idx = {\n    v: i for i, v in enumerate(article_ids)\n}\n\nhist_df = hist_df[\n    hist_df[\"customer_id\"].isin(customer_id_to_idx) &\n    hist_df[\"article_id\"].isin(product_id_to_idx)\n].copy()\n\nfuture_df = future_df[\n    future_df[\"customer_id\"].isin(customer_id_to_idx) &\n    future_df[\"article_id\"].isin(product_id_to_idx)\n].copy()\n\nif hist_df.empty:\n    raise RuntimeError(\"No historical customer-product matches remain.\")\n\nrows = (\n    hist_df[\"customer_id\"]\n    .map(customer_id_to_idx)\n    .to_numpy(np.int64)\n)\ncols = (\n    hist_df[\"article_id\"]\n    .map(product_id_to_idx)\n    .to_numpy(np.int64)\n)\n\nX_sparse = csr_matrix(\n    (\n        np.ones(len(rows), dtype=np.float32),\n        (rows, cols)\n    ),\n    shape=(len(customer_ids), len(article_ids)),\n    dtype=np.float32\n)\n\n# Peer graph is strictly restricted to training customer indices.\npeer_train = np.asarray(train_idx, dtype=np.int64)\n\nif len(peer_train) > 50_000:\n    rng = np.random.default_rng(SEED)\n    peer_train = np.sort(\n        rng.choice(\n            peer_train,\n            size=50_000,\n            replace=False\n        )\n    )\n\nif len(peer_train) < 8:\n    raise RuntimeError(\"Too few training customers for peer graph.\")\n\nnn_index = NNDescent(\n    X_sparse[peer_train],\n    n_neighbors=6,\n    metric=\"cosine\",\n    random_state=SEED,\n    low_memory=True,\n    n_trees=5,\n    n_iters=5,\n    verbose=False\n)\n\nneigh, _ = nn_index.neighbor_graph\nsrc_peer = peer_train[\n    np.repeat(np.arange(len(peer_train)), neigh.shape[1])\n]\ndst_peer = peer_train[\n    neigh.reshape(-1)\n]\n\ndef one_hot_series(s):\n    return pd.get_dummies(\n        s.fillna(\"Unknown\").astype(str),\n        dtype=np.float32\n    )\n\n# Customer attributes.\ncustomer_parts = []\nfor c in num_cols:\n    customer_parts.append(\n        customers_df[c].to_numpy(np.float32)[:, None]\n    )\n\nfor c in [\"FN\",\"Active\",\"club_member_status\",\"fashion_news_frequency\"]:\n    if c in customers_df.columns:\n        customer_parts.append(\n            one_hot_series(customers_df[c]).to_numpy(np.float32)\n        )\n\ncustomer_x = np.hstack(customer_parts).astype(np.float32)\n\n# Product attributes from the official H&M article table.\nproduct_parts = []\nproduct_feature_cols = [\n    \"product_type_no\",\n    \"graphical_appearance_no\",\n    \"colour_group_code\",\n    \"perceived_colour_value_id\",\n    \"department_no\",\n    \"index_code\",\n    \"index_group_no\",\n    \"section_no\",\n    \"garment_group_no\",\n    \"index_name\",\n    \"index_group_name\",\n    \"department_name\",\n    \"section_name\",\n    \"garment_group_name\",\n    \"product_type_name\"\n]\n\nfor c in product_feature_cols:\n    if c in articles_graph.columns:\n        product_parts.append(\n            one_hot_series(\n                articles_graph[c]\n            ).to_numpy(np.float32)\n        )\n\nproduct_x = np.hstack(product_parts).astype(np.float32)\n\ngraph_data = HeteroData()\n# Canonical graph variable used by all later cells.\ndata = graph_data\n\ngraph_data[\"customer\"].x = torch.tensor(\n    customer_x, dtype=torch.float32\n)\ngraph_data[\"product\"].x = torch.tensor(\n    product_x, dtype=torch.float32\n)\n\npurchase_ei = torch.tensor(\n    np.vstack([rows, cols]),\n    dtype=torch.long\n)\n\ngraph_data[\n    \"customer\", \"purchases\", \"product\"\n].edge_index = purchase_ei\n\ngraph_data[\n    \"product\", \"rev_purchases\", \"customer\"\n].edge_index = purchase_ei.flip(0)\n\npeer_ei = torch.tensor(\n    np.vstack([src_peer, dst_peer]),\n    dtype=torch.long\n)\n\ngraph_data[\n    \"customer\", \"peer_with\", \"customer\"\n].edge_index = peer_ei\n\n# Real article attribute nodes.\nattribute_specs = [\n    (\"product_type_name\", \"product_type_attr\"),\n    (\"department_name\", \"department_attr\"),\n    (\"section_name\", \"section_attr\"),\n    (\"garment_group_name\", \"garment_group_attr\"),\n    (\"index_group_name\", \"index_group_attr\"),\n    (\"colour_group_code\", \"colour_attr\")\n]\n\nfor col, node_type in attribute_specs:\n\n    if col not in articles_graph.columns:\n        continue\n\n    vals = (\n        articles_graph[col]\n        .fillna(\"Unknown\")\n        .astype(str)\n    )\n\n    uniques = pd.Index(vals.unique())\n    mapping = {v:i for i,v in enumerate(uniques)}\n    aidx = vals.map(mapping).to_numpy(np.int64)\n\n    graph_data[node_type].x = torch.ones(\n        (len(uniques), 1),\n        dtype=torch.float32\n    )\n\n    pidx = np.arange(\n        len(article_ids),\n        dtype=np.int64\n    )\n\n    ei = torch.tensor(\n        np.vstack([pidx, aidx]),\n        dtype=torch.long\n    )\n\n    graph_data[\n        \"product\",\n        f\"has_{node_type}\",\n        node_type\n    ].edge_index = ei\n\n    graph_data[\n        node_type,\n        f\"rev_has_{node_type}\",\n        \"product\"\n    ].edge_index = ei.flip(0)\n\n# Customer targets and masks.\ngraph_data[\"customer\"].retention_y = torch.tensor(\n    customers_df[\"retention_y\"].to_numpy(np.int64)\n)\ngraph_data[\"customer\"].vip_y = torch.tensor(\n    customers_df[\"vip_y\"].to_numpy(np.int64)\n)\ngraph_data[\"customer\"].train_mask = torch.tensor(train_mask_np)\ngraph_data[\"customer\"].val_mask = torch.tensor(val_mask_np)\ngraph_data[\"customer\"].test_mask = torch.tensor(test_mask_np)\n\n# Real future-item dictionary for recommendation evaluation.\ntest_items_by_customer = {}\ntrain_items_by_customer = {}\n\nfor cid, grp in future_df.groupby(\"customer_id\"):\n    u = customer_id_to_idx.get(cid)\n    if u is not None:\n        items = {\n            product_id_to_idx[a]\n            for a in grp[\"article_id\"]\n            if a in product_id_to_idx\n        }\n        if items:\n            test_items_by_customer[u] = items\n\nfor cid, grp in hist_df.groupby(\"customer_id\"):\n    u = customer_id_to_idx.get(cid)\n    if u is not None:\n        items = {\n            product_id_to_idx[a]\n            for a in grp[\"article_id\"]\n            if a in product_id_to_idx\n        }\n        if items:\n            train_items_by_customer[u] = items\n\n\n# CPU integrity checks BEFORE any CUDA transfer.\ndef _check_edge_bounds(g):\n    for et in g.edge_types:\n        ei = g[et].edge_index\n        if ei is None or ei.numel() == 0:\n            continue\n        src, _, dst = et\n        if int(ei[0].min()) < 0 or int(ei[0].max()) >= int(g[src].num_nodes):\n            raise RuntimeError(f\"Invalid source indices in edge type {et}.\")\n        if int(ei[1].min()) < 0 or int(ei[1].max()) >= int(g[dst].num_nodes):\n            raise RuntimeError(f\"Invalid destination indices in edge type {et}.\")\n\n_check_edge_bounds(data)\nassert data[\"customer\"].retention_y.dtype == torch.long\nassert data[\"customer\"].vip_y.dtype == torch.long\nassert set(torch.unique(data[\"customer\"].retention_y).tolist()).issubset({0,1})\nassert set(torch.unique(data[\"customer\"].vip_y).tolist()).issubset({0,1})\n\nprint(graph_data)\nprint(\"Node counts:\", {\n    k: int(graph_data[k].num_nodes)\n    for k in graph_data.node_types\n})\nprint(\"Edge counts:\", {\n    str(k): int(graph_data[k].edge_index.shape[1])\n    for k in graph_data.edge_types\n})\nprint(\n    \"Real future evaluation users:\",\n    sum(1 for u in test_idx if int(u) in test_items_by_customer)\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:03:48.372817Z","iopub.execute_input":"2026-10-02T13:03:48.37305Z","iopub.status.idle":"2026-10-02T13:04:26.28066Z","shell.execute_reply.started":"2026-10-02T13:03:48.373017Z","shell.execute_reply":"2026-10-02T13:04:26.279866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 8 — Directed outfit-compatibility DAG from official article metadata\n\n# The H&M dataset does not provide outfit-compatibility ground truth.\n# Therefore this is a deterministic metadata prior used only as an auxiliary\n# structural constraint, never as a fabricated evaluation label.\n#\n# The DAG is constructed efficiently by sampling a few products from a\n# strictly later outfit layer. This avoids any O(N^2) candidate construction.\n\ndef outfit_layer(row):\n    text = \" \".join(\n        str(row.get(c, \"\")).lower()\n        for c in [\n            \"product_type_name\",\n            \"department_name\",\n            \"section_name\",\n            \"garment_group_name\"\n        ]\n    )\n\n    if any(k in text for k in [\n        \"shoe\", \"sandal\", \"boot\", \"sneaker\", \"heel\"\n    ]):\n        return 4\n    if any(k in text for k in [\n        \"bag\", \"belt\", \"hat\", \"cap\", \"scarf\",\n        \"glove\", \"sock\", \"accessor\"\n    ]):\n        return 5\n    if any(k in text for k in [\n        \"coat\", \"jacket\", \"blazer\", \"cardigan\",\n        \"parka\", \"outerwear\", \"vest\", \"trench\"\n    ]):\n        return 3\n    if any(k in text for k in [\n        \"trouser\", \"trousers\", \"jeans\", \"skirt\",\n        \"shorts\", \"legging\", \"bottom\"\n    ]):\n        return 2\n    if any(k in text for k in [\n        \"dress\", \"top\", \"shirt\", \"t-shirt\",\n        \"tee\", \"blouse\", \"sweater\", \"hoodie\",\n        \"pullover\", \"bra\", \"tank\"\n    ]):\n        return 1\n    return 0\n\nlayers = articles_graph.apply(\n    outfit_layer,\n    axis=1\n).to_numpy(np.int64)\n\nrng = np.random.default_rng(SEED)\n\ngroup_cols = [\n    c for c in [\n        \"index_group_name\",\n        \"index_name\",\n        \"department_name\"\n    ]\n    if c in articles_graph.columns\n]\n\nif group_cols:\n    group_keys = (\n        articles_graph[group_cols]\n        .fillna(\"Unknown\")\n        .astype(str)\n        .agg(\"|\".join, axis=1)\n        .to_numpy()\n    )\nelse:\n    group_keys = np.array(\n        [\"all\"] * len(articles_graph),\n        dtype=object\n    )\n\n# Pre-index products by layer and broad metadata group.\nlayer_group_to_indices = {}\nlayer_to_indices = {}\n\nfor idx, (layer, group) in enumerate(\n    zip(layers, group_keys)\n):\n    layer = int(layer)\n    layer_to_indices.setdefault(layer, []).append(idx)\n    layer_group_to_indices.setdefault(\n        (layer, str(group)), []\n    ).append(idx)\n\ncompat_src = []\ncompat_dst = []\n\n# For each product, connect to at most 3 products in the next available\n# higher outfit layer, preferring the same broad metadata group.\nfor i in range(len(articles_graph)):\n    li = int(layers[i])\n\n    if li == 0:\n        continue\n\n    higher_layers = [\n        x for x in sorted(layer_to_indices)\n        if x > li\n    ]\n\n    if not higher_layers:\n        continue\n\n    target_layer = higher_layers[0]\n\n    same_group_pool = layer_group_to_indices.get(\n        (target_layer, str(group_keys[i])),\n        []\n    )\n\n    if same_group_pool:\n        pool = np.asarray(\n            same_group_pool,\n            dtype=np.int64\n        )\n    else:\n        pool = np.asarray(\n            layer_to_indices[target_layer],\n            dtype=np.int64\n        )\n\n    n_take = min(3, len(pool))\n\n    if n_take == 0:\n        continue\n\n    if len(pool) > n_take:\n        chosen = rng.choice(\n            pool,\n            size=n_take,\n            replace=False\n        )\n    else:\n        chosen = pool\n\n    for j in np.asarray(chosen):\n        if int(layers[j]) > li:\n            compat_src.append(i)\n            compat_dst.append(int(j))\n\nif compat_src:\n    compat_ei = torch.tensor(\n        np.vstack([compat_src, compat_dst]),\n        dtype=torch.long\n    )\nelse:\n    compat_ei = torch.empty(\n        (2, 0),\n        dtype=torch.long\n    )\n\ngraph_data[\n    \"product\", \"compat_to\", \"product\"\n].edge_index = compat_ei\n\n# Keep the canonical alias synchronized.\ndata = graph_data\n\nprint(\n    \"Outfit layers:\",\n    {\n        int(k): int((layers == k).sum())\n        for k in np.unique(layers)\n    }\n)\nprint(\n    \"Directed compatibility DAG edges:\",\n    compat_ei.shape[1]\n)\n\nif compat_ei.numel():\n    layer_tensor = torch.tensor(\n        layers,\n        dtype=torch.long\n    )\n    assert bool(\n        torch.all(\n            layer_tensor[compat_ei[0]]\n            <\n            layer_tensor[compat_ei[1]]\n        )\n    )\n\nprint(\"DAG acyclicity condition verified.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:04:26.281856Z","iopub.execute_input":"2026-10-02T13:04:26.282209Z","iopub.status.idle":"2026-10-02T13:04:27.079687Z","shell.execute_reply.started":"2026-10-02T13:04:26.282182Z","shell.execute_reply":"2026-10-02T13:04:27.079055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 9 — CC-HGNN architecture: multi-scale capsule encoder + decoupled recommendation head\n\nclass CapsuleGate(nn.Module):\n    def __init__(self, d):\n        super().__init__()\n        self.gate = nn.Linear(d, d)\n        self.value = nn.Linear(d, d)\n        self.norm = nn.LayerNorm(d)\n    def forward(self, x):\n        g = torch.sigmoid(self.gate(x))\n        v = torch.tanh(self.value(x))\n        z = self.norm(g * v + (1.0 - g) * x)\n        norm = torch.norm(z, p=2, dim=-1, keepdim=True).clamp_min(1e-6)\n        scale = (norm * norm) / (1.0 + norm * norm)\n        return scale * z / norm\n\nclass CCHGNN(nn.Module):\n    \"\"\"Recommendation-first CC-HGNN.\n\n    Encoder: heterogeneous multi-head GAT over customer/product/attribute\n    relations, with raw highway + hop0 + hop1 + hop2 fusion and capsule gate.\n\n    Recommendation head: normalized customer/product projections + learned\n    item popularity bias + compatibility interaction MLP.\n    \"\"\"\n    def __init__(self, metadata, in_dims, hidden=64, heads=2, dropout=0.15, use_multiscale=True, num_products=10000):\n        super().__init__()\n        self.hidden = hidden\n        self.dropout = dropout\n        self.use_multiscale = use_multiscale\n        self.proj = nn.ModuleDict({nt: nn.Linear(in_dims[nt], hidden) for nt in in_dims})\n        self.raw_highway = nn.ModuleDict({nt: nn.Sequential(nn.Linear(in_dims[nt], hidden), nn.LayerNorm(hidden), nn.GELU()) for nt in in_dims})\n        def make_conv():\n            rels = {}\n            for src, rel, dst in metadata[1]:\n                rels[(src, rel, dst)] = GATConv((hidden, hidden), hidden, heads=heads, concat=False, add_self_loops=False, dropout=dropout)\n            return HeteroConv(rels, aggr='sum')\n        self.conv1 = make_conv(); self.conv2 = make_conv()\n        self.norm1 = nn.ModuleDict({nt: nn.LayerNorm(hidden) for nt in in_dims})\n        self.norm2 = nn.ModuleDict({nt: nn.LayerNorm(hidden) for nt in in_dims})\n        self.jk = nn.ModuleDict({nt: nn.Sequential(nn.Linear(hidden*4, hidden), nn.LayerNorm(hidden), nn.GELU(), nn.Dropout(dropout)) for nt in in_dims})\n        self.capsule = nn.ModuleDict({nt: CapsuleGate(hidden) for nt in in_dims})\n\n        # Decoupled recommendation projections, following the recommendation\n        # head idea in the supplied reference CC-HGNN notebook.\n        rec_dim = 64\n        self.rec_proj_customer = nn.Sequential(nn.Linear(hidden, rec_dim), nn.LayerNorm(rec_dim), nn.GELU())\n        self.rec_proj_product = nn.Sequential(nn.Linear(hidden, rec_dim), nn.LayerNorm(rec_dim), nn.GELU())\n        self.rec_logit_scale = nn.Parameter(torch.tensor(2.0))\n        self.item_bias = nn.Parameter(torch.zeros(int(num_products)))\n        self.compat_head = nn.Sequential(nn.Linear(hidden*4, hidden), nn.GELU(), nn.Dropout(dropout), nn.Linear(hidden, 1))\n        self.retention_head = nn.Sequential(nn.Linear(hidden, hidden), nn.GELU(), nn.Dropout(dropout), nn.Linear(hidden, 2))\n        self.vip_head = nn.Sequential(nn.Linear(hidden, hidden), nn.GELU(), nn.Dropout(dropout), nn.Linear(hidden, 2))\n\n    def _project(self, x_dict):\n        return {nt: F.gelu(self.proj[nt](x)) for nt, x in x_dict.items()}\n    def _message_layer(self, conv, norm, h, edge_index_dict):\n        out = conv(h, edge_index_dict); result = {}\n        for nt in h:\n            z = out.get(nt, torch.zeros_like(h[nt]))\n            z = F.dropout(F.gelu(norm[nt](z + h[nt])), p=self.dropout, training=self.training)\n            result[nt] = z\n        return result\n    def encode(self, x_dict, edge_index_dict):\n        raw = {nt: self.raw_highway[nt](x) for nt, x in x_dict.items()}\n        h0 = self._project(x_dict)\n        h1 = self._message_layer(self.conv1, self.norm1, h0, edge_index_dict)\n        h2 = self._message_layer(self.conv2, self.norm2, h1, edge_index_dict)\n        h = {}\n        for nt in h0:\n            multi = torch.cat([raw[nt], h0[nt], h1[nt], h2[nt]], dim=-1)\n            fused = self.jk[nt](multi) if self.use_multiscale else h2[nt]\n            h[nt] = fused + self.capsule[nt](fused)\n        return h\n    def forward(self, x_dict, edge_index_dict):\n        h = self.encode(x_dict, edge_index_dict)\n        return self.retention_head(h['customer']), self.vip_head(h['customer']), h\n    def rec_embeddings(self, customer_h, product_h):\n        c = F.normalize(self.rec_proj_customer(customer_h), p=2, dim=-1)\n        p = F.normalize(self.rec_proj_product(product_h), p=2, dim=-1)\n        return c, p\n    def pair_score(self, customer_h, product_h, product_indices=None):\n        c, p = self.rec_embeddings(customer_h, product_h)\n        cosine = (c * p).sum(dim=-1) * self.rec_logit_scale.clamp(0.5, 8.0)\n        interaction = customer_h * product_h\n        diff = torch.abs(customer_h - product_h)\n        compat = self.compat_head(torch.cat([customer_h, product_h, interaction, diff], dim=-1)).squeeze(-1)\n        if product_indices is None:\n            bias = torch.zeros_like(cosine)\n        else:\n            bias = self.item_bias[product_indices]\n        return cosine + 0.25 * compat + 0.10 * bias\n    def score_all_products(self, customer_h, product_h, start=0, end=None):\n        if end is None: end = product_h.shape[0]\n        c, p = self.rec_embeddings(customer_h, product_h[start:end])\n        scores = (c @ p.t()) * self.rec_logit_scale.clamp(0.5, 8.0)\n        scores = scores + 0.10 * self.item_bias[start:end].unsqueeze(0)\n        return scores\n\nprint('CC-HGNN recommendation-first architecture compiled.')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:04:27.080566Z","iopub.execute_input":"2026-10-02T13:04:27.080955Z","iopub.status.idle":"2026-10-02T13:04:27.102278Z","shell.execute_reply.started":"2026-10-02T13:04:27.080928Z","shell.execute_reply":"2026-10-02T13:04:27.101587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 10 — Final graph validation + safe device transfer\n# ============================================================\n#\n# This cell is intentionally self-contained.\n# It defines validate_graph_cpu() before calling it.\n#\n# Expected from the previous graph-construction cell:\n#     graph_data\n# or:\n#     data\n#\n# After this cell:\n#     data       -> validated HeteroData on the selected device\n#     graph_data -> same validated HeteroData object\n#\n# ============================================================\n\nimport numpy as np\nimport torch\nfrom torch_geometric.data import HeteroData\n\n\n# ------------------------------------------------------------\n# 1. Locate the graph created by the previous cell\n# ------------------------------------------------------------\n\nif \"graph_data\" in globals():\n\n    data = graph_data\n\nelif \"data\" in globals() and isinstance(data, HeteroData):\n\n    graph_data = data\n\nelse:\n\n    raise RuntimeError(\n        \"Neither 'graph_data' nor a valid 'data' HeteroData object \"\n        \"exists. Run the previous graph-construction cell first.\"\n    )\n\n\nif not isinstance(data, HeteroData):\n\n    raise TypeError(\n        f\"Expected torch_geometric.data.HeteroData, \"\n        f\"but received {type(data)}.\"\n    )\n\n\n# ------------------------------------------------------------\n# 2. Use one explicit device everywhere\n# ------------------------------------------------------------\n\ndevice = torch.device(\n    \"cuda:0\" if torch.cuda.is_available() else \"cpu\"\n)\n\nprint(\"=\" * 70)\nprint(\"CC-HGNN GRAPH VALIDATION\")\nprint(\"=\" * 70)\nprint(\"Target device:\", device)\n\n\n# ------------------------------------------------------------\n# 3. CPU-side structural validation\n# ------------------------------------------------------------\n\ndef validate_graph_cpu(g):\n    \"\"\"\n    Validate a PyG HeteroData graph before CUDA transfer.\n\n    Checks:\n      - node features exist\n      - node feature row counts match node counts\n      - edge indices are integer tensors\n      - edge indices are [2, E]\n      - edge indices are within valid node ranges\n      - customer targets have correct length\n      - train/validation/test masks have correct length\n      - masks are mutually exclusive\n      - all three masks contain at least one node\n    \"\"\"\n\n    if not isinstance(g, HeteroData):\n\n        raise TypeError(\n            \"Graph must be a torch_geometric.data.HeteroData object.\"\n        )\n\n\n    # --------------------------------------------------------\n    # Node validation\n    # --------------------------------------------------------\n\n    if len(g.node_types) == 0:\n\n        raise RuntimeError(\n            \"The HeteroData graph contains no node types.\"\n        )\n\n\n    for node_type in g.node_types:\n\n        store = g[node_type]\n\n        if not hasattr(store, \"x\") or store.x is None:\n\n            raise RuntimeError(\n                f\"Node type '{node_type}' has no feature tensor x.\"\n            )\n\n        x = store.x\n\n        if not torch.is_tensor(x):\n\n            raise RuntimeError(\n                f\"Node type '{node_type}' feature x is not a tensor.\"\n            )\n\n        if x.ndim != 2:\n\n            raise RuntimeError(\n                f\"Node type '{node_type}' feature tensor must be 2-D. \"\n                f\"Got shape {tuple(x.shape)}.\"\n            )\n\n        if x.shape[0] <= 0:\n\n            raise RuntimeError(\n                f\"Node type '{node_type}' contains zero nodes.\"\n            )\n\n        if not torch.is_floating_point(x):\n\n            raise RuntimeError(\n                f\"Node type '{node_type}' features must be floating point. \"\n                f\"Got dtype {x.dtype}.\"\n            )\n\n        if not torch.isfinite(x).all():\n\n            raise RuntimeError(\n                f\"Node type '{node_type}' contains NaN or Inf features.\"\n            )\n\n        if int(store.num_nodes) != int(x.shape[0]):\n\n            raise RuntimeError(\n                f\"Node count mismatch for '{node_type}': \"\n                f\"num_nodes={store.num_nodes}, \"\n                f\"x rows={x.shape[0]}.\"\n            )\n\n\n    # --------------------------------------------------------\n    # Edge validation\n    # --------------------------------------------------------\n\n    if len(g.edge_types) == 0:\n\n        raise RuntimeError(\n            \"The HeteroData graph contains no edge types.\"\n        )\n\n\n    for edge_type in g.edge_types:\n\n        src_type, relation, dst_type = edge_type\n\n        store = g[edge_type]\n\n        if not hasattr(store, \"edge_index\"):\n\n            raise RuntimeError(\n                f\"Edge type {edge_type} has no edge_index.\"\n            )\n\n        ei = store.edge_index\n\n        if not torch.is_tensor(ei):\n\n            raise RuntimeError(\n                f\"edge_index for {edge_type} is not a tensor.\"\n            )\n\n        if ei.ndim != 2 or ei.shape[0] != 2:\n\n            raise RuntimeError(\n                f\"edge_index for {edge_type} must have shape [2, E]. \"\n                f\"Got {tuple(ei.shape)}.\"\n            )\n\n        if ei.dtype != torch.long:\n\n            raise RuntimeError(\n                f\"edge_index for {edge_type} must have dtype torch.long. \"\n                f\"Got {ei.dtype}.\"\n            )\n\n        if ei.numel() == 0:\n\n            # Empty relations are allowed.\n            continue\n\n        src_max = int(ei[0].max().item())\n        src_min = int(ei[0].min().item())\n\n        dst_max = int(ei[1].max().item())\n        dst_min = int(ei[1].min().item())\n\n        src_nodes = int(g[src_type].num_nodes)\n        dst_nodes = int(g[dst_type].num_nodes)\n\n        if src_min < 0 or src_max >= src_nodes:\n\n            raise RuntimeError(\n                f\"Invalid source indices in edge type {edge_type}. \"\n                f\"Valid range: [0, {src_nodes - 1}], \"\n                f\"observed [{src_min}, {src_max}].\"\n            )\n\n        if dst_min < 0 or dst_max >= dst_nodes:\n\n            raise RuntimeError(\n                f\"Invalid destination indices in edge type {edge_type}. \"\n                f\"Valid range: [0, {dst_nodes - 1}], \"\n                f\"observed [{dst_min}, {dst_max}].\"\n            )\n\n\n    # --------------------------------------------------------\n    # Customer target/mask validation\n    # --------------------------------------------------------\n\n    if \"customer\" not in g.node_types:\n\n        raise RuntimeError(\n            \"The graph does not contain the required 'customer' node type.\"\n        )\n\n\n    customer_store = g[\"customer\"]\n\n    required_customer_fields = [\n        \"retention_y\",\n        \"vip_y\",\n        \"train_mask\",\n        \"val_mask\",\n        \"test_mask\",\n    ]\n\n    for field in required_customer_fields:\n\n        if not hasattr(customer_store, field):\n\n            raise RuntimeError(\n                f\"data['customer'].{field} is missing.\"\n            )\n\n\n    n_customers = int(\n        customer_store.num_nodes\n    )\n\n\n    # Targets\n\n    retention_y = customer_store.retention_y\n\n    vip_y = customer_store.vip_y\n\n    if retention_y.ndim != 1:\n\n        raise RuntimeError(\n            \"retention_y must be a 1-D tensor.\"\n        )\n\n    if vip_y.ndim != 1:\n\n        raise RuntimeError(\n            \"vip_y must be a 1-D tensor.\"\n        )\n\n    if len(retention_y) != n_customers:\n\n        raise RuntimeError(\n            \"retention_y length does not match customer count.\"\n        )\n\n    if len(vip_y) != n_customers:\n\n        raise RuntimeError(\n            \"vip_y length does not match customer count.\"\n        )\n\n\n    retention_values = set(\n        retention_y.detach().cpu().numpy().tolist()\n    )\n\n    vip_values = set(\n        vip_y.detach().cpu().numpy().tolist()\n    )\n\n    if not retention_values.issubset({0, 1}):\n\n        raise RuntimeError(\n            f\"retention_y contains values outside {{0,1}}: \"\n            f\"{retention_values}\"\n        )\n\n    if not vip_values.issubset({0, 1}):\n\n        raise RuntimeError(\n            f\"vip_y contains values outside {{0,1}}: \"\n            f\"{vip_values}\"\n        )\n\n\n    # Masks\n\n    train_mask = customer_store.train_mask\n    val_mask = customer_store.val_mask\n    test_mask = customer_store.test_mask\n\n    for name, mask in [\n        (\"train_mask\", train_mask),\n        (\"val_mask\", val_mask),\n        (\"test_mask\", test_mask),\n    ]:\n\n        if mask.ndim != 1:\n\n            raise RuntimeError(\n                f\"{name} must be 1-D.\"\n            )\n\n        if len(mask) != n_customers:\n\n            raise RuntimeError(\n                f\"{name} length does not match customer count.\"\n            )\n\n        if mask.dtype != torch.bool:\n\n            raise RuntimeError(\n                f\"{name} must have dtype torch.bool. \"\n                f\"Got {mask.dtype}.\"\n            )\n\n        if int(mask.sum()) == 0:\n\n            raise RuntimeError(\n                f\"{name} contains zero customers.\"\n            )\n\n\n    # --------------------------------------------------------\n    # Ensure train / validation / test are disjoint\n    # --------------------------------------------------------\n\n    overlap_train_val = (\n        train_mask & val_mask\n    ).any().item()\n\n    overlap_train_test = (\n        train_mask & test_mask\n    ).any().item()\n\n    overlap_val_test = (\n        val_mask & test_mask\n    ).any().item()\n\n\n    if overlap_train_val:\n\n        raise RuntimeError(\n            \"train_mask and val_mask overlap.\"\n        )\n\n    if overlap_train_test:\n\n        raise RuntimeError(\n            \"train_mask and test_mask overlap.\"\n        )\n\n    if overlap_val_test:\n\n        raise RuntimeError(\n            \"val_mask and test_mask overlap.\"\n        )\n\n\n    total_masked = (\n        int(train_mask.sum())\n        + int(val_mask.sum())\n        + int(test_mask.sum())\n    )\n\n    if total_masked != n_customers:\n\n        raise RuntimeError(\n            \"Train/validation/test masks do not cover \"\n            \"exactly all customers.\"\n        )\n\n\n    print(\"CPU graph validation: PASSED\")\n\n    print(\n        \"Node types:\",\n        list(g.node_types)\n    )\n\n    print(\n        \"Edge types:\",\n        len(g.edge_types)\n    )\n\n    print(\n        \"Customers:\",\n        n_customers\n    )\n\n    print(\n        \"Products:\",\n        int(g[\"product\"].num_nodes)\n        if \"product\" in g.node_types\n        else \"N/A\"\n    )\n\n    print(\n        \"Train customers:\",\n        int(train_mask.sum())\n    )\n\n    print(\n        \"Validation customers:\",\n        int(val_mask.sum())\n    )\n\n    print(\n        \"Test customers:\",\n        int(test_mask.sum())\n    )\n\n    return True\n\n\n# ------------------------------------------------------------\n# 4. Validate BEFORE moving anything to GPU\n# ------------------------------------------------------------\n\nvalidate_graph_cpu(data)\n\n\n# ------------------------------------------------------------\n# 5. Explicitly move the COMPLETE graph to one device\n# ------------------------------------------------------------\n\ndata = data.to(device)\n\n# Keep graph_data synchronized with data.\ngraph_data = data\n\n\n# ------------------------------------------------------------\n# 6. Post-transfer device validation\n# ------------------------------------------------------------\n\nfor node_type in data.node_types:\n\n    store = data[node_type]\n\n    if hasattr(store, \"x\") and store.x is not None:\n\n        if store.x.device != device:\n\n            raise RuntimeError(\n                f\"{node_type}.x is on {store.x.device}, \"\n                f\"expected {device}.\"\n            )\n\n\nfor edge_type in data.edge_types:\n\n    ei = data[edge_type].edge_index\n\n    if ei.device != device:\n\n        raise RuntimeError(\n            f\"edge_index for {edge_type} is on {ei.device}, \"\n            f\"expected {device}.\"\n        )\n\n\ncustomer = data[\"customer\"]\n\nfor field in [\n    \"retention_y\",\n    \"vip_y\",\n    \"train_mask\",\n    \"val_mask\",\n    \"test_mask\",\n]:\n\n    tensor = getattr(customer, field)\n\n    if tensor.device != device:\n\n        raise RuntimeError(\n            f\"customer.{field} is on {tensor.device}, \"\n            f\"expected {device}.\"\n        )\n\n\n# ------------------------------------------------------------\n# 7. Final graph summary\n# ------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"GRAPH READY FOR CC-HGNN\")\nprint(\"=\" * 70)\n\nprint(\"Device:\", device)\n\nprint(\n    \"All node features on target device:\",\n    all(\n        data[nt].x.device == device\n        for nt in data.node_types\n    )\n)\n\nprint(\n    \"All edge indices on target device:\",\n    all(\n        data[et].edge_index.device == device\n        for et in data.edge_types\n    )\n)\n\nprint(\n    \"Customer targets/masks on target device:\",\n    all(\n        getattr(data[\"customer\"], field).device == device\n        for field in [\n            \"retention_y\",\n            \"vip_y\",\n            \"train_mask\",\n            \"val_mask\",\n            \"test_mask\",\n        ]\n    )\n)\n\nprint()\nprint(\"FINAL GRAPH VALIDATION: PASSED\")\nprint(\"Graph is ready for the CC-HGNN training cell.\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:04:27.103441Z","iopub.execute_input":"2026-10-02T13:04:27.103731Z","iopub.status.idle":"2026-10-02T13:04:27.522372Z","shell.execute_reply.started":"2026-10-02T13:04:27.103701Z","shell.execute_reply":"2026-10-02T13:04:27.521682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 10.5 — CC-HGNN TRAINING\n# ============================================================\n\nimport torch\nimport torch.nn.functional as F\n\nprint(\"=\" * 70)\nprint(\"CC-HGNN TRAINING\")\nprint(\"=\" * 70)\n\n# ------------------------------------------------------------\n# 1. Make sure required objects exist\n# ------------------------------------------------------------\n\nrequired_objects = [\n    \"data\",\n    \"device\",\n    \"CCHGNN\"\n]\n\nmissing = [\n    obj for obj in required_objects\n    if obj not in globals()\n]\n\nif missing:\n    raise RuntimeError(\n        \"Missing required objects: \"\n        + \", \".join(missing)\n        + \". Run Cells 9 and 10 first.\"\n    )\n\n\n# ------------------------------------------------------------\n# 2. Graph metadata and input dimensions\n# ------------------------------------------------------------\n\nmetadata = data.metadata()\n\nin_dims = {\n    node_type: data[node_type].x.shape[1]\n    for node_type in data.node_types\n}\n\n\n# ------------------------------------------------------------\n# 3. Number of products\n# ------------------------------------------------------------\n\nif \"product\" in data.node_types:\n    num_products = int(data[\"product\"].num_nodes)\nelse:\n    num_products = 10000\n\n\n# ------------------------------------------------------------\n# 4. Create CC-HGNN model\n# ------------------------------------------------------------\n\ncc_hgnn_model = CCHGNN(\n    metadata=metadata,\n    in_dims=in_dims,\n    hidden=64,\n    heads=2,\n    dropout=0.15,\n    use_multiscale=True,\n    num_products=num_products\n).to(device)\n\n\nprint(\"CC-HGNN model created.\")\nprint(\"Input dimensions:\", in_dims)\nprint(\"Number of products:\", num_products)\nprint(\"Device:\", device)\n\n\n# ------------------------------------------------------------\n# 5. Customer targets and masks\n# ------------------------------------------------------------\n\nretention_y = (\n    data[\"customer\"]\n    .retention_y\n    .long()\n    .to(device)\n)\n\nvip_y = (\n    data[\"customer\"]\n    .vip_y\n    .long()\n    .to(device)\n)\n\ntrain_mask = (\n    data[\"customer\"]\n    .train_mask\n    .bool()\n    .to(device)\n)\n\nval_mask = (\n    data[\"customer\"]\n    .val_mask\n    .bool()\n    .to(device)\n)\n\ntest_mask = (\n    data[\"customer\"]\n    .test_mask\n    .bool()\n    .to(device)\n)\n\n\n# ------------------------------------------------------------\n# 6. Class weights\n# ------------------------------------------------------------\n\ndef make_class_weights(labels, mask):\n\n    labels_train = labels[mask]\n\n    counts = torch.bincount(\n        labels_train,\n        minlength=2\n    ).float()\n\n    counts = counts.clamp_min(1.0)\n\n    weights = (\n        counts.sum()\n        /\n        (2.0 * counts)\n    )\n\n    weights = (\n        weights\n        /\n        weights.mean().clamp_min(1e-8)\n    )\n\n    return torch.clamp(\n        weights,\n        min=0.5,\n        max=4.0\n    )\n\n\nretention_weights = make_class_weights(\n    retention_y,\n    train_mask\n)\n\nvip_weights = make_class_weights(\n    vip_y,\n    train_mask\n)\n\n\nprint(\"Retention class weights:\",\n      retention_weights.detach().cpu().numpy())\n\nprint(\"VIP class weights:\",\n      vip_weights.detach().cpu().numpy())\n\n\n# ------------------------------------------------------------\n# 7. Optimizer\n# ------------------------------------------------------------\n\noptimizer = torch.optim.AdamW(\n    cc_hgnn_model.parameters(),\n    lr=0.002,\n    weight_decay=1e-4\n)\n\n\n# ------------------------------------------------------------\n# 8. Training configuration\n# ------------------------------------------------------------\n\nEPOCHS = 50\n\nbest_val_loss = float(\"inf\")\nbest_state = None\npatience = 10\npatience_counter = 0\n\n\n# ------------------------------------------------------------\n# 9. Training loop\n# ------------------------------------------------------------\n\nfor epoch in range(1, EPOCHS + 1):\n\n    cc_hgnn_model.train()\n\n    optimizer.zero_grad(set_to_none=True)\n\n    retention_logits, vip_logits, _ = cc_hgnn_model(\n        data.x_dict,\n        data.edge_index_dict\n    )\n\n    retention_loss = F.cross_entropy(\n        retention_logits[train_mask],\n        retention_y[train_mask],\n        weight=retention_weights\n    )\n\n    vip_loss = F.cross_entropy(\n        vip_logits[train_mask],\n        vip_y[train_mask],\n        weight=vip_weights\n    )\n\n    # Main objective: retention\n    # Auxiliary objective: VIP classification\n    loss = (\n        retention_loss\n        +\n        0.30 * vip_loss\n    )\n\n    if not torch.isfinite(loss):\n        raise RuntimeError(\n            f\"CC-HGNN produced a non-finite loss \"\n            f\"at epoch {epoch}.\"\n        )\n\n    loss.backward()\n\n    torch.nn.utils.clip_grad_norm_(\n        cc_hgnn_model.parameters(),\n        max_norm=1.0\n    )\n\n    optimizer.step()\n\n\n    # --------------------------------------------------------\n    # Validation\n    # --------------------------------------------------------\n\n    cc_hgnn_model.eval()\n\n    with torch.no_grad():\n\n        val_retention_logits, val_vip_logits, _ = (\n            cc_hgnn_model(\n                data.x_dict,\n                data.edge_index_dict\n            )\n        )\n\n        val_retention_loss = F.cross_entropy(\n            val_retention_logits[val_mask],\n            retention_y[val_mask],\n            weight=retention_weights\n        )\n\n        val_vip_loss = F.cross_entropy(\n            val_vip_logits[val_mask],\n            vip_y[val_mask],\n            weight=vip_weights\n        )\n\n        val_loss = (\n            val_retention_loss\n            +\n            0.30 * val_vip_loss\n        )\n\n\n    # --------------------------------------------------------\n    # Save best model\n    # --------------------------------------------------------\n\n    if val_loss.item() < best_val_loss:\n\n        best_val_loss = val_loss.item()\n\n        best_state = {\n            key: value.detach().cpu().clone()\n            for key, value\n            in cc_hgnn_model.state_dict().items()\n        }\n\n        patience_counter = 0\n\n    else:\n\n        patience_counter += 1\n\n\n    # --------------------------------------------------------\n    # Progress\n    # --------------------------------------------------------\n\n    if (\n        epoch == 1\n        or epoch % 5 == 0\n    ):\n\n        print(\n            f\"Epoch {epoch:02d}/{EPOCHS} \"\n            f\"| Train Loss: {loss.item():.6f} \"\n            f\"| Retention: {retention_loss.item():.6f} \"\n            f\"| VIP: {vip_loss.item():.6f} \"\n            f\"| Val Loss: {val_loss.item():.6f}\"\n        )\n\n\n    # --------------------------------------------------------\n    # Early stopping\n    # --------------------------------------------------------\n\n    if patience_counter >= patience:\n\n        print(\n            f\"Early stopping at epoch {epoch}.\"\n        )\n\n        break\n\n\n# ------------------------------------------------------------\n# 10. Restore best model\n# ------------------------------------------------------------\n\nif best_state is not None:\n\n    cc_hgnn_model.load_state_dict(\n        best_state\n    )\n\n    cc_hgnn_model = cc_hgnn_model.to(device)\n\n\n# ------------------------------------------------------------\n# 11. Create compatibility alias\n# ------------------------------------------------------------\n# Cell 11 originally expects the variable \"model\".\n# Keep this alias so the existing Cell 11 can work.\n\nmodel = cc_hgnn_model\n\n\n# ------------------------------------------------------------\n# 12. Final message\n# ------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CC-HGNN TRAINING COMPLETED\")\nprint(\"=\" * 70)\n\nprint(\"Best validation loss:\",\n      best_val_loss)\n\nprint(\"Model variable:\")\nprint(\"    cc_hgnn_model\")\n\nprint(\"Compatibility variable:\")\nprint(\"    model\")\n\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:12:03.194071Z","iopub.execute_input":"2026-10-02T13:12:03.194882Z","iopub.status.idle":"2026-10-02T13:12:44.38108Z","shell.execute_reply.started":"2026-10-02T13:12:03.194846Z","shell.execute_reply":"2026-10-02T13:12:44.380397Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 11 — BASELINE CLASSIFICATION MODELS\n# ============================================================\n#\n# Models:\n#   1. Logistic Regression\n#   2. Random Forest\n#   3. MLP\n#   4. GraphSAGE\n#   5. CC-HGNN\n#\n# IMPORTANT:\n#   This cell is completely independent of ret_w.\n#   It computes its own GraphSAGE class weights.\n#\n# All models use:\n#   - the same customer cohort\n#   - the same train/test split\n#   - the same retention target\n#\n# ============================================================\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    balanced_accuracy_score,\n    precision_recall_fscore_support,\n    roc_auc_score,\n    average_precision_score\n)\n\nfrom sklearn.pipeline import make_pipeline\nfrom sklearn.preprocessing import StandardScaler\nfrom sklearn.linear_model import LogisticRegression\nfrom sklearn.ensemble import RandomForestClassifier\nfrom sklearn.neural_network import MLPClassifier\n\nfrom torch_geometric.nn import SAGEConv\n\n\nprint(\"=\" * 70)\nprint(\"FINAL BASELINE CLASSIFICATION COMPARISON\")\nprint(\"=\" * 70)\n\n\n# ============================================================\n# 1. REQUIRED OBJECTS\n# ============================================================\n\nrequired_objects = [\n    \"data\",\n    \"train_idx\",\n    \"test_idx\",\n    \"SEED\",\n    \"device\"\n]\n\nmissing = [\n    x for x in required_objects\n    if x not in globals()\n]\n\nif missing:\n    raise RuntimeError(\n        \"Cell 11 is missing required objects: \"\n        + \", \".join(missing)\n        + \". Run Cells 1–10 first.\"\n    )\n\n\n# ============================================================\n# 2. CUSTOMER FEATURES AND LABELS\n# ============================================================\n\nX = (\n    data[\"customer\"]\n    .x\n    .detach()\n    .cpu()\n    .numpy()\n)\n\ny = (\n    data[\"customer\"]\n    .retention_y\n    .detach()\n    .cpu()\n    .numpy()\n    .astype(np.int64)\n)\n\ntrain_idx_cpu = np.asarray(\n    train_idx,\n    dtype=np.int64\n)\n\ntest_idx_cpu = np.asarray(\n    test_idx,\n    dtype=np.int64\n)\n\nX_train = X[train_idx_cpu]\nX_test = X[test_idx_cpu]\n\ny_train = y[train_idx_cpu]\ny_test = y[test_idx_cpu]\n\n\nif len(X_train) == 0:\n    raise RuntimeError(\n        \"Training customer set is empty.\"\n    )\n\nif len(X_test) == 0:\n    raise RuntimeError(\n        \"Test customer set is empty.\"\n    )\n\n\n# ============================================================\n# 3. COMMON METRIC FUNCTION\n# ============================================================\n\ndef classification_metric_row(\n    model_name,\n    y_true,\n    predictions,\n    probabilities\n):\n\n    return {\n        \"Model\": model_name,\n\n        \"Accuracy\": accuracy_score(\n            y_true,\n            predictions\n        ),\n\n        \"Balanced Accuracy\": balanced_accuracy_score(\n            y_true,\n            predictions\n        ),\n\n        \"Macro-F1\": precision_recall_fscore_support(\n            y_true,\n            predictions,\n            average=\"macro\",\n            zero_division=0\n        )[2],\n\n        \"ROC-AUC\": (\n            roc_auc_score(\n                y_true,\n                probabilities\n            )\n            if len(np.unique(y_true)) == 2\n            else np.nan\n        ),\n\n        \"PR-AUC\": (\n            average_precision_score(\n                y_true,\n                probabilities\n            )\n            if len(np.unique(y_true)) == 2\n            else np.nan\n        )\n    }\n\n\nbaseline_rows = []\n\n\n# ============================================================\n# 4. LOGISTIC REGRESSION\n# ============================================================\n\nprint(\"\\nTraining Logistic Regression...\")\n\nlr_model = make_pipeline(\n\n    StandardScaler(),\n\n    LogisticRegression(\n        max_iter=500,\n        random_state=SEED,\n        class_weight=\"balanced\"\n    )\n)\n\nlr_model.fit(\n    X_train,\n    y_train\n)\n\nlr_pred = lr_model.predict(\n    X_test\n)\n\nlr_prob = lr_model.predict_proba(\n    X_test\n)[:, 1]\n\nbaseline_rows.append(\n    classification_metric_row(\n        \"Logistic Regression\",\n        y_test,\n        lr_pred,\n        lr_prob\n    )\n)\n\n\n# ============================================================\n# 5. RANDOM FOREST\n# ============================================================\n\nprint(\"Training Random Forest...\")\n\nrf_model = RandomForestClassifier(\n\n    n_estimators=120,\n\n    max_depth=18,\n\n    min_samples_leaf=2,\n\n    random_state=SEED,\n\n    n_jobs=-1,\n\n    class_weight=\"balanced\"\n\n)\n\nrf_model.fit(\n    X_train,\n    y_train\n)\n\nrf_pred = rf_model.predict(\n    X_test\n)\n\nrf_prob = rf_model.predict_proba(\n    X_test\n)[:, 1]\n\nbaseline_rows.append(\n    classification_metric_row(\n        \"Random Forest\",\n        y_test,\n        rf_pred,\n        rf_prob\n    )\n)\n\n\n# ============================================================\n# 6. MLP\n# ============================================================\n\nprint(\"Training MLP...\")\n\nmlp_model = make_pipeline(\n\n    StandardScaler(),\n\n    MLPClassifier(\n\n        hidden_layer_sizes=(\n            64,\n            32\n        ),\n\n        max_iter=180,\n\n        early_stopping=True,\n\n        random_state=SEED\n\n    )\n\n)\n\nmlp_model.fit(\n    X_train,\n    y_train\n)\n\nmlp_pred = mlp_model.predict(\n    X_test\n)\n\nmlp_prob = mlp_model.predict_proba(\n    X_test\n)[:, 1]\n\nbaseline_rows.append(\n    classification_metric_row(\n        \"MLP\",\n        y_test,\n        mlp_pred,\n        mlp_prob\n    )\n)\n\n\n# ============================================================\n# 7. GRAPHSAGE\n# ============================================================\n\nprint(\"Training GraphSAGE...\")\n\n\n# ------------------------------------------------------------\n# Peer graph\n# ------------------------------------------------------------\n\nsage_edge_type = (\n    \"customer\",\n    \"peer_with\",\n    \"customer\"\n)\n\nif sage_edge_type not in data.edge_types:\n\n    raise RuntimeError(\n        \"Customer peer graph is missing. \"\n        \"GraphSAGE cannot be evaluated.\"\n    )\n\n\nsage_edge = (\n    data[sage_edge_type]\n    .edge_index\n    .to(device)\n)\n\n\n# ------------------------------------------------------------\n# GraphSAGE model\n# ------------------------------------------------------------\n\nclass SAGEClassifier(nn.Module):\n\n    def __init__(\n        self,\n        input_dim,\n        hidden_dim=48\n    ):\n\n        super().__init__()\n\n        self.sage1 = SAGEConv(\n            input_dim,\n            hidden_dim\n        )\n\n        self.sage2 = SAGEConv(\n            hidden_dim,\n            hidden_dim\n        )\n\n        self.classifier = nn.Linear(\n            hidden_dim,\n            2\n        )\n\n    def forward(\n        self,\n        x,\n        edge_index\n    ):\n\n        h = self.sage1(\n            x,\n            edge_index\n        )\n\n        h = F.relu(h)\n\n        h = F.dropout(\n            h,\n            p=0.15,\n            training=self.training\n        )\n\n        h = self.sage2(\n            h,\n            edge_index\n        )\n\n        h = F.relu(h)\n\n        return self.classifier(h)\n\n\nsage = SAGEClassifier(\n    input_dim=data[\"customer\"].x.shape[1],\n    hidden_dim=48\n).to(device)\n\n\n# ------------------------------------------------------------\n# IMPORTANT:\n# Compute GraphSAGE weights locally.\n# Do NOT use ret_w.\n# ------------------------------------------------------------\n\ntrain_labels_gpu = (\n    data[\"customer\"]\n    .retention_y[\n        data[\"customer\"].train_mask\n    ]\n    .long()\n)\n\nclass_counts = torch.bincount(\n    train_labels_gpu,\n    minlength=2\n).float()\n\nclass_counts = class_counts.clamp_min(\n    1.0\n)\n\nsage_class_weights = (\n    class_counts.sum()\n    /\n    (2.0 * class_counts)\n)\n\nsage_class_weights = (\n    sage_class_weights\n    /\n    sage_class_weights.mean().clamp_min(\n        1e-8\n    )\n)\n\nsage_class_weights = torch.clamp(\n    sage_class_weights,\n    min=0.5,\n    max=4.0\n).to(device)\n\n\n# ------------------------------------------------------------\n# Optimizer\n# ------------------------------------------------------------\n\nsage_optimizer = torch.optim.AdamW(\n\n    sage.parameters(),\n\n    lr=0.002,\n\n    weight_decay=1e-4\n\n)\n\n\n# ------------------------------------------------------------\n# Training\n# ------------------------------------------------------------\n\nsage.train()\n\nfor epoch in range(\n    1,\n    31\n):\n\n    sage_optimizer.zero_grad(\n        set_to_none=True\n    )\n\n    sage_logits = sage(\n        data[\"customer\"].x,\n        sage_edge\n    )\n\n    sage_train_mask = (\n        data[\"customer\"]\n        .train_mask\n    )\n\n    sage_loss = F.cross_entropy(\n\n        sage_logits[\n            sage_train_mask\n        ],\n\n        data[\"customer\"].retention_y[\n            sage_train_mask\n        ],\n\n        weight=sage_class_weights\n\n    )\n\n    if not torch.isfinite(\n        sage_loss\n    ):\n\n        raise RuntimeError(\n            \"GraphSAGE produced a non-finite loss.\"\n        )\n\n    sage_loss.backward()\n\n    torch.nn.utils.clip_grad_norm_(\n        sage.parameters(),\n        max_norm=1.0\n    )\n\n    sage_optimizer.step()\n\n\n# ------------------------------------------------------------\n# GraphSAGE test predictions\n# ------------------------------------------------------------\n\nsage.eval()\n\nwith torch.no_grad():\n\n    sage_logits = sage(\n        data[\"customer\"].x,\n        sage_edge\n    )\n\n    sage_prob_all = (\n        torch.softmax(\n            sage_logits,\n            dim=1\n        )[:, 1]\n        .detach()\n        .cpu()\n        .numpy()\n    )\n\n    sage_pred_all = (\n        sage_logits\n        .argmax(dim=1)\n        .detach()\n        .cpu()\n        .numpy()\n    )\n\n\nsage_prob_test = (\n    sage_prob_all[\n        test_idx_cpu\n    ]\n)\n\nsage_pred_test = (\n    sage_pred_all[\n        test_idx_cpu\n    ]\n)\n\n\nbaseline_rows.append(\n    classification_metric_row(\n        \"GraphSAGE\",\n        y_test,\n        sage_pred_test,\n        sage_prob_test\n    )\n)\n\n\n# ============================================================\n# 8. CC-HGNN\n# ============================================================\n\nprint(\"Evaluating CC-HGNN...\")\n\n\n# ------------------------------------------------------------\n# Make sure trained CC-HGNN exists.\n# ------------------------------------------------------------\n\nif \"model\" not in globals():\n\n    raise RuntimeError(\n        \"CC-HGNN model is missing. \"\n        \"Run the CC-HGNN training cell first.\"\n    )\n\n\nmodel.eval()\n\n\nwith torch.no_grad():\n\n    cc_ret_logits, _, _ = model(\n        data.x_dict,\n        data.edge_index_dict\n    )\n\n    cc_prob_all = (\n        torch.softmax(\n            cc_ret_logits,\n            dim=1\n        )[:, 1]\n        .detach()\n        .cpu()\n        .numpy()\n    )\n\n    cc_pred_all = (\n        cc_ret_logits\n        .argmax(dim=1)\n        .detach()\n        .cpu()\n        .numpy()\n    )\n\n\ncc_prob_test = (\n    cc_prob_all[\n        test_idx_cpu\n    ]\n)\n\ncc_pred_test = (\n    cc_pred_all[\n        test_idx_cpu\n    ]\n)\n\n\nbaseline_rows.append(\n    classification_metric_row(\n        \"CC-HGNN\",\n        y_test,\n        cc_pred_test,\n        cc_prob_test\n    )\n)\n\n\n# ============================================================\n# 9. FINAL CLASSIFICATION TABLE\n# ============================================================\n\nclassification_results = pd.DataFrame(\n    baseline_rows\n)\n\nprint(\"\\n\")\nprint(\"=\" * 70)\nprint(\"FINAL CLASSIFICATION COMPARISON\")\nprint(\"=\" * 70)\n\nprint(\n    classification_results\n    .round(4)\n    .to_string(index=False)\n)\n\nprint(\"=\" * 70)\n\nprint(\n    \"\\nModels evaluated:\",\n    classification_results[\"Model\"].tolist()\n)\n\nprint(\n    \"\\nBest Accuracy:\",\n    classification_results.loc[\n        classification_results[\"Accuracy\"].idxmax(),\n        [\"Model\", \"Accuracy\"]\n    ].to_dict()\n)\n\nprint(\n    \"\\nBest ROC-AUC:\",\n    classification_results.loc[\n        classification_results[\"ROC-AUC\"].idxmax(),\n        [\"Model\", \"ROC-AUC\"]\n    ].to_dict()\n)\n\nprint(\n    \"\\nBest PR-AUC:\",\n    classification_results.loc[\n        classification_results[\"PR-AUC\"].idxmax(),\n        [\"Model\", \"PR-AUC\"]\n    ].to_dict()\n)\n\nprint(\n    \"\\nBest Macro-F1:\",\n    classification_results.loc[\n        classification_results[\"Macro-F1\"].idxmax(),\n        [\"Model\", \"Macro-F1\"]\n    ].to_dict()\n)\n\nprint(\n    \"\\nBest Balanced Accuracy:\",\n    classification_results.loc[\n        classification_results[\"Balanced Accuracy\"].idxmax(),\n        [\"Model\", \"Balanced Accuracy\"]\n    ].to_dict()\n)\n\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:13:04.552159Z","iopub.execute_input":"2026-10-02T13:13:04.552866Z","iopub.status.idle":"2026-10-02T13:13:14.935594Z","shell.execute_reply.started":"2026-10-02T13:13:04.552833Z","shell.execute_reply":"2026-10-02T13:13:14.935003Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 12 — CC-HGNN retention evaluation; threshold selected on validation only\n\nfrom sklearn.metrics import (\n    accuracy_score, balanced_accuracy_score,\n    precision_recall_fscore_support, roc_auc_score,\n    confusion_matrix, average_precision_score\n)\n\nmodel.eval()\nwith torch.no_grad():\n    ret_logits, vip_logits, h_eval = model(\n        data.x_dict, data.edge_index_dict\n    )\n\nval_mask = data[\"customer\"].val_mask\ntest_mask = data[\"customer\"].test_mask\n\nvy = data[\"customer\"].retention_y[val_mask].detach().cpu().numpy()\nvp = F.softmax(ret_logits[val_mask], dim=1)[:, 1].detach().cpu().numpy()\n\nbest_threshold = 0.5\nbest_bacc = -np.inf\n\nfor threshold in np.linspace(0.10, 0.90, 161):\n    pred = (vp >= threshold).astype(np.int64)\n    bacc = balanced_accuracy_score(vy, pred)\n    if bacc > best_bacc:\n        best_bacc = float(bacc)\n        best_threshold = float(threshold)\n\nty = data[\"customer\"].retention_y[test_mask].detach().cpu().numpy()\ntp = F.softmax(ret_logits[test_mask], dim=1)[:, 1].detach().cpu().numpy()\ntest_pred = (tp >= best_threshold).astype(np.int64)\n\ncc_result = {\n    \"Model\": \"CC-HGNN (validation threshold)\",\n    \"Accuracy\": accuracy_score(ty, test_pred),\n    \"Balanced Accuracy\": balanced_accuracy_score(ty, test_pred),\n    \"Macro-F1\": precision_recall_fscore_support(\n        ty, test_pred, average=\"macro\", zero_division=0\n    )[2],\n    \"ROC-AUC\": roc_auc_score(ty, tp) if len(np.unique(ty)) == 2 else np.nan,\n    \"PR-AUC\": average_precision_score(ty, tp) if len(np.unique(ty)) == 2 else np.nan\n}\n\nprint(pd.DataFrame([cc_result]).round(4).to_string(index=False))\nprint(\"\\nValidation-selected threshold:\", round(best_threshold, 4))\nprint(\"Validation balanced accuracy:\", round(best_bacc, 4))\nprint(\"\\nCC-HGNN confusion matrix:\")\nprint(confusion_matrix(ty, test_pred))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:13:14.936948Z","iopub.execute_input":"2026-10-02T13:13:14.937528Z","iopub.status.idle":"2026-10-02T13:13:15.303822Z","shell.execute_reply.started":"2026-10-02T13:13:14.9375Z","shell.execute_reply":"2026-10-02T13:13:15.303223Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 13 — PRIMARY TEMPORAL RECOMMENDATION EVALUATION\n# ============================================================\n#\n# Full historical-only catalog is the PRIMARY result.\n# Fixed sampled-candidate stress test is reported separately.\n# It never affects model selection.\n#\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"PRIMARY TEMPORAL RECOMMENDATION EVALUATION\")\nprint(\"=\" * 70)\n\n\n# ------------------------------------------------------------\n# 1. Put CC-HGNN in evaluation mode and obtain embeddings\n# ------------------------------------------------------------\n\nmodel.eval()\n\nwith torch.no_grad():\n    _, _, h_final = model(\n        data.x_dict,\n        data.edge_index_dict\n    )\n\ncustomer_h = h_final[\"customer\"]\nproduct_h = h_final[\"product\"]\n\n\n# ------------------------------------------------------------\n# 2. Prepare future evaluation data\n# ------------------------------------------------------------\n\nfuture_eval = future_df.copy()\n\nfuture_eval[\"customer_idx\"] = (\n    future_eval[\"customer_id\"]\n    .map(customer_id_to_idx)\n)\n\nfuture_eval[\"product_idx\"] = (\n    future_eval[\"article_id\"]\n    .map(product_id_to_idx)\n)\n\nfuture_eval = future_eval.dropna(\n    subset=[\"customer_idx\", \"product_idx\"]\n).copy()\n\nfuture_eval[\"customer_idx\"] = (\n    future_eval[\"customer_idx\"]\n    .astype(np.int64)\n)\n\nfuture_eval[\"product_idx\"] = (\n    future_eval[\"product_idx\"]\n    .astype(np.int64)\n)\n\n\n# ------------------------------------------------------------\n# 3. Future relevant items by customer\n# ------------------------------------------------------------\n\nfuture_items = (\n    future_eval\n    .groupby(\"customer_idx\")[\"product_idx\"]\n    .apply(lambda s: set(s.tolist()))\n    .to_dict()\n)\n\n\n# ------------------------------------------------------------\n# 4. Evaluation users\n# ------------------------------------------------------------\n\ntest_users = np.flatnonzero(test_mask_np)\n\neval_users = [\n    int(u)\n    for u in test_users\n    if int(u) in future_items\n    and future_items[int(u)]\n]\n\nprint(\"Evaluation users:\", len(eval_users))\n\n\n# ------------------------------------------------------------\n# 5. Recommendation configuration\n# ------------------------------------------------------------\n\nK_LIST = [3, 5, 10]\n\nMAX_K = 10\n\nPRODUCT_BATCH = 512\n\n\n# ------------------------------------------------------------\n# 6. Historical popularity\n# ------------------------------------------------------------\n\nhist_pop = np.zeros(\n    data[\"product\"].num_nodes,\n    dtype=np.float64\n)\n\nfor aid, n in hist_df[\"article_id\"].value_counts().items():\n\n    pi = product_id_to_idx.get(aid)\n\n    if pi is not None:\n        hist_pop[int(pi)] = float(n)\n\n\n# ------------------------------------------------------------\n# 7. Historical product pool\n# ------------------------------------------------------------\n# FIX:\n# This variable was missing in your original Cell 13.\n\nhistorical_product_pool = np.arange(\n    data[\"product\"].num_nodes,\n    dtype=np.int64\n)\n\n\n# ------------------------------------------------------------\n# 8. Popularity probabilities\n# ------------------------------------------------------------\n# Use historical frequency for negative sampling.\n# Add a small floor so zero-frequency products do not create\n# invalid probabilities.\n\npop_weights = hist_pop.copy()\n\npop_weights = pop_weights + 1e-8\n\npop_prob = (\n    pop_weights\n    /\n    pop_weights.sum()\n)\n\n\n# Safety check\n\nif not np.isfinite(pop_prob).all():\n\n    raise RuntimeError(\n        \"pop_prob contains NaN or Inf values.\"\n    )\n\nif not np.isclose(\n    pop_prob.sum(),\n    1.0,\n    atol=1e-6\n):\n\n    pop_prob = (\n        pop_prob\n        /\n        pop_prob.sum()\n    )\n\n\n# ------------------------------------------------------------\n# 9. Popularity score for recommendation baseline\n# ------------------------------------------------------------\n\npop_log = np.log1p(hist_pop)\n\npop_z = (\n    pop_log - pop_log.mean()\n) / (\n    pop_log.std() + 1e-8\n)\n\npop_z_t = torch.tensor(\n    pop_z,\n    dtype=torch.float32,\n    device=device\n)\n\n\n# ------------------------------------------------------------\n# 10. Items already seen by each user\n# ------------------------------------------------------------\n\nseen_by_user = {\n    int(uid): set(items)\n    for uid, items\n    in train_items_by_customer.items()\n}\n\n\n# ============================================================\n# 11. BPR-MF BASELINE\n# ============================================================\n\nprint()\nprint(\"=\" * 70)\nprint(\"TRAINING BPR-MF BASELINE\")\nprint(\"=\" * 70)\n\n\nMF_DIM = 48\n\nMF_EPOCHS = 25\n\nMF_PAIRS_PER_EPOCH = 16000\n\n\nmf_user = nn.Embedding(\n    data[\"customer\"].num_nodes,\n    MF_DIM\n).to(device)\n\nmf_item = nn.Embedding(\n    data[\"product\"].num_nodes,\n    MF_DIM\n).to(device)\n\n\nnn.init.normal_(\n    mf_user.weight,\n    std=0.07\n)\n\nnn.init.normal_(\n    mf_item.weight,\n    std=0.07\n)\n\n\nmf_opt = torch.optim.AdamW(\n    list(mf_user.parameters())\n    +\n    list(mf_item.parameters()),\n    lr=0.004,\n    weight_decay=1e-5\n)\n\n\nall_hist_users = np.asarray(\n    list(train_items_by_customer.keys()),\n    dtype=np.int64\n)\n\n\nfor ep in range(1, MF_EPOCHS + 1):\n\n    rng = np.random.default_rng(\n        SEED + ep\n    )\n\n    us = rng.choice(\n        all_hist_users,\n        size=min(\n            MF_PAIRS_PER_EPOCH,\n            len(all_hist_users)\n        ),\n        replace=True\n    )\n\n    U = []\n    P = []\n    N = []\n\n\n    for uid in us:\n\n        obs = train_items_by_customer.get(\n            int(uid),\n            set()\n        )\n\n        if not obs:\n            continue\n\n\n        # Positive item\n\n        pos = int(\n            rng.choice(\n                list(obs)\n            )\n        )\n\n\n        # Negative item\n\n        neg = int(\n            rng.choice(\n                historical_product_pool,\n                p=pop_prob\n            )\n        )\n\n\n        tries = 0\n\n        while (\n            neg in obs\n            and tries < 30\n        ):\n\n            neg = int(\n                rng.choice(\n                    historical_product_pool,\n                    p=pop_prob\n                )\n            )\n\n            tries += 1\n\n\n        if neg in obs:\n            continue\n\n\n        U.append(int(uid))\n        P.append(pos)\n        N.append(neg)\n\n\n    if not U:\n        continue\n\n\n    u = torch.tensor(\n        U,\n        dtype=torch.long,\n        device=device\n    )\n\n    p = torch.tensor(\n        P,\n        dtype=torch.long,\n        device=device\n    )\n\n    n = torch.tensor(\n        N,\n        dtype=torch.long,\n        device=device\n    )\n\n\n    pos_score = (\n        mf_user(u)\n        *\n        mf_item(p)\n    ).sum(1)\n\n\n    neg_score = (\n        mf_user(u)\n        *\n        mf_item(n)\n    ).sum(1)\n\n\n    loss = -F.logsigmoid(\n        pos_score - neg_score\n    ).mean()\n\n\n    mf_opt.zero_grad(\n        set_to_none=True\n    )\n\n    loss.backward()\n\n\n    torch.nn.utils.clip_grad_norm_(\n        list(mf_user.parameters())\n        +\n        list(mf_item.parameters()),\n        1.0\n    )\n\n\n    mf_opt.step()\n\n\n    if ep == 1 or ep % 5 == 0:\n\n        print(\n            f\"Epoch {ep:02d}/{MF_EPOCHS} \"\n            f\"| BPR Loss: {loss.item():.6f}\"\n        )\n\n\nmf_user.eval()\nmf_item.eval()\n\n\nmf_user_w = (\n    mf_user.weight\n    .detach()\n)\n\nmf_item_w = (\n    mf_item.weight\n    .detach()\n)\n\n\nprint(\"BPR-MF training completed.\")\n\n\n# ============================================================\n# 12. FULL-CATALOG EVALUATION\n# ============================================================\n\ndef eval_full(method):\n\n    hits = {\n        k: 0\n        for k in K_LIST\n    }\n\n    rec = {\n        k: 0.0\n        for k in K_LIST\n    }\n\n    nd = {\n        k: 0.0\n        for k in K_LIST\n    }\n\n\n    if len(eval_users) == 0:\n\n        return pd.DataFrame([\n            {\n                \"Method\": method,\n                \"K\": k,\n                \"Recall@K\": 0.0,\n                \"HitRate@K\": 0.0,\n                \"NDCG@K\": 0.0\n            }\n\n            for k in K_LIST\n        ])\n\n\n    with torch.no_grad():\n\n        for uid in eval_users:\n\n            relevant = future_items[uid]\n\n            nrel = len(relevant)\n\n            seen = seen_by_user.get(\n                uid,\n                set()\n            )\n\n\n            best_s = None\n            best_i = None\n\n\n            if method == \"CC-HGNN\":\n\n                uemb = customer_h[\n                    uid:uid + 1\n                ]\n\n            elif method == \"BPR-MF\":\n\n                uemb = mf_user_w[\n                    uid:uid + 1\n                ]\n\n\n            # ------------------------------------------------\n            # Score the COMPLETE product catalog in batches\n            # ------------------------------------------------\n\n            for st in range(\n                0,\n                product_h.shape[0],\n                PRODUCT_BATCH\n            ):\n\n                en = min(\n                    st + PRODUCT_BATCH,\n                    product_h.shape[0]\n                )\n\n\n                pids = torch.arange(\n                    st,\n                    en,\n                    device=device\n                )\n\n\n                if method == \"CC-HGNN\":\n\n                    scores = (\n                        model.score_all_products(\n                            uemb,\n                            product_h,\n                            st,\n                            en\n                        )\n                        .reshape(-1)\n                    )\n\n\n                elif method == \"BPR-MF\":\n\n                    scores = (\n                        uemb.expand(\n                            en - st,\n                            -1\n                        )\n                        *\n                        mf_item_w[\n                            st:en\n                        ]\n                    ).sum(1)\n\n\n                else:\n\n                    # Popularity baseline\n\n                    scores = pop_z_t[\n                        st:en\n                    ]\n\n\n                # --------------------------------------------\n                # Remove products already seen during training\n                # --------------------------------------------\n\n                local = [\n                    x - st\n                    for x in seen\n                    if st <= x < en\n                ]\n\n\n                if local:\n\n                    scores = scores.clone()\n\n                    scores[\n                        torch.tensor(\n                            local,\n                            dtype=torch.long,\n                            device=device\n                        )\n                    ] = -torch.inf\n\n\n                # --------------------------------------------\n                # Keep local top-K\n                # --------------------------------------------\n\n                kk = min(\n                    MAX_K,\n                    scores.numel()\n                )\n\n                ls, li = torch.topk(\n                    scores,\n                    kk\n                )\n\n                li = li + st\n\n\n                # --------------------------------------------\n                # Merge with global top-K\n                # --------------------------------------------\n\n                if best_s is None:\n\n                    best_s = ls\n                    best_i = li\n\n                else:\n\n                    ms = torch.cat(\n                        [best_s, ls]\n                    )\n\n                    mi = torch.cat(\n                        [best_i, li]\n                    )\n\n                    keep = min(\n                        MAX_K,\n                        ms.numel()\n                    )\n\n                    best_s, pos = torch.topk(\n                        ms,\n                        keep\n                    )\n\n                    best_i = mi[pos]\n\n\n            # ------------------------------------------------\n            # Evaluate ranking\n            # ------------------------------------------------\n\n            ranked = (\n                best_i\n                .cpu()\n                .tolist()\n            )\n\n\n            for k in K_LIST:\n\n                top = ranked[:k]\n\n\n                ranks = [\n                    r\n                    for r, item\n                    in enumerate(top, 1)\n                    if int(item) in relevant\n                ]\n\n\n                nh = len(ranks)\n\n\n                rec[k] += (\n                    nh / nrel\n                )\n\n\n                if nh:\n\n                    hits[k] += 1\n\n\n                    dcg = sum(\n                        1 / np.log2(r + 1)\n                        for r in ranks\n                    )\n\n\n                    ideal = min(\n                        k,\n                        nrel\n                    )\n\n\n                    idcg = sum(\n                        1 / np.log2(r + 1)\n                        for r in range(\n                            1,\n                            ideal + 1\n                        )\n                    )\n\n\n                    if idcg:\n\n                        nd[k] += (\n                            dcg / idcg\n                        )\n\n\n    return pd.DataFrame([\n        {\n            \"Method\": method,\n            \"K\": k,\n            \"Recall@K\": (\n                rec[k] / len(eval_users)\n            ),\n            \"HitRate@K\": (\n                hits[k] / len(eval_users)\n            ),\n            \"NDCG@K\": (\n                nd[k] / len(eval_users)\n            )\n        }\n\n        for k in K_LIST\n    ])\n\n\n# ------------------------------------------------------------\n# 13. Primary comparison\n# ------------------------------------------------------------\n\ncomparison_results = pd.concat(\n    [\n        eval_full(\"CC-HGNN\"),\n        eval_full(\"BPR-MF\"),\n        eval_full(\"Popularity\")\n    ],\n    ignore_index=True\n)\n\n\nrecommendation_results = (\n    comparison_results[\n        comparison_results.Method\n        == \"CC-HGNN\"\n    ]\n    .drop(columns=\"Method\")\n    .reset_index(drop=True)\n)\n\n\n# ============================================================\n# 14. FIXED SAMPLED-CANDIDATE STRESS TEST\n# ============================================================\n\nSAMPLED_NEG = 99\n\nsampled_rows = []\n\nrng = np.random.default_rng(\n    20261001\n)\n\n\nfor method in [\n    \"CC-HGNN\",\n    \"BPR-MF\",\n    \"Popularity\"\n]:\n\n    # [Recall sum, HitRate sum, NDCG sum, user count]\n    sums = {\n        k: [0.0, 0.0, 0.0, 0]\n        for k in K_LIST\n    }\n\n\n    with torch.no_grad():\n\n        for uid in eval_users:\n\n            rel = list(\n                future_items[uid]\n            )\n\n            seen = seen_by_user.get(\n                uid,\n                set()\n            )\n\n\n            pos = rel[:]\n\n\n            neg = []\n\n            attempts = 0\n\n\n            while (\n                len(neg) < SAMPLED_NEG\n                and\n                attempts < SAMPLED_NEG * 20\n            ):\n\n                x = int(\n                    rng.choice(\n                        historical_product_pool,\n                        p=pop_prob\n                    )\n                )\n\n                attempts += 1\n\n\n                if (\n                    x not in seen\n                    and\n                    x not in rel\n                    and\n                    x not in neg\n                ):\n\n                    neg.append(x)\n\n\n            if not neg or not pos:\n                continue\n\n\n            cand = pos + neg\n\n\n            pp = torch.tensor(\n                cand,\n                dtype=torch.long,\n                device=device\n            )\n\n\n            uu = torch.full(\n                (len(cand),),\n                uid,\n                dtype=torch.long,\n                device=device\n            )\n\n\n            # --------------------------------------------\n            # Score candidates\n            # --------------------------------------------\n\n            if method == \"CC-HGNN\":\n\n                scores = model.pair_score(\n                    customer_h[uu],\n                    product_h[pp],\n                    pp\n                )\n\n\n            elif method == \"BPR-MF\":\n\n                scores = (\n                    mf_user_w[uu]\n                    *\n                    mf_item_w[pp]\n                ).sum(1)\n\n\n            else:\n\n                scores = pop_z_t[pp]\n\n\n            order = (\n                torch.argsort(\n                    scores,\n                    descending=True\n                )\n                .cpu()\n                .tolist()\n            )\n\n\n            ranked = [\n                cand[i]\n                for i in order\n            ]\n\n\n            relset = set(rel)\n\n            nrel = len(relset)\n\n\n            # --------------------------------------------\n            # Metrics\n            # --------------------------------------------\n\n            for k in K_LIST:\n\n                top = ranked[:k]\n\n\n                ranks = [\n                    r\n                    for r, it\n                    in enumerate(top, 1)\n                    if it in relset\n                ]\n\n\n                nh = len(ranks)\n\n\n                sums[k][0] += (\n                    nh / nrel\n                )\n\n                sums[k][1] += (\n                    1 if nh else 0\n                )\n\n\n                if nh:\n\n                    dcg = sum(\n                        1 / np.log2(r + 1)\n                        for r in ranks\n                    )\n\n\n                    ideal = min(\n                        k,\n                        nrel\n                    )\n\n\n                    idcg = sum(\n                        1 / np.log2(r + 1)\n                        for r in range(\n                            1,\n                            ideal + 1\n                        )\n                    )\n\n\n                    if idcg:\n\n                        sums[k][2] += (\n                            dcg / idcg\n                        )\n\n\n                # FIX:\n                # Count the evaluated user.\n                sums[k][3] += 1\n\n\n    # --------------------------------------------------------\n    # Aggregate sampled results\n    # --------------------------------------------------------\n\n    for k in K_LIST:\n\n        denom = sums[k][3]\n\n\n        if denom == 0:\n            denom = 1\n\n\n        sampled_rows.append(\n            {\n                \"Method\": method,\n                \"K\": k,\n                \"Recall@K\": (\n                    sums[k][0] / denom\n                ),\n                \"HitRate@K\": (\n                    sums[k][1] / denom\n                ),\n                \"NDCG@K\": (\n                    sums[k][2] / denom\n                )\n            }\n        )\n\n\nsampled_candidate_results = pd.DataFrame(\n    sampled_rows\n)\n\n\n# ============================================================\n# 15. FINAL RESULTS\n# ============================================================\n\nprint()\nprint(\"=\" * 70)\nprint(\"PRIMARY FULL-CATALOG RECOMMENDATION\")\nprint(\"=\" * 70)\n\nprint(\n    comparison_results\n    .round(6)\n    .to_string(index=False)\n)\n\nprint()\n\nprint(\n    \"Evaluated users:\",\n    len(eval_users)\n)\n\nprint(\n    \"Candidate products:\",\n    data[\"product\"].num_nodes\n)\n\n\nprint()\nprint(\n    \"FIXED SAMPLED-CANDIDATE STRESS TEST \"\n    \"(99 negatives + future positives; descriptive only)\"\n)\n\nprint(\n    sampled_candidate_results\n    .round(6)\n    .to_string(index=False)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:13:15.304811Z","iopub.execute_input":"2026-10-02T13:13:15.305114Z","iopub.status.idle":"2026-10-02T13:15:55.925346Z","shell.execute_reply.started":"2026-10-02T13:13:15.305087Z","shell.execute_reply":"2026-10-02T13:15:55.924766Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 14 — MEMORY-SAFE FINAL ARCHITECTURAL ABLATION\n# ============================================================\n#\n# Full CC-HGNN remains unchanged.\n# Ablation models are deliberately smaller/shorter because\n# ablation is a supporting experiment and full-graph GAT\n# consumes substantial GPU memory.\n#\n# Validation AUC = model-selection metric\n# Test AUC       = descriptive only\n# ============================================================\n\nimport gc\nimport copy\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn.functional as F\n\n\nprint(\"=\" * 70)\nprint(\"MEMORY-SAFE ARCHITECTURAL ABLATION\")\nprint(\"=\" * 70)\n\n\n# ------------------------------------------------------------\n# 1. GPU cleanup before starting\n# ------------------------------------------------------------\n\nif torch.cuda.is_available():\n\n    torch.cuda.empty_cache()\n\n    gc.collect()\n\n    print(\n        f\"GPU memory allocated before ablation: \"\n        f\"{torch.cuda.memory_allocated() / 1024**3:.2f} GB\"\n    )\n\n\n# ------------------------------------------------------------\n# 2. Loss functions\n# ------------------------------------------------------------\n\ndef ret_loss_fn(logits, targets):\n\n    return F.cross_entropy(\n        logits,\n        targets.long()\n    )\n\n\ndef vip_loss_fn(logits, targets):\n\n    return F.cross_entropy(\n        logits,\n        targets.long()\n    )\n\n\n# ------------------------------------------------------------\n# 3. Ablation variants\n# ------------------------------------------------------------\n\nvariants = [\n    (\"Full CC-HGNN\", True, True, True),\n    (\"w/o Multi-scale JK\", False, True, True),\n    (\"w/o Peer Graph\", True, False, True),\n    (\"w/o Compatibility DAG\", True, True, False)\n]\n\n\n# ============================================================\n# 4. Pair sampler\n# ============================================================\n\ndef sample_pairs_safe(\n    h,\n    rng,\n    num_pairs=1500\n):\n\n    if \"train_items_by_customer\" not in globals():\n\n        return []\n\n\n    if \"customer\" not in h or \"product\" not in h:\n\n        return []\n\n\n    product_count = h[\"product\"].shape[0]\n\n    users = list(\n        train_items_by_customer.keys()\n    )\n\n    if not users:\n\n        return []\n\n\n    pairs = []\n\n    attempts = 0\n\n    max_attempts = num_pairs * 5\n\n\n    while (\n        len(pairs) < num_pairs\n        and\n        attempts < max_attempts\n    ):\n\n        attempts += 1\n\n        uid = int(\n            rng.choice(users)\n        )\n\n\n        positives = train_items_by_customer.get(\n            uid,\n            set()\n        )\n\n\n        if not positives:\n\n            continue\n\n\n        pos = int(\n            rng.choice(\n                list(positives)\n            )\n        )\n\n\n        neg = int(\n            rng.integers(\n                0,\n                product_count\n            )\n        )\n\n\n        tries = 0\n\n        while (\n            neg in positives\n            and\n            tries < 20\n        ):\n\n            neg = int(\n                rng.integers(\n                    0,\n                    product_count\n                )\n            )\n\n            tries += 1\n\n\n        if neg in positives:\n\n            continue\n\n\n        pairs.append(\n            (uid, pos, neg)\n        )\n\n\n    return pairs\n\n\n# ============================================================\n# 5. Lightweight validation AUC\n# ============================================================\n\ndef validation_auc_safe(\n    model,\n    h,\n    rng,\n    max_users=200,\n    negatives_per_positive=4\n):\n\n    if \"train_items_by_customer\" not in globals():\n\n        return np.nan\n\n\n    users = list(\n        train_items_by_customer.keys()\n    )\n\n\n    if not users:\n\n        return np.nan\n\n\n    if len(users) > max_users:\n\n        users = rng.choice(\n            users,\n            size=max_users,\n            replace=False\n        )\n\n\n    wins = 0.0\n\n    total = 0\n\n\n    with torch.no_grad():\n\n        for uid in users:\n\n            uid = int(uid)\n\n\n            positives = train_items_by_customer.get(\n                uid,\n                set()\n            )\n\n\n            if not positives:\n\n                continue\n\n\n            pos_list = list(\n                positives\n            )\n\n\n            # Evaluate only a small number of positives.\n            pos_list = pos_list[:3]\n\n\n            for pos in pos_list:\n\n                negs = []\n\n                attempts = 0\n\n\n                while (\n                    len(negs)\n                    <\n                    negatives_per_positive\n                    and\n                    attempts\n                    <\n                    negatives_per_positive * 15\n                ):\n\n                    neg = int(\n                        rng.integers(\n                            0,\n                            h[\"product\"].shape[0]\n                        )\n                    )\n\n                    attempts += 1\n\n\n                    if (\n                        neg not in positives\n                        and\n                        neg not in negs\n                    ):\n\n                        negs.append(\n                            neg\n                        )\n\n\n                if not negs:\n\n                    continue\n\n\n                u = torch.tensor(\n                    [uid],\n                    dtype=torch.long,\n                    device=device\n                )\n\n\n                p = torch.tensor(\n                    [int(pos)],\n                    dtype=torch.long,\n                    device=device\n                )\n\n\n                pos_score = model.pair_score(\n                    h[\"customer\"][u],\n                    h[\"product\"][p],\n                    p\n                )\n\n\n                n = torch.tensor(\n                    negs,\n                    dtype=torch.long,\n                    device=device\n                )\n\n\n                user_expand = h[\"customer\"][u].expand(\n                    len(negs),\n                    -1\n                )\n\n\n                neg_score = model.pair_score(\n                    user_expand,\n                    h[\"product\"][n],\n                    n\n                )\n\n\n                ps = float(\n                    pos_score.item()\n                )\n\n\n                ns = neg_score.detach().cpu().numpy()\n\n\n                wins += float(\n                    (ps > ns).sum()\n                )\n\n\n                total += len(ns)\n\n\n                del u, p, n\n                del pos_score, neg_score\n\n\n    if total == 0:\n\n        return np.nan\n\n\n    return wins / total\n\n\n# ============================================================\n# 6. Lightweight sampled softmax\n# ============================================================\n\ndef sampled_softmax_safe(\n    model,\n    h,\n    rng,\n    max_users=64,\n    negatives=5\n):\n\n    if \"train_items_by_customer\" not in globals():\n\n        return torch.zeros(\n            (),\n            device=device\n        )\n\n\n    users = list(\n        train_items_by_customer.keys()\n    )\n\n\n    if not users:\n\n        return torch.zeros(\n            (),\n            device=device\n        )\n\n\n    if len(users) > max_users:\n\n        users = rng.choice(\n            users,\n            size=max_users,\n            replace=False\n        )\n\n\n    losses = []\n\n\n    for uid in users:\n\n        uid = int(uid)\n\n\n        positives = train_items_by_customer.get(\n            uid,\n            set()\n        )\n\n\n        if not positives:\n\n            continue\n\n\n        pos = int(\n            rng.choice(\n                list(positives)\n            )\n        )\n\n\n        negs = []\n\n\n        for _ in range(negatives):\n\n            neg = int(\n                rng.integers(\n                    0,\n                    h[\"product\"].shape[0]\n                )\n            )\n\n\n            tries = 0\n\n            while (\n                neg in positives\n                and\n                tries < 20\n            ):\n\n                neg = int(\n                    rng.integers(\n                        0,\n                        h[\"product\"].shape[0]\n                    )\n                )\n\n                tries += 1\n\n\n            if neg not in positives:\n\n                negs.append(\n                    neg\n                )\n\n\n        if not negs:\n\n            continue\n\n\n        u = torch.tensor(\n            [uid],\n            dtype=torch.long,\n            device=device\n        )\n\n\n        p = torch.tensor(\n            [pos],\n            dtype=torch.long,\n            device=device\n        )\n\n\n        n = torch.tensor(\n            negs,\n            dtype=torch.long,\n            device=device\n        )\n\n\n        pos_score = model.pair_score(\n            h[\"customer\"][u],\n            h[\"product\"][p],\n            p\n        )\n\n\n        user_expand = h[\"customer\"][u].expand(\n            len(negs),\n            -1\n        )\n\n\n        neg_score = model.pair_score(\n            user_expand,\n            h[\"product\"][n],\n            n\n        )\n\n\n        logits = torch.cat(\n            [\n                pos_score.reshape(1),\n                neg_score.reshape(-1)\n            ]\n        ).unsqueeze(0)\n\n\n        target = torch.zeros(\n            1,\n            dtype=torch.long,\n            device=device\n        )\n\n\n        losses.append(\n            F.cross_entropy(\n                logits,\n                target\n            )\n        )\n\n\n        del u, p, n\n        del pos_score, neg_score\n        del logits, target\n\n\n    if not losses:\n\n        return torch.zeros(\n            (),\n            device=device\n        )\n\n\n    return torch.stack(\n        losses\n    ).mean()\n\n\n# ============================================================\n# 7. Run one ablation\n# ============================================================\n\ndef quick_ablation_safe(\n    name,\n    use_multiscale,\n    use_peer,\n    use_dag,\n    seed=42\n):\n\n    print()\n    print(\"-\" * 70)\n    print(\"Running:\", name)\n    print(\"-\" * 70)\n\n\n    # --------------------------------------------------------\n    # Aggressive cleanup before each variant\n    # --------------------------------------------------------\n\n    gc.collect()\n\n    if torch.cuda.is_available():\n\n        torch.cuda.empty_cache()\n\n\n    # --------------------------------------------------------\n    # Clone graph\n    # --------------------------------------------------------\n\n    gd = data.clone()\n\n\n    # --------------------------------------------------------\n    # Remove peer graph\n    # --------------------------------------------------------\n\n    if not use_peer:\n\n        et = (\n            \"customer\",\n            \"peer_with\",\n            \"customer\"\n        )\n\n\n        if et in gd.edge_types:\n\n            gd[et].edge_index = torch.empty(\n                (2, 0),\n                dtype=torch.long,\n                device=device\n            )\n\n\n    # --------------------------------------------------------\n    # Remove compatibility DAG\n    # --------------------------------------------------------\n\n    if not use_dag:\n\n        et = (\n            \"product\",\n            \"compat_to\",\n            \"product\"\n        )\n\n\n        if et in gd.edge_types:\n\n            gd[et].edge_index = torch.empty(\n                (2, 0),\n                dtype=torch.long,\n                device=device\n            )\n\n\n    # --------------------------------------------------------\n    # IMPORTANT:\n    # Smaller model ONLY for ablation.\n    #\n    # Your main CC-HGNN is NOT changed.\n    # --------------------------------------------------------\n\n    ABLATION_HIDDEN = 32\n\n\n    m = CCHGNN(\n        gd.metadata(),\n        {\n            nt: int(\n                gd[nt].x.shape[1]\n            )\n            for nt in gd.node_types\n        },\n        hidden=ABLATION_HIDDEN,\n        heads=1,\n        dropout=0.15,\n        use_multiscale=use_multiscale,\n        num_products=int(\n            gd[\"product\"].num_nodes\n        )\n    ).to(device)\n\n\n    opt = torch.optim.AdamW(\n        m.parameters(),\n        lr=0.0018,\n        weight_decay=1e-4\n    )\n\n\n    rng = np.random.default_rng(\n        seed\n    )\n\n\n    best = -np.inf\n\n    best_state = None\n\n\n    # --------------------------------------------------------\n    # Only 8 epochs for supporting ablation\n    # --------------------------------------------------------\n\n    ABLATION_EPOCHS = 8\n\n\n    for ep in range(\n        ABLATION_EPOCHS\n    ):\n\n        m.train()\n\n        opt.zero_grad(\n            set_to_none=True\n        )\n\n\n        # --------------------------------------------\n        # Full graph forward\n        # --------------------------------------------\n\n        rl, vl, h = m(\n            gd.x_dict,\n            gd.edge_index_dict\n        )\n\n\n        tr = gd[\"customer\"].train_mask\n\n\n        rloss = ret_loss_fn(\n            rl[tr],\n            gd[\"customer\"].retention_y[tr]\n        )\n\n\n        vloss = vip_loss_fn(\n            vl[tr],\n            gd[\"customer\"].vip_y[tr]\n        )\n\n\n        # --------------------------------------------\n        # BPR pairs\n        # --------------------------------------------\n\n        pairs = sample_pairs_safe(\n            h,\n            rng,\n            num_pairs=1500\n        )\n\n\n        if pairs:\n\n            u = torch.tensor(\n                [x[0] for x in pairs],\n                dtype=torch.long,\n                device=device\n            )\n\n\n            p = torch.tensor(\n                [x[1] for x in pairs],\n                dtype=torch.long,\n                device=device\n            )\n\n\n            n = torch.tensor(\n                [x[2] for x in pairs],\n                dtype=torch.long,\n                device=device\n            )\n\n\n            pos_score = m.pair_score(\n                h[\"customer\"][u],\n                h[\"product\"][p],\n                p\n            )\n\n\n            neg_score = m.pair_score(\n                h[\"customer\"][u],\n                h[\"product\"][n],\n                n\n            )\n\n\n            b = -F.logsigmoid(\n                pos_score - neg_score\n            ).mean()\n\n\n            del pos_score\n            del neg_score\n            del u, p, n\n\n        else:\n\n            b = torch.zeros(\n                (),\n                device=device\n            )\n\n\n        # --------------------------------------------\n        # Lightweight softmax\n        # --------------------------------------------\n\n        ssl = sampled_softmax_safe(\n            m,\n            h,\n            rng,\n            max_users=64,\n            negatives=5\n        )\n\n\n        loss = (\n            0.70 * b\n            +\n            0.20 * ssl\n            +\n            0.08 * rloss\n            +\n            0.02 * vloss\n        )\n\n\n        loss.backward()\n\n\n        torch.nn.utils.clip_grad_norm_(\n            m.parameters(),\n            1.5\n        )\n\n\n        opt.step()\n\n\n        # --------------------------------------------\n        # Validation\n        # --------------------------------------------\n\n        m.eval()\n\n\n        with torch.no_grad():\n\n            _, _, hv = m(\n                gd.x_dict,\n                gd.edge_index_dict\n            )\n\n\n        auc = validation_auc_safe(\n            m,\n            hv,\n            rng,\n            max_users=200,\n            negatives_per_positive=4\n        )\n\n\n        if (\n            np.isfinite(auc)\n            and\n            auc > best\n        ):\n\n            best = float(auc)\n\n            best_state = copy.deepcopy(\n                m.state_dict()\n            )\n\n\n        print(\n            f\"Epoch {ep + 1:02d}/{ABLATION_EPOCHS} \"\n            f\"| Loss: {loss.item():.5f} \"\n            f\"| Val AUC: {auc:.5f}\"\n        )\n\n\n        # ------------------------------------------------\n        # Release temporary tensors\n        # ------------------------------------------------\n\n        del rl, vl, h, hv\n        del loss, b, ssl, rloss, vloss\n\n\n        gc.collect()\n\n        if torch.cuda.is_available():\n\n            torch.cuda.empty_cache()\n\n\n    # --------------------------------------------------------\n    # Restore best validation model\n    # --------------------------------------------------------\n\n    if best_state is not None:\n\n        m.load_state_dict(\n            best_state\n        )\n\n\n    # --------------------------------------------------------\n    # Descriptive test AUC\n    # --------------------------------------------------------\n\n    m.eval()\n\n\n    with torch.no_grad():\n\n        _, _, hf = m(\n            gd.x_dict,\n            gd.edge_index_dict\n        )\n\n\n    test_auc = validation_auc_safe(\n        m,\n        hf,\n        np.random.default_rng(\n            seed + 7000\n        ),\n        max_users=300,\n        negatives_per_positive=5\n    )\n\n\n    result = {\n        \"Configuration\": name,\n        \"Validation Ranking AUC\": best,\n        \"Test Pair AUC\": test_auc\n    }\n\n\n    # --------------------------------------------------------\n    # CRITICAL MEMORY CLEANUP\n    # --------------------------------------------------------\n\n    del hf\n    del gd\n\n    del m\n    del opt\n    del best_state\n\n\n    gc.collect()\n\n    if torch.cuda.is_available():\n\n        torch.cuda.empty_cache()\n\n\n    print(\n        \"Finished:\",\n        name\n    )\n\n\n    return result\n\n\n# ============================================================\n# 8. Run all variants\n# ============================================================\n\nablation_results = pd.DataFrame(\n    [\n        quick_ablation_safe(\n            name,\n            ms,\n            peer,\n            dag,\n            seed=42 + i\n        )\n\n        for i, (\n            name,\n            ms,\n            peer,\n            dag\n        )\n        in enumerate(\n            variants\n        )\n    ]\n)\n\n\n# ============================================================\n# 9. Display results\n# ============================================================\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL ARCHITECTURAL ABLATION RESULTS\")\nprint(\"=\" * 70)\n\nprint(\n    ablation_results\n    .round(5)\n    .to_string(index=False)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:20:02.403403Z","iopub.execute_input":"2026-10-02T13:20:02.404171Z","iopub.status.idle":"2026-10-02T13:26:18.418888Z","shell.execute_reply.started":"2026-10-02T13:20:02.404137Z","shell.execute_reply":"2026-10-02T13:26:18.418164Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 15 — Lightweight recommendation explanation\ndef explain_target(idx=0, top_k=3):\n    model.eval()\n    with torch.no_grad():\n        _,_,hh=model(data.x_dict,data.edge_index_dict)\n        scores=[]\n        for st in range(0,hh[\"product\"].shape[0],512):\n            en=min(st+512,hh[\"product\"].shape[0])\n            s=model.score_all_products(hh[\"customer\"][idx:idx+1],hh[\"product\"],st,en).reshape(-1)\n            scores.append(s)\n        top=torch.topk(torch.cat(scores),top_k).indices.cpu().tolist()\n    print(\"Customer index:\",idx)\n    print(\"Top recommended article IDs:\",[articles_graph.iloc[i][\"article_id\"] for i in top])\n    return top\nexplain_target(0,3)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:26:24.447554Z","iopub.execute_input":"2026-10-02T13:26:24.448373Z","iopub.status.idle":"2026-10-02T13:26:24.70679Z","shell.execute_reply.started":"2026-10-02T13:26:24.448338Z","shell.execute_reply":"2026-10-02T13:26:24.706159Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 16 — FINAL PUBLICATION TABLES + HARD INTEGRITY CHECKS\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"PRIMARY RECOMMENDATION RESULTS\")\nprint(\"=\" * 70)\n\nprint(\n    comparison_results\n    .round(6)\n    .to_string(index=False)\n)\n\n\nprint(\"\\nSAMPLED-CANDIDATE STRESS TEST\")\nprint(\"=\" * 70)\n\nprint(\n    sampled_candidate_results\n    .round(6)\n    .to_string(index=False)\n)\n\n\nprint(\"\\nSECONDARY CLASSIFICATION RESULTS\")\nprint(\"=\" * 70)\n\nprint(\n    classification_results\n    .round(4)\n    .to_string(index=False)\n)\n\n\nprint(\"\\nABLATION RESULTS\")\nprint(\"=\" * 70)\n\nprint(\n    ablation_results\n    .round(5)\n    .to_string(index=False)\n)\n\n\n# ============================================================\n# HARD INTEGRITY CHECKS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"RUNNING FINAL INTEGRITY CHECKS\")\nprint(\"=\" * 70)\n\n\n# ------------------------------------------------------------\n# 1. Node-count integrity\n# ------------------------------------------------------------\n\nassert (\n    data[\"customer\"].num_nodes\n    ==\n    len(customers_df)\n), (\n    \"Customer node count does not match customers_df.\"\n)\n\n\nassert (\n    data[\"product\"].num_nodes\n    ==\n    len(articles_graph)\n), (\n    \"Product node count does not match articles_graph.\"\n)\n\n\n# ------------------------------------------------------------\n# 2. Train / validation / test split integrity\n# ------------------------------------------------------------\n\nassert (\n    int(data[\"customer\"].train_mask.sum())\n    ==\n    70000\n), \"Unexpected number of training customers.\"\n\n\nassert (\n    int(data[\"customer\"].val_mask.sum())\n    ==\n    15000\n), \"Unexpected number of validation customers.\"\n\n\nassert (\n    int(data[\"customer\"].test_mask.sum())\n    ==\n    15000\n), \"Unexpected number of test customers.\"\n\n\n# ------------------------------------------------------------\n# 3. Essential purchase graph check\n# ------------------------------------------------------------\n\npurchase_edge_type = (\n    \"customer\",\n    \"purchases\",\n    \"product\"\n)\n\n\nassert (\n    purchase_edge_type in data.edge_types\n), \"Purchase edge type is missing.\"\n\n\nassert (\n    data[purchase_edge_type]\n    .edge_index\n    .shape[1]\n    > 0\n), \"Purchase graph contains no edges.\"\n\n\n# ------------------------------------------------------------\n# 4. Compatibility DAG check\n# ------------------------------------------------------------\n\ncompat_edge_type = (\n    \"product\",\n    \"compat_to\",\n    \"product\"\n)\n\n\nassert (\n    compat_edge_type in data.edge_types\n), \"Compatibility DAG edge type is missing.\"\n\n\nassert (\n    data[compat_edge_type]\n    .edge_index\n    .shape[1]\n    > 0\n), \"Compatibility DAG contains no edges.\"\n\n\n# ------------------------------------------------------------\n# 5. Edge-index range integrity\n# ------------------------------------------------------------\n\nfor et in data.edge_types:\n\n    ei = data[et].edge_index\n\n    if ei.numel() == 0:\n        continue\n\n    src, _, dst = et\n\n\n    # Source indices\n    assert (\n        int(ei[0].min()) >= 0\n    ), (\n        f\"Negative source index found in {et}\"\n    )\n\n\n    assert (\n        int(ei[0].max())\n        <\n        data[src].num_nodes\n    ), (\n        f\"Source index out of range in {et}\"\n    )\n\n\n    # Destination indices\n    assert (\n        int(ei[1].min()) >= 0\n    ), (\n        f\"Negative destination index found in {et}\"\n    )\n\n\n    assert (\n        int(ei[1].max())\n        <\n        data[dst].num_nodes\n    ), (\n        f\"Destination index out of range in {et}\"\n    )\n\n\n# ------------------------------------------------------------\n# 6. Secondary classification model integrity\n# ------------------------------------------------------------\n\nexpected_classification_models = {\n    \"CC-HGNN\",\n    \"GraphSAGE\",\n    \"MLP\",\n    \"Logistic Regression\",\n    \"Random Forest\"\n}\n\n\nactual_classification_models = set(\n    classification_results[\"Model\"]\n    .astype(str)\n)\n\n\nassert (\n    actual_classification_models\n    ==\n    expected_classification_models\n), (\n    \"Classification model set does not match expected models.\\n\"\n    f\"Found: {actual_classification_models}\"\n)\n\n\n# ------------------------------------------------------------\n# 7. Recommendation method integrity\n# ------------------------------------------------------------\n\nexpected_methods = {\n    \"CC-HGNN\",\n    \"BPR-MF\",\n    \"Popularity\"\n}\n\n\nactual_methods = set(\n    comparison_results[\"Method\"]\n    .unique()\n)\n\n\nassert (\n    actual_methods\n    ==\n    expected_methods\n), (\n    \"Recommendation methods do not match expected methods.\\n\"\n    f\"Found: {actual_methods}\"\n)\n\n\n# ------------------------------------------------------------\n# 8. Three K values for each recommendation method\n# ------------------------------------------------------------\n\nmethod_counts = (\n    comparison_results\n    .groupby(\"Method\")\n    .size()\n    .to_dict()\n)\n\n\nassert (\n    method_counts\n    ==\n    {\n        \"CC-HGNN\": 3,\n        \"BPR-MF\": 3,\n        \"Popularity\": 3\n    }\n), (\n    \"Each recommendation method must have exactly \"\n    \"3 K values.\\n\"\n    f\"Found: {method_counts}\"\n)\n\n\n# ------------------------------------------------------------\n# 9. Recommendation K integrity\n# ------------------------------------------------------------\n\nassert (\n    recommendation_results[\"K\"].tolist()\n    ==\n    [3, 5, 10]\n), (\n    \"Recommendation results must contain K = [3, 5, 10].\"\n)\n\n\n# ------------------------------------------------------------\n# 10. Seed-results check\n# ------------------------------------------------------------\n#\n# IMPORTANT:\n# The current notebook does NOT generate genuine\n# three-seed experimental results.\n#\n# Therefore we do NOT fabricate seed metrics.\n#\n# If seed_results exists, validate it.\n# Otherwise simply record that no multi-seed result table\n# is available in the current run.\n# ------------------------------------------------------------\n\nif \"seed_results\" in globals():\n\n    assert (\n        len(seed_results) == 3\n    ), (\n        \"seed_results exists but does not contain 3 rows.\"\n    )\n\n\n    assert (\n        set(\n            seed_results[\"seed\"]\n            .astype(int)\n        )\n        ==\n        {42, 123, 2024}\n    ), (\n        \"seed_results does not contain the expected seeds \"\n        \"{42, 123, 2024}.\"\n    )\n\n\n    print(\n        \"Seed-results check: PASSED \"\n        \"(42, 123, 2024)\"\n    )\n\nelse:\n\n    print(\n        \"Seed-results check: SKIPPED \"\n        \"(no genuine multi-seed seed_results table was generated)\"\n    )\n\n\n# ------------------------------------------------------------\n# 11. Primary recommendation metric integrity\n# ------------------------------------------------------------\n\nrecommendation_metrics = [\n    \"Recall@K\",\n    \"HitRate@K\",\n    \"NDCG@K\"\n]\n\n\nassert np.isfinite(\n    comparison_results[\n        recommendation_metrics\n    ].to_numpy()\n).all(), (\n    \"comparison_results contains NaN or infinite values.\"\n)\n\n\n# ------------------------------------------------------------\n# 12. Sampled-candidate metric integrity\n# ------------------------------------------------------------\n\nassert np.isfinite(\n    sampled_candidate_results[\n        recommendation_metrics\n    ].to_numpy()\n).all(), (\n    \"sampled_candidate_results contains NaN or infinite values.\"\n)\n\n\n# ------------------------------------------------------------\n# 13. Metric range checks\n# ------------------------------------------------------------\n\nassert (\n    comparison_results[\n        recommendation_metrics\n    ].to_numpy()\n    >= 0\n).all(), (\n    \"Negative recommendation metric detected.\"\n)\n\n\nassert (\n    comparison_results[\n        recommendation_metrics\n    ].to_numpy()\n    <= 1\n).all(), (\n    \"Recommendation metric greater than 1 detected.\"\n)\n\n\nassert (\n    sampled_candidate_results[\n        recommendation_metrics\n    ].to_numpy()\n    >= 0\n).all(), (\n    \"Negative sampled-candidate metric detected.\"\n)\n\n\nassert (\n    sampled_candidate_results[\n        recommendation_metrics\n    ].to_numpy()\n    <= 1\n).all(), (\n    \"Sampled-candidate metric greater than 1 detected.\"\n)\n\n\n# ============================================================\n# FINAL STATUS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"ALL FINAL INTEGRITY CHECKS PASSED\")\nprint(\"=\" * 70)\n\nprint(\"Primary task: temporal next-item recommendation\")\nprint(\"Secondary task: retention classification\")\nprint(\"History cap: 30 unique products/customer\")\nprint(\"Candidate catalog: 10,000 historical-only products\")\nprint(\"No test metric used for model/seed selection: YES\")\nprint(\"No dense customer x product matrix: YES\")\n\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:26:27.880446Z","iopub.execute_input":"2026-10-02T13:26:27.88079Z","iopub.status.idle":"2026-10-02T13:26:27.920522Z","shell.execute_reply.started":"2026-10-02T13:26:27.880759Z","shell.execute_reply":"2026-10-02T13:26:27.919706Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 18 — FINAL publication-readiness gate\nprint(\"=\"*70);print(\"FINAL PUBLICATION-READINESS GATE\");print(\"=\"*70)\nfor metric in [\"Recall@K\",\"HitRate@K\",\"NDCG@K\"]:\n    print(\"\\n\",metric)\n    for k in [3,5,10]:\n        t=comparison_results[comparison_results.K==k].set_index(\"Method\");print(f\"K={k}: winner={t[metric].idxmax()} | CC-HGNN={t.loc['CC-HGNN',metric]:.6f} | BPR-MF={t.loc['BPR-MF',metric]:.6f} | Popularity={t.loc['Popularity',metric]:.6f}\")\nprint(\"\\nSampled-candidate stress-test winners:\")\nfor metric in [\"Recall@K\",\"HitRate@K\",\"NDCG@K\"]:\n    for k in [3,5,10]:\n        t=sampled_candidate_results[sampled_candidate_results.K==k].set_index(\"Method\");print(f\"{metric} K={k}: winner={t[metric].idxmax()}\")\nprint(\"\\nCC-HGNN superiority is NOT hard-coded. The measured test results determine the claim.\")\nprint(\"=\"*70)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-02T13:26:43.41253Z","iopub.execute_input":"2026-10-02T13:26:43.412872Z","iopub.status.idle":"2026-10-02T13:26:43.437073Z","shell.execute_reply.started":"2026-10-02T13:26:43.412841Z","shell.execute_reply":"2026-10-02T13:26:43.436222Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}