{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":31254,"databundleVersionId":3103714}],"dockerImageVersionId":31328,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"37592dd4","cell_type":"markdown","source":"# H&M 1st Stage Candidate Generation Demo\n\nThis lightweight notebook demonstrates the first stage of a two-stage recommender pipeline:\n\n1. Load a recent slice of the H&M transactions with Polars.\n2. Encode `customer_id` and `article_id` to integer keys.\n3. Generate candidates from history, global popularity, age-bucket popularity, and category popularity.\n4. Rank candidates with a weighted linear `candidate_score`.\n5. Measure Recall@k and MAP@k at `k = 12, 24, 48, 96, 120`.\n6. Plot recall customer mean and MAP by cutoff.\n","metadata":{}},{"id":"bb8214b0","cell_type":"code","source":"try:\n    import polars as pl\nexcept ImportError:\n    %pip -q install 'polars[gpu]'\n    import polars as pl\n\ntry:\n    import matplotlib.pyplot as plt\nexcept ImportError:\n    %pip -q install matplotlib\n    import matplotlib.pyplot as plt\n\nimport math\nfrom collections import defaultdict\nfrom datetime import date, timedelta\nfrom pathlib import Path\n\n# Polars GPU engine is used for LazyFrame materialization when available.\n# On CPU-only Kaggle sessions this falls back to the normal Polars engine.\nUSE_POLARS_GPU = True\n\n\ndef collect_pl(lf: pl.LazyFrame) -> pl.DataFrame:\n    global USE_POLARS_GPU\n    if USE_POLARS_GPU:\n        try:\n            return lf.collect(engine=\"gpu\")\n        except Exception as exc:\n            USE_POLARS_GPU = False\n            print(f\"Polars GPU unavailable; falling back to CPU. Reason: {type(exc).__name__}: {exc}\")\n    return lf.collect()\n\n\npl.__version__","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T14:26:26.296021Z","iopub.execute_input":"2026-05-14T14:26:26.296331Z","iopub.status.idle":"2026-05-14T14:26:26.992527Z","shell.execute_reply.started":"2026-05-14T14:26:26.296283Z","shell.execute_reply":"2026-05-14T14:26:26.991730Z"}},"outputs":[],"execution_count":null},{"id":"0bbd02fd","cell_type":"code","source":"# Kaggle path first, local repo path second.\nKAGGLE_DATA_DIR = Path(\"/kaggle/input/competitions/h-and-m-personalized-fashion-recommendations\")\nLOCAL_DATA_DIR = Path(\"../data/raw\")\nDATA_DIR = KAGGLE_DATA_DIR if KAGGLE_DATA_DIR.exists() else LOCAL_DATA_DIR\n\nTRANSACTIONS_FILE = DATA_DIR / \"transactions_train.csv\"\nCUSTOMERS_FILE = DATA_DIR / \"customers.csv\"\nARTICLES_FILE = DATA_DIR / \"articles.csv\"\n\nassert TRANSACTIONS_FILE.exists(), f\"Missing {TRANSACTIONS_FILE}\"\nassert CUSTOMERS_FILE.exists(), f\"Missing {CUSTOMERS_FILE}\"\nassert ARTICLES_FILE.exists(), f\"Missing {ARTICLES_FILE}\"\n\n# Teaching/demo parameters. Set MAX_VALID_CUSTOMERS=None for a fuller run.\nVALIDATION_CUTOFF = date(2020, 9, 8)\nLABEL_DAYS = 7\nEVAL_CUTOFFS = [12, 24, 48, 96, 120]\nMAX_VALID_CUSTOMERS = 5_000\n\nCANDIDATE_LIMIT = 120\nHISTORY_DAYS = 45\nHISTORY_TOP_N = 24\nGLOBAL_POPULARITY_DAYS = [7, 14, 30, 60]\nGLOBAL_TOP_N = 50\nAGE_POPULARITY_DAYS = 30\nAGE_TOP_N = 40\nCATEGORY_HISTORY_DAYS = 120\nCATEGORY_POPULARITY_DAYS = 45\nCATEGORY_TOP_N = 16\nCATEGORY_PREF_TOP_N = 2\nCATEGORY_COLS_RAW = [\"section_name\", \"garment_group_name\", \"index_group_name\"]\n\nSOURCE_WEIGHTS = {\n    \"history\": 4.0,\n    \"category\": 2.0,\n    \"global7\": 2.0,\n    \"global14\": 1.4,\n    \"global30\": 0.9,\n    \"global60\": 0.5,\n    \"age14\": 0.8,\n}\nRANK_WEIGHT = 2.0\nCOUNT_WEIGHT = 0.25\nRECENCY_WEIGHT = 3.0\n\nLABEL_END = VALIDATION_CUTOFF + timedelta(days=LABEL_DAYS)\nEARLIEST_NEEDED_DATE = VALIDATION_CUTOFF - timedelta(days=max(HISTORY_DAYS, max(GLOBAL_POPULARITY_DAYS), AGE_POPULARITY_DAYS, CATEGORY_HISTORY_DAYS, CATEGORY_POPULARITY_DAYS))\n\nDATA_DIR, EARLIEST_NEEDED_DATE, VALIDATION_CUTOFF, LABEL_END","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T14:28:05.302647Z","iopub.execute_input":"2026-05-14T14:28:05.303277Z","iopub.status.idle":"2026-05-14T14:28:05.314816Z","shell.execute_reply.started":"2026-05-14T14:28:05.303246Z","shell.execute_reply":"2026-05-14T14:28:05.314191Z"}},"outputs":[],"execution_count":null},{"id":"a0367db3","cell_type":"markdown","source":"## Load and encode a recent data slice\n\nFor the production repo we prebuild shared encoded parquet files. In this standalone Kaggle demo, we build a small in-notebook encoding from raw CSV so students can see the full flow.","metadata":{}},{"id":"1e6aff49","cell_type":"code","source":"transactions_raw = collect_pl(\n    pl.scan_csv(\n        TRANSACTIONS_FILE,\n        schema_overrides={\"t_dat\": pl.Utf8, \"customer_id\": pl.Utf8, \"article_id\": pl.Utf8},\n    )\n    .select([\"t_dat\", \"customer_id\", \"article_id\"])\n    .with_columns(pl.col(\"t_dat\").str.strptime(pl.Date, strict=False))\n    .filter((pl.col(\"t_dat\") > pl.lit(EARLIEST_NEEDED_DATE)) & (pl.col(\"t_dat\") <= pl.lit(LABEL_END)))\n)\n\ncustomers_raw = collect_pl(\n    pl.scan_csv(\n        CUSTOMERS_FILE,\n        schema_overrides={\"customer_id\": pl.Utf8, \"age\": pl.Float32},\n    ).select([\"customer_id\", \"age\"])\n)\n\narticles_raw = collect_pl(\n    pl.scan_csv(\n        ARTICLES_FILE,\n        schema_overrides={\"article_id\": pl.Utf8, **{col: pl.Utf8 for col in CATEGORY_COLS_RAW}},\n    ).select([\"article_id\", *CATEGORY_COLS_RAW])\n)\n\ncustomer_map = collect_pl(\n    pl.concat([transactions_raw.select(\"customer_id\"), customers_raw.select(\"customer_id\")])\n    .lazy()\n    .unique()\n    .sort(\"customer_id\")\n    .with_row_index(\"customer_idx\", offset=1)\n    .select([\"customer_idx\", \"customer_id\"])\n)\narticle_map = collect_pl(\n    pl.concat([transactions_raw.select(\"article_id\"), articles_raw.select(\"article_id\")])\n    .lazy()\n    .unique()\n    .sort(\"article_id\")\n    .with_row_index(\"article_idx\", offset=1)\n    .select([\"article_idx\", \"article_id\"])\n)\n\ntransactions = collect_pl(\n    transactions_raw\n    .lazy()\n    .join(customer_map.lazy(), on=\"customer_id\", how=\"left\")\n    .join(article_map.lazy(), on=\"article_id\", how=\"left\")\n    .select([\"t_dat\", \"customer_idx\", \"article_idx\"])\n)\ncustomers = collect_pl(\n    customers_raw\n    .lazy()\n    .join(customer_map.lazy(), on=\"customer_id\", how=\"left\")\n    .select([\"customer_idx\", \"age\"])\n)\narticles = collect_pl(articles_raw.lazy().join(article_map.lazy(), on=\"article_id\", how=\"left\").select([\"article_idx\", *CATEGORY_COLS_RAW]))\n\nfor col in CATEGORY_COLS_RAW:\n    map_df = collect_pl(articles.lazy().select(col).fill_null(\"__missing__\").unique().sort(col).with_row_index(f\"{col}_idx\", offset=1))\n    articles = collect_pl(articles.lazy().with_columns(pl.col(col).fill_null(\"__missing__\")).join(map_df.lazy(), on=col, how=\"left\").drop(col))\n\nCATEGORY_COLS = [f\"{col}_idx\" for col in CATEGORY_COLS_RAW]\n\ntransactions.shape, customers.shape, articles.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T14:28:08.115267Z","iopub.execute_input":"2026-05-14T14:28:08.116035Z","iopub.status.idle":"2026-05-14T14:29:01.332384Z","shell.execute_reply.started":"2026-05-14T14:28:08.116003Z","shell.execute_reply":"2026-05-14T14:29:01.331467Z"}},"outputs":[],"execution_count":null},{"id":"feb834dc","cell_type":"code","source":"def add_age_bucket(customers_df: pl.DataFrame) -> pl.DataFrame:\n    return collect_pl(customers_df.lazy().with_columns(\n        pl.when(pl.col(\"age\").is_null())\n        .then(pl.lit(\"unknown\"))\n        .when(pl.col(\"age\") < 20)\n        .then(pl.lit(\"<20\"))\n        .when(pl.col(\"age\") < 30)\n        .then(pl.lit(\"20s\"))\n        .when(pl.col(\"age\") < 40)\n        .then(pl.lit(\"30s\"))\n        .when(pl.col(\"age\") < 50)\n        .then(pl.lit(\"40s\"))\n        .when(pl.col(\"age\") < 60)\n        .then(pl.lit(\"50s\"))\n        .otherwise(pl.lit(\"60+\"))\n        .alias(\"age_bucket\")\n    ))\n\n\ndef sorted_count_frame(df: pl.DataFrame, group_cols: list[str]) -> pl.DataFrame:\n    return collect_pl(\n        df.lazy()\n        .group_by(group_cols)\n        .agg(pl.len().alias(\"cnt\"))\n        .sort(group_cols[:-1] + [\"cnt\", group_cols[-1]], descending=[False] * (len(group_cols) - 1) + [True, False])\n    )\n\n\ndef actuals_for_window(transactions_df: pl.DataFrame, cutoff: date, label_days: int) -> dict[int, list[int]]:\n    label_end = cutoff + timedelta(days=label_days)\n    grouped = collect_pl(\n        transactions_df\n        .lazy()\n        .filter((pl.col(\"t_dat\") > cutoff) & (pl.col(\"t_dat\") <= label_end))\n        .group_by(\"customer_idx\")\n        .agg(pl.col(\"article_idx\").unique().alias(\"actual\"))\n    )\n    return {int(row[\"customer_idx\"]): [int(x) for x in row[\"actual\"]] for row in grouped.iter_rows(named=True)}\n\n\ndef build_top_list(transactions_df: pl.DataFrame, cutoff: date, days: int, limit: int) -> list[tuple[int, int]]:\n    start = cutoff - timedelta(days=days)\n    counts = collect_pl(\n        transactions_df\n        .lazy()\n        .filter((pl.col(\"t_dat\") > start) & (pl.col(\"t_dat\") <= cutoff))\n        .group_by(\"article_idx\")\n        .agg(pl.len().alias(\"cnt\"))\n        .sort([\"cnt\", \"article_idx\"], descending=[True, False])\n        .head(limit)\n    )\n    return [(int(row[\"article_idx\"]), int(row[\"cnt\"])) for row in counts.iter_rows(named=True)]\n\n\ndef build_customer_history(transactions_df: pl.DataFrame, cutoff: date, days: int, top_n: int) -> dict[int, list[tuple[int, int, int]]]:\n    start = cutoff - timedelta(days=days)\n    latest = collect_pl(\n        transactions_df\n        .lazy()\n        .filter((pl.col(\"t_dat\") > start) & (pl.col(\"t_dat\") <= cutoff))\n        .group_by([\"customer_idx\", \"article_idx\"])\n        .agg(pl.len().alias(\"cnt\"), pl.max(\"t_dat\").alias(\"last_date\"))\n        .sort([\"customer_idx\", \"last_date\", \"article_idx\"], descending=[False, True, False])\n    )\n    history = defaultdict(list)\n    for row in latest.iter_rows(named=True):\n        customer_idx = int(row[\"customer_idx\"])\n        if len(history[customer_idx]) < top_n:\n            history[customer_idx].append((int(row[\"article_idx\"]), int(row[\"cnt\"]), (cutoff - row[\"last_date\"]).days))\n    return dict(history)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T14:29:01.334022Z","iopub.execute_input":"2026-05-14T14:29:01.334333Z","iopub.status.idle":"2026-05-14T14:29:01.348128Z","shell.execute_reply.started":"2026-05-14T14:29:01.334307Z","shell.execute_reply":"2026-05-14T14:29:01.347296Z"}},"outputs":[],"execution_count":null},{"id":"2875a062","cell_type":"code","source":"def build_age_top(transactions_df, customers_df, cutoff, days, top_n):\n    start = cutoff - timedelta(days=days)\n    recent = collect_pl(\n        transactions_df\n        .lazy()\n        .filter((pl.col(\"t_dat\") > start) & (pl.col(\"t_dat\") <= cutoff))\n        .select([\"customer_idx\", \"article_idx\"])\n        .join(customers_df.lazy().select([\"customer_idx\", \"age_bucket\"]), on=\"customer_idx\", how=\"left\")\n        .with_columns(pl.col(\"age_bucket\").fill_null(\"unknown\"))\n    )\n    counts = sorted_count_frame(recent, [\"age_bucket\", \"article_idx\"])\n    out = defaultdict(list)\n    for row in counts.iter_rows(named=True):\n        bucket = str(row[\"age_bucket\"])\n        if len(out[bucket]) < top_n:\n            out[bucket].append((int(row[\"article_idx\"]), int(row[\"cnt\"])))\n    return dict(out)\n\n\ndef build_category_top(transactions_df, articles_df, cutoff, days, category_cols, top_n):\n    start = cutoff - timedelta(days=days)\n    recent = collect_pl(\n        transactions_df\n        .lazy()\n        .filter((pl.col(\"t_dat\") > start) & (pl.col(\"t_dat\") <= cutoff))\n        .select([\"article_idx\"])\n        .join(articles_df.lazy(), on=\"article_idx\", how=\"left\")\n    )\n    category_top = {}\n    for col in category_cols:\n        counts = sorted_count_frame(collect_pl(recent.lazy().filter(pl.col(col).is_not_null()).select([col, \"article_idx\"])), [col, \"article_idx\"])\n        top_for_col = defaultdict(list)\n        for row in counts.iter_rows(named=True):\n            category = int(row[col])\n            if len(top_for_col[category]) < top_n:\n                top_for_col[category].append((int(row[\"article_idx\"]), int(row[\"cnt\"])))\n        category_top[col] = dict(top_for_col)\n    return category_top\n\n\ndef build_category_preferences(transactions_df, articles_df, cutoff, days, category_cols, pref_top_n):\n    start = cutoff - timedelta(days=days)\n    recent = collect_pl(\n        transactions_df\n        .lazy()\n        .filter((pl.col(\"t_dat\") > start) & (pl.col(\"t_dat\") <= cutoff))\n        .select([\"customer_idx\", \"article_idx\", \"t_dat\"])\n        .join(articles_df.lazy(), on=\"article_idx\", how=\"left\")\n    )\n    preferences = {}\n    for col in category_cols:\n        counts = collect_pl(\n            recent\n            .lazy()\n            .filter(pl.col(col).is_not_null())\n            .group_by([\"customer_idx\", col])\n            .agg(pl.len().alias(\"cnt\"), pl.max(\"t_dat\").alias(\"last_date\"))\n            .sort([\"customer_idx\", \"cnt\", \"last_date\", col], descending=[False, True, True, False])\n        )\n        pref_for_col = defaultdict(list)\n        for row in counts.iter_rows(named=True):\n            customer_idx = int(row[\"customer_idx\"])\n            if len(pref_for_col[customer_idx]) < pref_top_n:\n                pref_for_col[customer_idx].append(int(row[col]))\n        preferences[col] = dict(pref_for_col)\n    return preferences","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T14:29:10.868151Z","iopub.execute_input":"2026-05-14T14:29:10.868579Z","iopub.status.idle":"2026-05-14T14:29:10.880264Z","shell.execute_reply.started":"2026-05-14T14:29:10.868549Z","shell.execute_reply":"2026-05-14T14:29:10.879422Z"}},"outputs":[],"execution_count":null},{"id":"3bea12bf","cell_type":"code","source":"customers = add_age_bucket(customers)\nactuals_all = actuals_for_window(transactions, VALIDATION_CUTOFF, LABEL_DAYS)\nvalid_customers = sorted(actuals_all)\nif MAX_VALID_CUSTOMERS is not None:\n    valid_customers = valid_customers[:MAX_VALID_CUSTOMERS]\nactuals = {customer_idx: actuals_all[customer_idx] for customer_idx in valid_customers}\n\ncomponents = {\n    \"history\": build_customer_history(transactions, VALIDATION_CUTOFF, HISTORY_DAYS, HISTORY_TOP_N),\n    \"global_top\": {days: build_top_list(transactions, VALIDATION_CUTOFF, days, GLOBAL_TOP_N) for days in GLOBAL_POPULARITY_DAYS},\n    \"age_top\": build_age_top(transactions, customers, VALIDATION_CUTOFF, AGE_POPULARITY_DAYS, AGE_TOP_N),\n    \"category_top\": build_category_top(transactions, articles, VALIDATION_CUTOFF, CATEGORY_POPULARITY_DAYS, CATEGORY_COLS, CATEGORY_TOP_N),\n    \"category_preferences\": build_category_preferences(transactions, articles, VALIDATION_CUTOFF, CATEGORY_HISTORY_DAYS, CATEGORY_COLS, CATEGORY_PREF_TOP_N),\n    \"customer_age_bucket\": dict(zip(customers[\"customer_idx\"].to_list(), customers[\"age_bucket\"].fill_null(\"unknown\").to_list())),\n}\n\nlen(valid_customers), len(actuals_all)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T14:29:15.489638Z","iopub.execute_input":"2026-05-14T14:29:15.490353Z","iopub.status.idle":"2026-05-14T14:29:29.425198Z","shell.execute_reply.started":"2026-05-14T14:29:15.490319Z","shell.execute_reply":"2026-05-14T14:29:29.424542Z"}},"outputs":[],"execution_count":null},{"id":"9b23a4e5","cell_type":"code","source":"SOURCE_COLUMNS = [\"source_history\", \"source_age14\", \"source_global7\", \"source_global14\", \"source_global30\", \"source_global60\", \"source_category\"]\nRANK_COLUMNS = [\"history_rank\", \"age_rank\", \"global7_rank\", \"global14_rank\", \"global30_rank\", \"global60_rank\", \"category_rank\"]\nMISSING_RANK = 999\n\n\ndef new_candidate_features():\n    return {\n        \"candidate_score\": 0.0,\n        **{col: 0 for col in SOURCE_COLUMNS},\n        **{col: MISSING_RANK for col in RANK_COLUMNS},\n    }\n\n\ndef score_piece(source, rank, count, recency_days=None):\n    score = SOURCE_WEIGHTS[source] + RANK_WEIGHT / max(rank, 1) + COUNT_WEIGHT * math.log1p(max(count, 0))\n    if recency_days is not None:\n        score += RECENCY_WEIGHT / (1.0 + max(recency_days, 0))\n    return score\n\n\ndef add_candidate(candidates, article_idx, source, rank, count, recency_days=None):\n    if article_idx not in candidates:\n        if len(candidates) >= CANDIDATE_LIMIT:\n            return\n        candidates[article_idx] = new_candidate_features()\n    row = candidates[article_idx]\n    row[f\"source_{source}\"] = 1\n    rank_col = \"age_rank\" if source == \"age14\" else f\"{source}_rank\"\n    row[rank_col] = min(row[rank_col], rank)\n    row[\"candidate_score\"] += score_piece(source, rank, count, recency_days)\n\n\ndef candidate_map_for_customer(customer_idx):\n    candidates = {}\n    for rank, (article_idx, count, recency_days) in enumerate(components[\"history\"].get(customer_idx, []), start=1):\n        add_candidate(candidates, article_idx, \"history\", rank, count, recency_days)\n    bucket = components[\"customer_age_bucket\"].get(customer_idx, \"unknown\")\n    for rank, (article_idx, count) in enumerate(components[\"age_top\"].get(bucket, []), start=1):\n        add_candidate(candidates, article_idx, \"age14\", rank, count)\n    for rank, (article_idx, count) in enumerate(components[\"global_top\"].get(7, []), start=1):\n        add_candidate(candidates, article_idx, \"global7\", rank, count)\n    for col in CATEGORY_COLS:\n        for category in components[\"category_preferences\"][col].get(customer_idx, []):\n            for rank, (article_idx, count) in enumerate(components[\"category_top\"][col].get(category, []), start=1):\n                add_candidate(candidates, article_idx, \"category\", rank, count)\n    for days, source in [(14, \"global14\"), (30, \"global30\"), (60, \"global60\")]:\n        for rank, (article_idx, count) in enumerate(components[\"global_top\"].get(days, []), start=1):\n            add_candidate(candidates, article_idx, source, rank, count)\n    return candidates\n\n\ndef ordered_candidates(candidates):\n    return sorted(candidates.items(), key=lambda item: (-item[1][\"candidate_score\"], min(item[1][col] for col in RANK_COLUMNS), item[0]))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T14:29:29.426293Z","iopub.execute_input":"2026-05-14T14:29:29.426607Z","iopub.status.idle":"2026-05-14T14:29:29.453763Z","shell.execute_reply.started":"2026-05-14T14:29:29.426584Z","shell.execute_reply":"2026-05-14T14:29:29.452751Z"}},"outputs":[],"execution_count":null},{"id":"8cdb02ab","cell_type":"code","source":"def apk(actual_set, predicted, k):\n    if not actual_set:\n        return 0.0\n    score = 0.0\n    hits = 0\n    seen = set()\n    for rank, article_idx in enumerate(predicted[:k], start=1):\n        if article_idx in seen:\n            continue\n        seen.add(article_idx)\n        if article_idx in actual_set:\n            hits += 1\n            score += hits / rank\n    return score / min(len(actual_set), k)\n\n\nmetric_acc = {k: {\"customer_recall\": [], \"map\": []} for k in EVAL_CUTOFFS}\ncandidate_rows = []\n\nfor customer_idx in valid_customers:\n    candidates = candidate_map_for_customer(customer_idx)\n    ordered_items = ordered_candidates(candidates)\n    prediction = [article_idx for article_idx, _ in ordered_items]\n    actual_set = set(actuals[customer_idx])\n\n    for k in EVAL_CUTOFFS:\n        hits = len(set(prediction[:k]) & actual_set)\n        metric_acc[k][\"customer_recall\"].append(hits / len(actual_set) if actual_set else 0.0)\n        metric_acc[k][\"map\"].append(apk(actual_set, prediction, k))\n\n    for order, (article_idx, features) in enumerate(ordered_items, start=1):\n        candidate_rows.append({\n            \"customer_idx\": customer_idx,\n            \"article_idx\": article_idx,\n            \"candidate_order\": order,\n            \"candidate_score\": features[\"candidate_score\"],\n            \"label\": int(article_idx in actual_set),\n        })\n\nmetrics = pl.DataFrame([\n    {\n        \"k\": k,\n        \"recall_customer_mean\": sum(metric_acc[k][\"customer_recall\"]) / len(metric_acc[k][\"customer_recall\"]),\n        \"map\": sum(metric_acc[k][\"map\"]) / len(metric_acc[k][\"map\"]),\n    }\n    for k in EVAL_CUTOFFS\n])\ncandidate_df = pl.DataFrame(candidate_rows)\n\nmetrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T14:29:41.049786Z","iopub.execute_input":"2026-05-14T14:29:41.050077Z","iopub.status.idle":"2026-05-14T14:29:45.352972Z","shell.execute_reply.started":"2026-05-14T14:29:41.050053Z","shell.execute_reply":"2026-05-14T14:29:45.352286Z"}},"outputs":[],"execution_count":null},{"id":"50e0f814","cell_type":"code","source":"ks = metrics[\"k\"].to_list()\nrecall_values = metrics[\"recall_customer_mean\"].to_list()\nmap_values = metrics[\"map\"].to_list()\n\nfig, ax1 = plt.subplots(figsize=(8, 4.5))\nax1.plot(ks, recall_values, marker=\"o\", linewidth=2, label=\"Recall customer mean\", color=\"tab:blue\")\nax1.set_xlabel(\"k\")\nax1.set_ylabel(\"Recall customer mean\", color=\"tab:blue\")\nax1.tick_params(axis=\"y\", labelcolor=\"tab:blue\")\nax1.set_xticks(ks)\nax1.grid(True, alpha=0.25)\n\nax2 = ax1.twinx()\nax2.plot(ks, map_values, marker=\"s\", linewidth=2, label=\"MAP\", color=\"tab:orange\")\nax2.set_ylabel(\"MAP\", color=\"tab:orange\")\nax2.tick_params(axis=\"y\", labelcolor=\"tab:orange\")\n\nfig.suptitle(\"1st-stage candidate quality by cutoff\")\nfig.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T14:29:57.354145Z","iopub.execute_input":"2026-05-14T14:29:57.354584Z","iopub.status.idle":"2026-05-14T14:29:57.589876Z","shell.execute_reply.started":"2026-05-14T14:29:57.354553Z","shell.execute_reply":"2026-05-14T14:29:57.588963Z"}},"outputs":[],"execution_count":null},{"id":"50e0f814","cell_type":"code","source":"candidate_df.shape, candidate_df.head(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T14:31:00.845732Z","iopub.execute_input":"2026-05-14T14:31:00.846609Z","iopub.status.idle":"2026-05-14T14:31:00.852433Z","shell.execute_reply.started":"2026-05-14T14:31:00.846574Z","shell.execute_reply":"2026-05-14T14:31:00.851758Z"}},"outputs":[],"execution_count":null},{"id":"11afab37","cell_type":"markdown","source":"## What students should notice\n\n- Recall rises as `k` increases because the candidate set is wider.\n- MAP is sensitive to ordering, so `candidate_score` matters even before the 2nd-stage reranker.\n- The 1st stage does not need to be a heavy model. A transparent weighted linear score is enough to teach the candidate-generation idea.\n- The 2nd stage should receive this fixed candidate set and only rerank it.","metadata":{}}]}