{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.17","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpuV5e8","dataSources":[{"sourceId":105399,"databundleVersionId":12733338,"sourceType":"competition"}],"dockerImageVersionId":31042,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false},"colab":{"provenance":[],"gpuType":"T4"},"accelerator":"GPU"},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# %%capture\n!pip install --upgrade pip\n!pip install polars","metadata":{"_uuid":"3f03a46b-3930-4a76-a8b8-f2bac10535ab","_cell_guid":"539292b0-37ce-483a-952e-7186f4d8c507","trusted":true,"collapsed":false,"id":"dE63opamDxtJ","outputId":"ac1d1165-8e97-45be-ed5a-13d99468e50c","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:40:03.284656Z","iopub.execute_input":"2025-12-04T00:40:03.284953Z","iopub.status.idle":"2025-12-04T00:40:08.384948Z","shell.execute_reply.started":"2025-12-04T00:40:03.284935Z","shell.execute_reply":"2025-12-04T00:40:08.380314Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import kagglehub\nimport os\n\npath = kagglehub.competition_download(\"aeroclub-recsys-2025\")\nprint(\"✅ Dataset downloaded to:\", path)","metadata":{"_uuid":"73162f5f-34aa-4e82-928d-b2b9ccfe17bc","_cell_guid":"6e755549-de95-4fb3-a249-bcfc5bf72c6e","trusted":true,"collapsed":false,"id":"V1-piopMEHHz","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:40:08.386613Z","iopub.execute_input":"2025-12-04T00:40:08.386781Z","iopub.status.idle":"2025-12-04T00:40:08.903055Z","shell.execute_reply.started":"2025-12-04T00:40:08.386764Z","shell.execute_reply":"2025-12-04T00:40:08.899576Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import polars as pl\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport time\n\nRANDOM_STATE = 42\nnp.random.seed(RANDOM_STATE)","metadata":{"_uuid":"b6b3b072-ca04-4edb-bfb8-a44787e1770b","_cell_guid":"1f02085b-5d49-4ca5-a667-d1d0dbfe8730","trusted":true,"collapsed":false,"id":"gNNQlxrRDxtL","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:40:08.904341Z","iopub.execute_input":"2025-12-04T00:40:08.904512Z","iopub.status.idle":"2025-12-04T00:40:08.911585Z","shell.execute_reply.started":"2025-12-04T00:40:08.904497Z","shell.execute_reply":"2025-12-04T00:40:08.908072Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import polars as pl\nimport numpy as np\nimport time\n\nprint(\"=\" * 60)\nprint(\"Starting Simplified Feature Engineering Pipeline for DL\")\nprint(\"=\" * 60)\n\n# Load data\nprint(\"\\n[1/8] Loading data...\")\nstart_time = time.time()\ntrain = pl.read_parquet(f'{path}/train.parquet').drop('__index_level_0__')\ntest = pl.read_parquet(f'{path}/test.parquet').drop('__index_level_0__').with_columns(pl.lit(0, dtype=pl.Int64).alias(\"selected\"))\n\ntest_start_idx = len(train)\n\ndata_raw = pl.concat((train, test))\n\nprint(f\"   ✓ Loaded {len(data_raw)} rows\")\nprint(f\"   ✓ Total dataset: {len(data_raw)} rows, {len(data_raw.columns)} columns\")\nprint(f\"   Time: {time.time() - start_time:.2f}s\")\n\ndf = data_raw.clone()","metadata":{"_uuid":"ccfdd54d-3f65-4c60-aa81-b68a41352c8f","_cell_guid":"e41a607e-1418-4ae1-8e61-1d7f41fe834f","trusted":true,"collapsed":false,"id":"3hKhZsdmDxtL","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:40:08.912051Z","iopub.execute_input":"2025-12-04T00:40:08.912934Z","iopub.status.idle":"2025-12-04T00:40:10.711452Z","shell.execute_reply.started":"2025-12-04T00:40:08.912918Z","shell.execute_reply":"2025-12-04T00:40:10.707576Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Helpers","metadata":{"_uuid":"683a7b8d-ce95-4545-b287-9c9bb0a12fde","_cell_guid":"7d1ae6d0-73ed-4247-959d-6796ae8cb44f","trusted":true,"collapsed":false,"id":"hk1nYWaPDxtM","jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def hitrate_at_3(y_true, y_pred, groups):\n    df = pl.DataFrame({\n        'group': groups,\n        'pred': y_pred,\n        'true': y_true\n    })\n\n    return (\n        df.filter(pl.col(\"group\").count().over(\"group\") > 10)\n        .sort([\"group\", \"pred\"], descending=[False, True])\n        .group_by(\"group\", maintain_order=True)\n        .head(3)\n        .group_by(\"group\")\n        .agg(pl.col(\"true\").max())\n        .select(pl.col(\"true\").mean())\n        .item()\n    )","metadata":{"_uuid":"4d4633b9-6352-41fd-b885-0905330ab689","_cell_guid":"387435ae-bdc0-4faa-8b22-09bce104ce7a","trusted":true,"collapsed":false,"id":"zZoKursCDxtM","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:40:10.713392Z","iopub.execute_input":"2025-12-04T00:40:10.713570Z","iopub.status.idle":"2025-12-04T00:40:10.722881Z","shell.execute_reply.started":"2025-12-04T00:40:10.713552Z","shell.execute_reply":"2025-12-04T00:40:10.718652Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Feature Engineering","metadata":{"_uuid":"08214025-0128-46e2-aba2-bf3cd6d6dcc9","_cell_guid":"f543a311-5e5c-4941-8995-4445bf473cb2","trusted":true,"collapsed":false,"id":"rlK8kru-DxtN","jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def dur_to_hour(col: pl.Expr) -> pl.Expr:\n    # Always treat input as string during processing\n    col_str = col.cast(pl.Utf8)\n\n    # Extract days part (e.g., \"1.\" from \"1.10:30:00\")\n    days = (\n        col_str\n        .str.extract(r\"^(\\d+)\\.\", 1)\n        .cast(pl.Float64)\n        .fill_null(0) * 24\n    )\n\n    # Remove \"X.\" prefix\n    time_str = col_str.str.replace(r\"^\\d+\\.\", \"\")\n\n    # Extract hours/minutes\n    hours = (\n        time_str.str.extract(r\"(\\d+):\", 1)\n        .cast(pl.Float64)\n        .fill_null(0)\n    )\n\n    minutes = (\n        time_str.str.extract(r\":(\\d+):\", 1)\n        .cast(pl.Float64)\n        .fill_null(0) / 60\n    )\n\n    return (days + hours + minutes).fill_null(0)\n# Process duration columns\nprint(\"\\n[2/8] Processing duration columns...\")\nstart_time = time.time()\ndur_cols = [\"legs0_duration\", \"legs1_duration\"] + [f\"legs{l}_segments{s}_duration\" for l in (0, 1) for s in (0, 1, 2, 3)]\ndur_exprs = [dur_to_hour(pl.col(c)).alias(c) for c in dur_cols if c in df.columns]\n\nif dur_exprs:\n    df = df.with_columns(dur_exprs)\n    print(f\"   ✓ Converted {len(dur_exprs)} duration columns to minutes\")\nprint(f\"   Time: {time.time() - start_time:.2f}s\")\n\n# === CORE NUMERICAL FEATURES ===\nprint(\"\\n[3/8] Creating core numerical features...\")\nstart_time = time.time()\n\n# Get all segment columns for aggregation\nall_seg_cols = [c for c in df.columns if '_segments' in c]\ncabin_cols = [c for c in all_seg_cols if c.endswith('_cabinClass')]\nbaggage_cols = [c for c in all_seg_cols if 'baggageAllowance_quantity' in c]\nduration_seg_cols = [c for c in all_seg_cols if c.endswith('_duration')]\n\n\n\ndf = df.with_columns([\n    # Price features\n    (pl.col(\"totalPrice\") / (pl.col(\"taxes\") + 1)).alias(\"price_per_tax\"),\n    (pl.col(\"taxes\") * 100 / (pl.col(\"totalPrice\") + 1)).alias(\"tax_ratex100\"),\n    pl.col(\"totalPrice\").log1p().alias(\"log_price\"),\n\n    # Duration features\n    (pl.col(\"legs0_duration\").fill_null(0) + pl.col(\"legs1_duration\").fill_null(0)).alias(\"total_duration\"),\n    pl.when(pl.col(\"legs1_duration\").fill_null(0) > 0)\n        .then(pl.col(\"legs0_duration\") / (pl.col(\"legs1_duration\") + 0.01))\n        .otherwise(1.0).alias(\"duration_ratio\"),\n\n    # Baggage - aggregate ALL segments\n    (pl.mean_horizontal([pl.col(c).cast(pl.Float64).fill_null(0) for c in baggage_cols]) if baggage_cols else pl.lit(0)).alias(\"baggage_mean\"),\n\n    # Fees\n    (pl.col(\"miniRules0_monetaryAmount\").fill_null(0) +\n     pl.col(\"miniRules1_monetaryAmount\").fill_null(0)).alias(\"total_fees\"),\n\n    # Cabin class - average across ALL segments (not just first)\n    (pl.mean_horizontal([pl.col(c).cast(pl.Float64).fill_null(0) for c in cabin_cols]) if cabin_cols else pl.lit(0)).alias(\"avg_cabin_class_all\"),\n\n    # Cabin class difference between legs (if round trip)\n    pl.when(pl.col(\"legs1_duration\").is_not_null())\n        .then(\n            pl.mean_horizontal([pl.col(c).cast(pl.Float64).fill_null(0) for c in cabin_cols if 'legs0_' in c]) -\n            pl.mean_horizontal([pl.col(c).cast(pl.Float64).fill_null(0) for c in cabin_cols if 'legs1_' in c])\n        )\n        .otherwise(0.0).alias(\"cabin_class_diff_legs\"),\n])\n\n# We use .over(\"ranker_id\") to calculate stats per search query\ndf = df.with_columns([\n    # 1. Calculate the minimum and mean duration for this specific search (ranker_id)\n    pl.col(\"total_duration\").min().over(\"ranker_id\").alias(\"min_dur_group\"),\n    pl.col(\"total_duration\").mean().over(\"ranker_id\").alias(\"mean_dur_group\"),\n    pl.col(\"total_duration\").std().over(\"ranker_id\").fill_null(1).alias(\"std_dur_group\"),\n])\n\ndf = df.with_columns([\n    # === APPROACH C: Z-Score (Standardized) ===\n    # Positive values = Faster than average (Good)\n    # Negative values = Slower than average (Bad)\n    ((pl.col(\"mean_dur_group\") - pl.col(\"total_duration\")) / (pl.col(\"std_dur_group\") + 1e-6)).alias(\"duration_z_score\")\n])\n\n# Drop the temporary aggregate columns\ndf = df.drop([\"min_dur_group\", \"mean_dur_group\", \"std_dur_group\"])\n\n# Segment counting\nmc_cols = [f'legs{l}_segments{s}_marketingCarrier_code' for l in (0, 1) for s in range(4)]\nmc_exists = [col for col in mc_cols if col in df.columns]\n\ndf = df.with_columns([\n    (pl.sum_horizontal([pl.col(col).is_not_null().cast(pl.UInt8) for col in mc_exists])\n     if mc_exists else pl.lit(0)).alias(\"total_segments\"),\n])\n\n# Derived fee features\ndf = df.with_columns([\n    (pl.col(\"total_fees\") / (pl.col(\"totalPrice\") + 1)).alias(\"fee_rate\"),\n])\n\nprint(f\"   ✓ Created aggregated price, duration, baggage, and cabin features\")\nprint(f\"   Time: {time.time() - start_time:.2f}s\")\n\n# === BINARY FEATURES ===\nprint(\"\\n[4/8] Creating binary features...\")\nstart_time = time.time()\n\ndf = df.with_columns([\n    # Trip type\n    (pl.col(\"legs1_duration\").is_null() |\n     (pl.col(\"legs1_duration\") == 0) |\n     pl.col(\"legs1_segments0_departureFrom_airport_iata\").is_null()).cast(pl.Int32).alias(\"is_one_way\"),\n\n    # Direct flights\n    pl.when(pl.col(\"legs1_duration\").is_not_null())\n        .then((pl.sum_horizontal([pl.col(c).is_not_null() for c in mc_exists if 'legs1_' in c]) == 1).cast(pl.Int32))\n        .otherwise(0).alias(\"is_direct_leg1\"),\n\n    # Corporate & VIP\n    pl.col(\"corporateTariffCode\").is_not_null().cast(pl.Int32).alias(\"has_corporate_tariff\"),\n    (pl.col(\"pricingInfo_isAccessTP\") == 1).cast(pl.Int32).alias(\"has_access_tp\"),\n    ((pl.col(\"isVip\") == 1) | (pl.col(\"frequentFlyer\").fill_null(\"\") != \"\")).cast(pl.Int32).alias(\"is_vip_freq\"),\n\n    # Baggage & fees\n    (pl.col(\"total_fees\") > 0).cast(pl.Int32).alias(\"has_fees\"),\n\n\n    # Major carriers\n    (pl.col(\"legs0_segments0_marketingCarrier_code\").is_in([\"SU\", \"S7\", \"U6\"]) if \"legs0_segments0_marketingCarrier_code\" in df.columns\n     else pl.lit(False)).cast(pl.Int32).alias(\"is_major_carrier\"),\n])\n\n\nprint(f\"   ✓ Created binary trip type, VIP, baggage, and carrier features\")\nprint(f\"   Time: {time.time() - start_time:.2f}s\")\n\n# === DATETIME FEATURES WITH CYCLIC ENCODING ===\nprint(\"\\n[5/8] Processing datetime features with cyclic encoding...\")\nstart_time = time.time()\n\n# Cyclic encoding for hour and weekday\ntime_exprs = []\nfor col in (\"legs0_departureAt\", \"legs0_arrivalAt\", \"legs1_departureAt\", \"legs1_arrivalAt\"):\n    if col in df.columns:\n        dt = pl.col(col).str.to_datetime(strict=False)\n        h = dt.dt.hour().fill_null(12)\n        wd = dt.dt.weekday().fill_null(0)\n\n        # Sin/cos for hour (24-hour cycle)\n        time_exprs.extend([\n            (np.sin(2 * np.pi * h / 24)).alias(f\"{col}_hour_sin\"),\n            (np.cos(2 * np.pi * h / 24)).alias(f\"{col}_hour_cos\"),\n            # Sin/cos for weekday (7-day cycle)\n            (np.sin(2 * np.pi * wd / 7)).alias(f\"{col}_weekday_sin\"),\n            (np.cos(2 * np.pi * wd / 7)).alias(f\"{col}_weekday_cos\"),\n            # Business time flag\n            (((h >= 6) & (h <= 9)) | ((h >= 17) & (h <= 20))).cast(pl.Int32).alias(f\"{col}_business_time\")\n        ])\n\nif time_exprs:\n    df = df.with_columns(time_exprs)\n    print(f\"   ✓ Created cyclic datetime features (sin/cos encoding)\")\n\nprint(f\"   Time: {time.time() - start_time:.2f}s\")\n\n# === TRUST VALUE (FREQUENT FLYER ALIGNMENT) ===\nprint(\"\\n[6/8] Computing trust/alignment features...\")\nstart_time = time.time()\n\n# Extract frequent flyer carrier codes\ndf = df.with_columns([\n    pl.col(\"frequentFlyer\").fill_null(\"\").str.split(\"/\").alias(\"ff_carriers_list\"),\n])\n\n# Count matching carriers in segments\ncarrier_cols = [c for c in mc_exists]\ndf = df.with_columns([\n    pl.sum_horizontal([\n        pl.col(c).is_in(pl.col(\"ff_carriers_list\")).cast(pl.Int32)\n        for c in carrier_cols\n    ]).alias(\"ff_matches\")\n])\n\n# Trust value = matching carriers / total segments\ndf = df.with_columns([\n    (pl.col(\"ff_matches\") / (pl.col(\"total_segments\") + 1)).alias(\"trust_value\"),\n])\n\ndf = df.drop([\"ff_carriers_list\", \"ff_matches\"])\nprint(f\"   ✓ Created trust value based on frequent flyer alignment\")\nprint(f\"   Time: {time.time() - start_time:.2f}s\")\n\n# === CATEGORICAL FEATURES FOR EMBEDDING ===\nprint(\"\\n[7/8] Extracting categorical features for embeddings...\")\nstart_time = time.time()\n\n# First and last airports\ndf = df.with_columns([\n    pl.col(\"legs0_segments0_departureFrom_airport_iata\").fill_null(\"UNK\").alias(\"first_departure_airport\"),\n\n    # Last arrival - check all possible segments\n    pl.coalesce([\n        pl.col(f\"legs1_segments{i}_arrivalTo_airport_iata\")\n        for i in range(3, -1, -1)\n    ] + [\n        pl.col(f\"legs0_segments{i}_arrivalTo_airport_iata\")\n        for i in range(3, -1, -1)\n    ]).fill_null(\"UNK\").alias(\"last_arrival_airport\"),\n])\n\n# Collect all carriers used (as comma-separated string for embedding)\ncarrier_list_expr = pl.concat_str([\n    pl.col(c).fill_null(\"\") for c in carrier_cols\n], separator=\",\").str.replace_all(r\",+\", \",\").str.strip_chars(\",\")\n\ndf = df.with_columns([\n    carrier_list_expr.alias(\"carriers_used\"),\n])\n\nprint(f\"   ✓ Created categorical features: airports and carriers for embedding\")\nprint(f\"   Time: {time.time() - start_time:.2f}s\")\n\n# === HISTORICAL EMBEDDING ====\n# === CARRIER POPULARITY (FROM TRAINING DATA) ===\nprint(\"\\n[8/8] Computing carrier popularity features...\")\nstart_time = time.time()\n\n# Aggregate carrier popularity across ALL segments (not just first)\nall_carrier_pops = []\nfor l in (0, 1):\n    for s in range(4):\n        col_name = f'legs{l}_segments{s}_marketingCarrier_code'\n        if col_name in df.columns:\n            # carrier_pop = train.group_by(col_name).agg(\n            carrier_pop = df.group_by(col_name).agg(\n                pl.mean('selected').alias(f'carrier_pop_{l}_{s}')\n            )\n            df = df.join(carrier_pop, on=col_name, how='left')\n            df = df.with_columns(pl.col(f'carrier_pop_{l}_{s}').fill_null(0.0))\n            all_carrier_pops.append(f'carrier_pop_{l}_{s}')\n\n# Average carrier popularity across all segments\nif all_carrier_pops:\n    df = df.with_columns([\n        pl.mean_horizontal([pl.col(c) for c in all_carrier_pops]).alias(\"avg_carrier_popularity\")\n    ])\n    # Clean up individual carrier pop columns\n    df = df.drop(all_carrier_pops)\nelse:\n    df = df.with_columns(pl.lit(0.0).alias(\"avg_carrier_popularity\"))\n\nprint(f\"   ✓ Created aggregated carrier popularity from training data\")\nprint(f\"   Time: {time.time() - start_time:.2f}s\")\n\n# === FINAL CLEANUP: DROP RAW SEGMENT COLUMNS ===\nprint(\"\\nDropping raw segment-level columns...\")\nstart_time = time.time()\n\n# Identify columns to drop (all segment-level detail columns)\ncols_to_drop = [c for c in df.columns if any([\n    '_segments' in c and c.endswith('_code'),  # aircraft codes, carrier codes\n    '_segments' in c and 'flightNumber' in c,\n    '_segments' in c and 'seatsAvailable' in c,\n    '_segments' in c and 'baggageAllowance_weightMeasurementType' in c,\n    '_segments' in c and 'airport_city_iata' in c,\n    '_segments' in c and '_duration' in c,  # Already aggregated\n    '_segments' in c and 'baggageAllowance_quantity' in c,  # Already aggregated\n    '_segments' in c and 'departureFrom_airport_iata' in c and c != 'legs0_segments0_departureFrom_airport_iata',\n    '_segments' in c and 'arrivalTo_airport_iata' in c,\n])]\n\n# Also drop legs-level durations (already have total_duration)\ncols_to_drop.extend(['legs0_duration', 'legs1_duration'])\n\ndf = df.drop([c for c in cols_to_drop if c in df.columns])\nprint(f\"   ✓ Dropped {len([c for c in cols_to_drop if c in df.columns])} redundant segment-level columns\")\nprint(f\"   Time: {time.time() - start_time:.2f}s\")\n\n\n# === FINAL CLEANUP: DROP RAW SEGMENT COLUMNS ===\nprint(\"\\nDropping raw segment-level columns...\")\nstart_time = time.time()\n\n# Identify columns to drop (all segment-level detail columns)\ncols_to_drop = [c for c in df.columns if any([\n    '_segments' in c and c.endswith('_code'),  # aircraft codes, carrier codes\n    '_segments' in c and 'flightNumber' in c,\n    '_segments' in c and 'seatsAvailable' in c,\n    '_segments' in c and 'baggageAllowance_weightMeasurementType' in c,\n    '_segments' in c and 'airport_city_iata' in c,\n    '_segments' in c and '_duration' in c,  # Already aggregated\n    '_segments' in c and 'baggageAllowance_quantity' in c,  # Already aggregated\n    '_segments' in c and 'departureFrom_airport_iata' in c and c != 'legs0_segments0_departureFrom_airport_iata',\n    '_segments' in c and 'arrivalTo_airport_iata' in c,\n])]\n\n# Also drop legs-level durations (already have total_duration)\ncols_to_drop.extend(['legs0_duration', 'legs1_duration'])\n\ndf = df.drop([c for c in cols_to_drop if c in df.columns])\nprint(f\"   ✓ Dropped {len([c for c in cols_to_drop if c in df.columns])} redundant segment-level columns\")\nprint(f\"   Time: {time.time() - start_time:.2f}s\")","metadata":{"_uuid":"b695fb21-f39b-499e-a441-f6235159917a","_cell_guid":"4f72ad0d-9d47-4bd2-bf4e-b79379e12c6a","trusted":true,"collapsed":false,"id":"pkrlVlijDxtN","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:40:10.724837Z","iopub.execute_input":"2025-12-04T00:40:10.725021Z","iopub.status.idle":"2025-12-04T00:41:23.915647Z","shell.execute_reply.started":"2025-12-04T00:40:10.725004Z","shell.execute_reply":"2025-12-04T00:41:23.910986Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Feature Selection","metadata":{"_uuid":"9b3d8e25-106c-4039-8490-c55a844cf3ee","_cell_guid":"1058f4d0-cbcc-4965-b824-5113b77d2041","trusted":true,"collapsed":false,"id":"NmcPTUOvDxtO","jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# === CREATE CLEAN FEATURE DATAFRAME ===\nprint(\"\\nCreating clean feature-only DataFrame...\")\nstart_time = time.time()\n\ngroup_feature = [\n    \"Id\",\n    \"ranker_id\",\n    \"selected\",\n]\n\n# Define the core input features (X)\nfeature_columns = [\n    # === PERSONAL/USER DATA ===\n\n    # === NUMERICAL FEATURES ===\n    'log_price',\n    'tax_ratex100',\n    'fee_rate',\n    'duration_z_score',\n    'duration_ratio',\n    'baggage_mean',\n    'avg_cabin_class_all',\n    'cabin_class_diff_legs',\n    'total_segments',\n    'trust_value',\n    'avg_carrier_popularityx10',\n\n    # === BINARY FEATURES ===\n    # 'is_one_way',\n    'is_vip_freq', # Or use 'is_vip_freq' if you prefer the engineered version created in step 4\n    'has_corporate_tariff',\n    'has_access_tp',\n\n    # === CATEGORICAL FEATURES (Uncomment if your DL model uses Embeddings) ===\n    # 'carriers_used',\n    # 'searchRoute'\n]\n\n# === DATETIME FEATURES (Cyclic) ===\n# Dynamically add the sin/cos/business_time columns created in step [5/8]\ndatetime_features = [c for c in df.columns if any(\n    suffix in c for suffix in ['_hour_sin', '_hour_cos', '_weekday_sin', '_weekday_cos', '_business_time']\n)]\n\n# Combine to get the final model input list\ninput_features = feature_columns + datetime_features\n\n# Select only the engineered features\ndf_clean = df.select([c for c in input_features if c in df.columns])\ndf_id = df.select([c for c in group_feature if c in df.columns])\n\n\nprint(f\"   ✓ Created clean DataFrame with {len(df_clean.columns)} engineered features\")\nprint(f\"   Time: {time.time() - start_time:.2f}s\")\n\nprint(\"\\n\" + \"=\" * 60)\nprint(\"Simplified Feature Engineering Complete!\")\nprint(\"=\" * 60)\nprint(f\"Final dataset shape: {len(df_clean)} rows × {len(df_clean.columns)} columns\")\nprint(\"\\nFeature categories:\")\nprint(f\"  - IDs & Target: Id, ranker_id, selected\")\nprint(f\"  - Numerical ({15}): price, duration, baggage, fees, cabin, segments, trust\")\nprint(f\"  - Binary ({11}): trip type, direct flights, VIP, baggage, fees, carriers, routes\")\nprint(f\"  - Datetime ({len(datetime_features)}): cyclic hour/weekday encoding, business time, days_to_departure\")\nprint(f\"  - Categorical ({4}): airports (2), carriers_used, searchRoute\")\nprint(f\"  - Popularity ({1}): avg_carrier_popularity\")\nprint(f\"\\nTotal engineered features: {len(df_clean.columns)}\")","metadata":{"_uuid":"29c2b829-f97a-4c50-91b7-313325e88f9b","_cell_guid":"5a0442fd-5ec2-4959-b1cb-1eb1367471c2","trusted":true,"collapsed":false,"id":"430-fZ0gDxtO","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:41:23.916965Z","iopub.execute_input":"2025-12-04T00:41:23.917184Z","iopub.status.idle":"2025-12-04T00:41:23.928562Z","shell.execute_reply.started":"2025-12-04T00:41:23.917165Z","shell.execute_reply":"2025-12-04T00:41:23.926270Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Fill nulls\ndf_clean = df_clean.with_columns(\n    [pl.col(c).fill_null(0) for c in df_clean.select(pl.selectors.numeric()).columns] +\n    [pl.col(c).fill_null(\"missing\") for c in df_clean.select(pl.selectors.string()).columns]\n)","metadata":{"_uuid":"f4820b6d-9e70-4ea2-8ee8-af0f5053d708","_cell_guid":"e478d17c-a328-4f1e-93e2-61d254479366","trusted":true,"collapsed":false,"id":"IKNupXB7O1mP","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:41:23.930346Z","iopub.execute_input":"2025-12-04T00:41:23.930514Z","iopub.status.idle":"2025-12-04T00:41:23.960624Z","shell.execute_reply.started":"2025-12-04T00:41:23.930496Z","shell.execute_reply":"2025-12-04T00:41:23.957578Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(df_id)\nprint(df_clean)","metadata":{"_uuid":"0514d478-6048-40ab-9f89-d114cc44d93e","_cell_guid":"9b24f27c-9968-4f4c-90af-c868596e4e20","trusted":true,"collapsed":false,"id":"e03KC-q_DxtO","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:41:23.961944Z","iopub.execute_input":"2025-12-04T00:41:23.962110Z","iopub.status.idle":"2025-12-04T00:41:23.968459Z","shell.execute_reply.started":"2025-12-04T00:41:23.962083Z","shell.execute_reply":"2025-12-04T00:41:23.966194Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Write output\nprint(\"\\nWriting features to CSV...\")\nstart_time = time.time()\ndf_clean[:1000].write_csv(\"feature_dl_simplified.csv\")\nprint(f\"   ✓ Saved {len(df_clean)} rows, {len(df_clean.columns)} columns to 'feature_dl_simplified.csv'\")\nprint(f\"   Time: {time.time() - start_time:.2f}s\")","metadata":{"_uuid":"846d7bda-0e32-417a-9ad2-62cc388c0d05","_cell_guid":"5c8eea42-1dfc-4519-b811-71bdd7be981d","trusted":true,"collapsed":false,"id":"IcjxjwrYDxtO","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:41:23.971023Z","iopub.execute_input":"2025-12-04T00:41:23.971214Z","iopub.status.idle":"2025-12-04T00:41:23.984435Z","shell.execute_reply.started":"2025-12-04T00:41:23.971197Z","shell.execute_reply":"2025-12-04T00:41:23.979727Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\ndef plot_feature_values(X, max_samples=500, figsize=(16, 8)):\n    \"\"\"\n    X: numpy array or torch tensor [num_samples, num_features]\n    max_samples: limit how many samples to plot (prevents overload)\n    \"\"\"\n    if hasattr(X, \"cpu\"):\n        X = X.cpu().numpy()\n\n    # sample limit for readability\n    if X.shape[0] > max_samples:\n        X = X[:max_samples]\n\n    num_samples, num_features = X.shape\n\n    plt.figure(figsize=figsize)\n    for f in range(num_features):\n        plt.plot(X[:, f], alpha=0.4)\n\n    plt.title(f\"Feature values for {num_features} features across {num_samples} samples\")\n    plt.xlabel(\"Sample index\")\n    plt.ylabel(\"Feature value\")\n    plt.grid(True)\n    plt.tight_layout()\n    plt.show()\n\nimport numpy as np\nimport torch\n\ndef find_nan_features(X):\n    \"\"\"\n    Check for NaNs in a feature matrix.\n\n    X: np.ndarray or torch.Tensor [num_samples, num_features]\n\n    Returns:\n        nan_rows: indices of rows with NaNs\n        nan_cols: indices of columns with NaNs\n    \"\"\"\n    if isinstance(X, torch.Tensor):\n        nan_mask = torch.isnan(X)\n        nan_rows = torch.any(nan_mask, dim=1).nonzero(as_tuple=True)[0].cpu().numpy()\n        nan_cols = torch.any(nan_mask, dim=0).nonzero(as_tuple=True)[0].cpu().numpy()\n    else:  # assume numpy\n        nan_mask = np.isnan(X)\n        nan_rows = np.where(np.any(nan_mask, axis=1))[0]\n        nan_cols = np.where(np.any(nan_mask, axis=0))[0]\n\n    return nan_rows, nan_cols\n\n\nX = df_clean.to_numpy()\n# X = scaler.fit_transform(df_clean.to_numpy())\n# plot_feature_values(X)\n\nnan_rows, nan_cols = find_nan_features(X)\n\nprint(f\"Rows with NaN: {nan_rows}\")\nprint(f\"Columns with NaN: {nan_cols}\")\nprint(f\"Number of NaNs in total: {np.isnan(X).sum() if isinstance(X, np.ndarray) else torch.isnan(X).sum()}\")\n\nif len(nan_cols) > 0:\n    print(\"Feature indices containing NaN:\", nan_cols)\nelse:\n    print(\"No NaNs in features!\")","metadata":{"_uuid":"d21c2536-d9d8-4263-9bba-026e8cc477f4","_cell_guid":"c8c64283-3bf4-4a8d-8c55-e87c10dcc959","trusted":true,"collapsed":false,"id":"pBQJbzpVHk7X","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:41:23.985936Z","iopub.execute_input":"2025-12-04T00:41:23.986120Z","iopub.status.idle":"2025-12-04T00:41:26.107863Z","shell.execute_reply.started":"2025-12-04T00:41:23.986087Z","shell.execute_reply":"2025-12-04T00:41:26.104859Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Preparation","metadata":{"_uuid":"ce2116a1-280b-498f-ba9e-2470636cfe55","_cell_guid":"701575ee-6b04-4ba9-a63b-87280f2f480a","trusted":true,"collapsed":false,"id":"ahS3kplFDxtO","jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\n\nclass RankingDataset(Dataset):\n    def __init__(self, samples):\n        self.samples = samples\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        entry = self.samples[idx]\n        features = torch.tensor(entry[\"features\"], dtype=torch.float32)\n        pos_idx = torch.tensor(entry[\"positive_idx\"], dtype=torch.long)\n        return features, pos_idx, entry[\"Id\"], entry[\"ranker_id\"]\n","metadata":{"_uuid":"84aac078-1e22-4d44-849a-059a0c2599a3","_cell_guid":"29beb0f9-aebe-42b1-aefd-1550f95e1475","trusted":true,"collapsed":false,"id":"MusM0QqhDxtP","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:41:26.110153Z","iopub.execute_input":"2025-12-04T00:41:26.110865Z","iopub.status.idle":"2025-12-04T00:41:26.118052Z","shell.execute_reply.started":"2025-12-04T00:41:26.110844Z","shell.execute_reply":"2025-12-04T00:41:26.115455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom collections import defaultdict\nimport numpy as np\nfrom sklearn.preprocessing import StandardScaler\n\n\n\n# ============================================================\n# Build samples from groups\n# ============================================================\n\ndef build_samples(df_clean, df_id, test_sample = False):\n    print(f\"Data with {len(df_clean)} rows\")\n    \n    groups = defaultdict(list)\n\n    scaler = StandardScaler()\n    X = scaler.fit_transform(df_clean.to_numpy())\n    \n    # Process ALL rows, not just first 20\n    for i in range(len(df_id)):\n        id = df_id['Id'][i]\n        rid = df_id['ranker_id'][i]\n        label = df_id['selected'][i]\n        feat = X[i]\n        groups[rid].append((feat, label, id))\n\n    print(f\"\\nNumber of unique rankers: {len(groups)}\")\n\n    samples = []\n    \n    for ranker_id, items in groups.items():\n        features_list = []\n        id_list = []\n        positive_idx = None\n        for i, (feat, label, ids) in enumerate(items):\n            features_list.append(feat)\n            id_list.append(ids)\n            if label == 1:\n                positive_idx = i\n\n        if not test_sample and positive_idx is None:\n            continue\n        if test_sample: positive_idx = -1\n\n        features_array = np.stack(features_list, axis=0)\n\n        samples.append({\n            \"features\": features_array,\n            \"positive_idx\": positive_idx,\n            \"Id\": id_list,\n            \"ranker_id\": ranker_id              # <── added\n        })    \n    return samples\n\n# ============================================================\n# Train/Val/Test Split\n# ============================================================\n\ntrain_df_feat, test_df_feat = df_clean[:test_start_idx], df_clean[test_start_idx:]\ntrain_df_id, test_df_id = df_id[:test_start_idx], df_id[test_start_idx:]\n\n\n\ntrain_samples = build_samples(train_df_feat, train_df_id)\ntrain_s, val_s = train_test_split(train_samples, test_size=0.1, shuffle=True)\n\ntest_s = build_samples(test_df_feat, test_df_id, test_sample = True)\n\ntrain_ds = RankingDataset(train_s)\nval_ds   = RankingDataset(val_s)\ntest_ds  = RankingDataset(test_s)\nprint(f\"Train: {len(train_ds)} samples, Val: {len(val_ds)} samples, Test: {len(test_ds)} samples\")","metadata":{"_uuid":"780fdd4c-e969-4c36-b842-8e0e420307b2","_cell_guid":"7615c98f-b777-4fc3-8365-f4acc7d7eaf1","trusted":true,"collapsed":false,"id":"JMQilDsoCHl-","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:41:54.281020Z","iopub.execute_input":"2025-12-04T00:41:54.281282Z","iopub.status.idle":"2025-12-04T00:43:57.761726Z","shell.execute_reply.started":"2025-12-04T00:41:54.281264Z","shell.execute_reply":"2025-12-04T00:43:57.757135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\n\nclass MLPScoreModel(nn.Module):\n    def __init__(self, input_dim, hidden_dim=256, layers=3):\n        super().__init__()\n        blocks = []\n        dim = input_dim\n        for _ in range(layers):\n            blocks.append(nn.Linear(dim, hidden_dim))\n            blocks.append(nn.ReLU())\n            dim = hidden_dim\n        blocks.append(nn.Linear(dim, 1))  # output score per item\n        self.net = nn.Sequential(*blocks)\n\n    def forward(self, x):\n        # x = [N, D]\n        return self.net(x).squeeze(-1)  # [N]","metadata":{"_uuid":"3815c017-c141-4d25-8af0-9c865c93f18a","_cell_guid":"0021c657-4f77-4f23-9e6b-97ee539e4c72","trusted":true,"collapsed":false,"id":"GLFWev7BDxtP","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:43:57.762257Z","iopub.execute_input":"2025-12-04T00:43:57.762429Z","iopub.status.idle":"2025-12-04T00:43:57.774409Z","shell.execute_reply.started":"2025-12-04T00:43:57.762413Z","shell.execute_reply":"2025-12-04T00:43:57.771034Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Training","metadata":{"_uuid":"26ed510f-46a9-4f6f-8aa6-e3eb0a4f8729","_cell_guid":"fae6312b-f2b8-41cd-8f98-a1719dacf236","trusted":true,"collapsed":false,"id":"w0O0daIADxtP","jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom accelerate import Accelerator\nfrom dataclasses import dataclass\nimport numpy as np\n\n\n# ============================================================\n# Training Configuration (Dataclass)\n# ============================================================\n\n@dataclass\nclass TrainingConfig:\n    batch_size: int = 128\n    hidden_dim: int = 512\n    layers: int = 4\n    lr: float = 1e-4\n    num_epochs: int = 30\n\n    # Dynamically set later\n    input_dim: int = None\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nconfig = TrainingConfig()","metadata":{"_uuid":"9cc5eeb1-c7fb-4c0c-a5c4-ae8669b7b324","_cell_guid":"a071144d-979e-4f02-8299-2893fb6d32a3","trusted":true,"collapsed":false,"id":"ZEqxjaFUBbR-","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:43:57.775648Z","iopub.execute_input":"2025-12-04T00:43:57.775803Z","iopub.status.idle":"2025-12-04T00:44:10.116591Z","shell.execute_reply.started":"2025-12-04T00:43:57.775787Z","shell.execute_reply":"2025-12-04T00:44:10.111540Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader\n\n# TPU/XLA imports\nimport torch_xla\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as pl\n\n# --- 1. Optimized Collate Function (Padding & Masking) ---\ndef bucketed_collate_fn(batch):\n    \"\"\"\n    Factory function to create a bucket-based collate function.\n    \"\"\"\n    bucket_boundaries = [5, 20, 50, 160, 400, 620, 8300]\n    \n    # Sort boundaries to ensure correct bucketing\n    bucket_boundaries = sorted(bucket_boundaries)\n    \n    batch_size = len(batch)\n    \n    # Extract components\n    raw_feats = [b[0] for b in batch]\n    pos_idx = torch.tensor([b[1] for b in batch], dtype=torch.long)\n    \n    # Get dimensions\n    lengths = torch.tensor([f.shape[0] for f in raw_feats], dtype=torch.long)\n    max_len_in_batch = lengths.max().item()\n    feature_dim = raw_feats[0].shape[1]\n    \n    # Determine the appropriate bucket size\n    padded_len = max_len_in_batch\n    for boundary in bucket_boundaries:\n        if max_len_in_batch <= boundary:\n            padded_len = boundary\n            break\n    else:\n        # If exceeds all boundaries, use the max length in batch\n        # or optionally round up to nearest multiple\n        padded_len = max_len_in_batch\n    \n    # Pre-allocate tensors with bucket size\n    padded_features = torch.zeros((batch_size, padded_len, feature_dim), dtype=torch.float32)\n    attention_masks = torch.zeros((batch_size, padded_len), dtype=torch.bool)\n    \n    # Vectorized padding\n    for i, f in enumerate(raw_feats):\n        curr_len = f.shape[0]\n        padded_features[i, :curr_len] = f\n        attention_masks[i, :curr_len] = True\n    \n    return padded_features, pos_idx, attention_masks\n    \n# --- 2. Vectorized Loss Function ---\ndef batch_ranking_loss(scores, pos_idxs, mask):\n    \"\"\"\n    scores:   [Batch, Max_Len]\n    pos_idxs: [Batch]\n    mask:     [Batch, Max_Len]\n    \"\"\"\n    scores = scores.masked_fill(mask == 0, -1e9)\n    return F.cross_entropy(scores, pos_idxs)\n\n# --- 3. FIXED Single-Core TPU Training ---\ndef train_model(config, train_dataset, val_dataset=None):\n    device = xm.xla_device()\n    xm.master_print(f\"Training on TPU device: {device}\")\n    \n    xm.master_print(f\"Batch size: {config.batch_size}\")\n    \n    # DataLoader settings\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.batch_size,\n        shuffle=True,\n        collate_fn=bucketed_collate_fn,\n        num_workers=4,\n        drop_last=True,\n        prefetch_factor=2,\n        persistent_workers=True,  # FIX: Reuse workers\n    )\n    \n    val_device_loader = None\n    if val_dataset:\n        val_loader = DataLoader(\n            val_dataset,\n            batch_size=config.batch_size,\n            shuffle=False,\n            collate_fn=bucketed_collate_fn,\n            num_workers=4,\n            drop_last=True,\n            prefetch_factor=2,\n            persistent_workers=True,\n        )\n    \n    config.input_dim = train_dataset[0][0].shape[-1]\n    \n    model = MLPScoreModel(\n        input_dim=config.input_dim,\n        hidden_dim=config.hidden_dim,\n        layers=config.layers\n    ).to(device)\n    \n    # Enable bfloat16 for 2-3x speedup on TPU\n    use_bfloat16 = getattr(config, 'use_bfloat16', True)\n    if use_bfloat16:\n        xm.master_print(\"Converting model to bfloat16 for faster training...\")\n        model = model.to(torch.bfloat16)\n    \n    optimizer = torch.optim.AdamW(\n        model.parameters(), \n        lr=config.lr,\n        weight_decay=0.01\n    )\n    \n    xm.master_print(\"Setup complete. Start training...\")\n    \n    # FIX: More frequent marking to prevent graph explosion\n    log_steps = max(len(train_loader) // 5, 1)  # Log 20 times per epoch\n    \n    xm.master_print(f\"Log every {log_steps} steps\")\n    \n    try:\n        for epoch in range(config.num_epochs):\n            model.train()\n            \n            # FIX: Use scalar accumulation instead of tensor accumulation\n            epoch_loss_sum = 0.0\n            step_count = 0\n            \n            # FIX: Recreate ParallelLoader each epoch to avoid exhaustion\n            train_device_loader = pl.ParallelLoader(train_loader, [device]).per_device_loader(device)\n            \n            for step, (batch_feats, batch_pos, batch_masks) in enumerate(train_device_loader):\n                \n                # Cast inputs to bfloat16 if enabled\n                if use_bfloat16:\n                    batch_feats = batch_feats.to(torch.bfloat16)\n                \n                out = model(batch_feats)\n                scores = out.squeeze(-1)\n                \n                # Loss calculation stays in float32 for numerical stability\n                if use_bfloat16:\n                    scores = scores.to(torch.float32)\n                \n                loss = batch_ranking_loss(scores, batch_pos, batch_masks)\n                \n                loss.backward()\n                \n                # Gradient clipping\n                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                \n                xm.optimizer_step(optimizer)\n                optimizer.zero_grad()\n                \n                # FIX: Convert to Python scalar immediately to avoid memory buildup\n                epoch_loss_sum += loss.item()\n                step_count += 1\n                \n                # FIX: Mark step regularly to prevent graph explosion\n                xm.mark_step()\n                \n                # # Logging\n                # if step_count % log_steps == 0:\n                #     avg_loss = epoch_loss_sum / step_count\n                #     xm.master_print(f\"[Epoch {epoch+1}] Step {step}/{len(train_loader)} - Loss: {avg_loss:.4f}\")\n            \n            # Mark step at epoch end\n            xm.mark_step()\n            \n            final_avg_loss = epoch_loss_sum / step_count\n            xm.master_print(f\"[Epoch {epoch+1}] Train Loss: {final_avg_loss:.4f}\")\n            \n            # ---- Validation ----\n            if val_dataset and val_loader is not None:\n                model.eval()\n                val_loss_sum = 0.0\n                val_steps = 0\n\n                xm.master_print(\"Start Validating\")\n                \n                # FIX: Recreate ParallelLoader for validation each epoch\n                val_device_loader = pl.ParallelLoader(val_loader, [device]).per_device_loader(device)\n                \n                with torch.no_grad():\n                    for step, (batch_feats, batch_pos, batch_masks) in enumerate(val_device_loader):\n                        \n                        # Cast to bfloat16 if enabled\n                        if use_bfloat16:\n                            batch_feats = batch_feats.to(torch.bfloat16)\n                        \n                        out = model(batch_feats)\n                        scores = out.squeeze(-1)\n                        \n                        # Convert back to float32 for loss\n                        if use_bfloat16:\n                            scores = scores.to(torch.float32)\n                        \n                        loss = batch_ranking_loss(scores, batch_pos, batch_masks)\n                        \n                        # FIX: Immediate conversion to Python scalar\n                        val_loss_sum += loss.item()\n                        val_steps += 1\n                        \n                        # FIX: Mark step every iteration in validation\n                        xm.mark_step()\n                \n                xm.mark_step()\n                avg_val_loss = val_loss_sum / val_steps\n                xm.master_print(f\"            Val Loss: {avg_val_loss:.4f}\")\n            \n            # Save checkpoint less frequently\n            if (epoch + 1) % 5 == 0 or epoch == config.num_epochs - 1:\n                xm.master_print(f\"Saving checkpoint at epoch {epoch+1}...\")\n                xm.save(model.state_dict(), f'model_checkpoint_epoch_{epoch+1}.pth')\n                xm.mark_step()  # FIX: Mark step after save\n\n    except Exception as e:\n        xm.master_print(f\"Training interrupted: {e}\")\n        raise\n    finally:\n        # FIX: Cleanup to prevent issues on reruns\n        xm.master_print(\"Cleaning up...\")\n        xm.mark_step()\n        torch.cuda.empty_cache() if torch.cuda.is_available() else None\n\n    xm.master_print(\"Training complete!\")\n    \n    # Convert model back to float32 for saving/inference\n    if use_bfloat16:\n        model = model.to(torch.float32)\n    \n    return model","metadata":{"_uuid":"1398ed61-3dd1-43de-abcc-80b96629b43a","_cell_guid":"846c1eb5-1594-48bc-9ea8-d7ff8f336113","trusted":true,"collapsed":false,"id":"jvTHXXAaDxtP","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:44:46.479284Z","iopub.execute_input":"2025-12-04T00:44:46.479592Z","iopub.status.idle":"2025-12-04T00:44:46.505723Z","shell.execute_reply.started":"2025-12-04T00:44:46.479573Z","shell.execute_reply":"2025-12-04T00:44:46.501282Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = train_model(config, train_ds, val_ds)","metadata":{"_uuid":"2bf05c87-41a7-4fb4-8264-608de0f9beda","_cell_guid":"f9cade41-cc09-417d-b077-3852cfa2c0c7","trusted":true,"collapsed":false,"id":"DjTl5PbUDE_P","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T00:44:49.170356Z","iopub.execute_input":"2025-12-04T00:44:49.170652Z","iopub.status.idle":"2025-12-04T01:12:47.330882Z","shell.execute_reply.started":"2025-12-04T00:44:49.170632Z","shell.execute_reply":"2025-12-04T01:12:47.327273Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Error analysis and visualization","metadata":{"_uuid":"95d634c1-c299-4a13-926b-eae2903bdd5a","_cell_guid":"63de2b12-90e9-4f86-b014-73fe1e7de68e","trusted":true,"collapsed":false,"id":"_tuvBCqTDxtP","jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as pl\nimport polars\nimport numpy as np\n\n# Get TPU device\ndevice = xm.xla_device()\n\ndef bucketed_test_collate_fn(batch):\n    \"\"\"\n    Bucket-based collate function for efficient padding.\n    \"\"\"\n    bucket_boundaries = [5, 20, 50, 160, 400, 620, 8300]\n    bucket_boundaries = sorted(bucket_boundaries)\n    \n    batch_size = len(batch)\n    \n    # Extract components\n    raw_feats = [torch.tensor(b[0], dtype=torch.float32) for b in batch]\n    pos_idx = [b[1] for b in batch]\n    ids = [b[2] for b in batch]\n    rank_ids = [b[3] for b in batch]\n    \n    # Get dimensions\n    lengths = torch.tensor([f.shape[0] for f in raw_feats], dtype=torch.long)\n    max_len_in_batch = lengths.max().item()\n    feature_dim = raw_feats[0].shape[1]\n    \n    # Determine the appropriate bucket size\n    padded_len = max_len_in_batch\n    for boundary in bucket_boundaries:\n        if max_len_in_batch <= boundary:\n            padded_len = boundary\n            break\n    \n    # Pre-allocate tensors with bucket size\n    padded_features = torch.zeros((batch_size, padded_len, feature_dim), dtype=torch.float32)\n    attention_masks = torch.zeros((batch_size, padded_len), dtype=torch.bool)\n    \n    # Fill in the data\n    for i, f in enumerate(raw_feats):\n        curr_len = f.shape[0]\n        padded_features[i, :curr_len] = f\n        attention_masks[i, :curr_len] = True\n    \n    return padded_features, lengths, attention_masks, pos_idx, ids, rank_ids\n\n# DataLoader with larger batch size\nBATCH_SIZE = 32  # Adjust based on your TPU memory\nval_loader = DataLoader(\n    val_ds, \n    collate_fn=bucketed_test_collate_fn, \n    batch_size=BATCH_SIZE, \n    shuffle=False\n)\n\n# Wrap with MpDeviceLoader for TPU\nval_loader = pl.MpDeviceLoader(val_loader, device)\n\n# ============================================================\n# Inference on val_ds → build val_df\n# ============================================================\nmodel.eval()\nmodel = model.to(device)  # Ensure model is on TPU\n\nall_ids = []\nall_ranker_ids = []\nall_pred_scores = []\nall_selected = []\n\nwith torch.no_grad():\n    for batch_idx, (padded_features, lengths, attention_masks, pos_idx_list, ids_list, rank_ids_list) in enumerate(val_loader):\n        # All tensors are already on TPU device via MpDeviceLoader\n        # padded_features: (batch_size, padded_len, feature_dim)\n        # lengths: (batch_size,)\n        # attention_masks: (batch_size, padded_len)\n        \n        batch_size = padded_features.size(0)\n        \n        # Forward pass - pass attention mask if your model supports it\n        preds_batch = model(padded_features)  # (batch_size, padded_len)\n        \n        # Mark step every few batches for optimal TPU performance\n        if batch_idx % 5 == 0:\n            xm.mark_step()\n        \n        # Move to CPU for processing\n        preds_cpu = preds_batch.cpu()\n        lengths_cpu = lengths.cpu()\n        xm.mark_step()  # Ensure transfer completes\n        \n        preds_np = preds_cpu.numpy()\n        lengths_np = lengths_cpu.numpy()\n        \n        # Process each item in the batch\n        for i in range(batch_size):\n            L = lengths_np[i]  # Actual length (before padding)\n            preds_item = preds_np[i, :L]  # Only take non-padded predictions\n            pos_idx = pos_idx_list[i]\n            id_group = ids_list[i]\n            rid = rank_ids_list[i]\n            \n            # Build labels for the whole group\n            selected = np.zeros(L, dtype=int)\n            selected[pos_idx] = 1\n            \n            # Store - handle if id_group is a list or single value\n            if isinstance(id_group, list):\n                all_ids.extend(id_group)\n            else:\n                all_ids.extend([id_group] * L)\n            \n            all_ranker_ids.extend([rid] * L)\n            all_pred_scores.extend(preds_item.tolist())\n            all_selected.extend(selected.tolist())\n\n# Final mark_step to complete all pending operations\nxm.mark_step()\n\n# Convert to Polars DataFrame\nval_df = polars.DataFrame({\n    \"Id\": all_ids,\n    \"ranker_id\": all_ranker_ids,\n    \"pred_score\": all_pred_scores,\n    \"selected\": all_selected\n})\n\n# Add group_size\nval_df = val_df.join(\n    val_df.group_by(\"ranker_id\").agg(polars.len().alias(\"group_size\")),\n    on=\"ranker_id\"\n)\n\nprint(val_df)","metadata":{"_uuid":"f53e241b-49a9-40f1-9d09-b252d2ab0c41","_cell_guid":"874d2eec-9a94-4e3f-b324-f893e4ecc24f","trusted":true,"collapsed":false,"id":"fc-oLzVKODP7","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T01:14:11.664012Z","iopub.execute_input":"2025-12-04T01:14:11.664316Z","iopub.status.idle":"2025-12-04T01:14:31.234469Z","shell.execute_reply.started":"2025-12-04T01:14:11.664298Z","shell.execute_reply":"2025-12-04T01:14:31.229596Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Visualization on val_df\n# ============================================================\n\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport polars as pl\n\n# Color palette\nred = (0.86, 0.08, 0.24)\nblue = (0.12, 0.56, 1.0)\n\n\n\n# ============================================================\n# val_df requirement\n# Columns available:\n#   - ranker_id\n#   - pred_score\n#   - selected\n#   - group_size\n# ============================================================\n\n# Keep only groups with more than 10 items\nva_df = val_df.filter(pl.col(\"group_size\") > 10)\n\n# Compute quantiles of group sizes (unique per ranker_id)\nsize_quantiles = (\n    va_df.select(\"ranker_id\", \"group_size\")\n         .unique()\n         .select(\n             pl.col(\"group_size\").quantile(0.25).alias(\"q25\"),\n             pl.col(\"group_size\").quantile(0.50).alias(\"q50\"),\n             pl.col(\"group_size\").quantile(0.75).alias(\"q75\")\n         )\n         .to_dicts()[0]\n)\n\n\n\n# ============================================================\n# HitRate@k Curve Function\n# ============================================================\ndef calculate_hitrate_curve(df, k_values):\n    sorted_df = df.sort([\"ranker_id\", \"pred_score\"], descending=[False, True])\n    return [\n        (\n            sorted_df.group_by(\"ranker_id\", maintain_order=True)\n            .head(k)\n            .group_by(\"ranker_id\")\n            .agg(pl.col(\"selected\").max().alias(\"hit\"))\n            .select(pl.col(\"hit\").mean())\n            .item()\n        )\n        for k in k_values\n    ]\n\n\nk_values = list(range(1, 21))\n\ncurves = {\n    'All groups (>10)': calculate_hitrate_curve(va_df, k_values),\n    f\"Small (11-{int(size_quantiles['q25'])})\": calculate_hitrate_curve(\n        va_df.filter(pl.col('group_size') <= size_quantiles['q25']), k_values\n    ),\n    f\"Medium ({int(size_quantiles['q25']+1)}-{int(size_quantiles['q75'])})\": calculate_hitrate_curve(\n        va_df.filter(\n            (pl.col('group_size') > size_quantiles['q25']) &\n            (pl.col('group_size') <= size_quantiles['q75'])\n        ),\n        k_values\n    ),\n    f\"Large (>{int(size_quantiles['q75'])})\": calculate_hitrate_curve(\n        va_df.filter(pl.col('group_size') > size_quantiles['q75']),\n        k_values\n    ),\n}\n\n\n\n# ============================================================\n# HitRate@3 vs Group Size (log-binned)\n# ============================================================\n\n# Create log bins\nmin_size = va_df['group_size'].min()\nmax_size = va_df['group_size'].max()\nbins = np.logspace(np.log10(min_size), np.log10(max_size), 51)\n\n# One HitRate@3 per ranker\nranker_hr3 = (\n    va_df.sort([\"ranker_id\", \"pred_score\"], descending=[False, True])\n         .group_by(\"ranker_id\", maintain_order=True)\n         .agg([\n             pl.col(\"selected\").head(3).max().alias(\"hit_top3\"),\n             pl.col(\"group_size\").first()\n         ])\n)\n\n# Assign bins\nbin_centers = (bins[:-1] + bins[1:]) / 2\nbin_indices = np.digitize(ranker_hr3[\"group_size\"].to_numpy(), bins) - 1\n\nsize_analysis = (\n    pl.DataFrame({\n        \"bin_idx\": bin_indices,\n        \"bin_center\": bin_centers[np.clip(bin_indices, 0, len(bin_centers)-1)],\n        \"hit_top3\": ranker_hr3[\"hit_top3\"]\n    })\n    .group_by([\"bin_idx\", \"bin_center\"])\n    .agg([\n        pl.col(\"hit_top3\").mean().alias(\"hitrate3\"),\n        pl.len().alias(\"n_groups\")\n    ])\n    .filter(pl.col(\"n_groups\") >= 3)\n    .sort(\"bin_center\")\n)\n\n\n\n# ============================================================\n# Plot Figures\n# ============================================================\nfig, (ax1, ax2) = plt.subplots(1, 2, figsize=(8, 4), dpi=400)\n\n# Left: HitRate@k Curves\ncolors = [\"black\"]  # All groups\nfor i in range(3):\n    t = i / 2\n    color = tuple(blue[j] * (1 - t) + red[j] * t for j in range(3))\n    colors.append(color)\n\nfor (label, hitrates), color in zip(curves.items(), colors):\n    ax1.plot(k_values, hitrates, marker='o', label=label, color=color, markersize=3)\n\nax1.set_xlabel(\"k (top-k predictions)\")\nax1.set_ylabel(\"HitRate@k\")\nax1.set_title(\"HitRate@k by Group Size\")\nax1.legend(fontsize=8)\nax1.grid(True, alpha=0.3)\nax1.set_xlim(0, 21)\nax1.set_ylim(-0.025, 1.025)\n\n# Right: HitRate@3 vs Group Size\nax2.scatter(\n    size_analysis[\"bin_center\"],\n    size_analysis[\"hitrate3\"],\n    s=30,\n    alpha=0.6,\n    color=blue\n)\nax2.set_xlabel(\"Group Size\")\nax2.set_ylabel(\"HitRate@3\")\nax2.set_title(\"HitRate@3 vs Group Size\")\nax2.set_xscale(\"log\")\nax2.grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.show()","metadata":{"_uuid":"0322ac8b-0c78-4056-988f-1b96d60b0a65","_cell_guid":"b7b062ca-ac28-4ccd-a972-59d2f61f220a","trusted":true,"id":"V8UA7dAODxtP","execution":{"iopub.status.busy":"2025-12-04T01:20:20.734899Z","iopub.execute_input":"2025-12-04T01:20:20.735195Z","iopub.status.idle":"2025-12-04T01:20:23.408012Z","shell.execute_reply.started":"2025-12-04T01:20:20.735177Z","shell.execute_reply":"2025-12-04T01:20:23.403923Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Summary\nprint(f\"HitRate@1: {curves['All groups (>10)'][0]:.3f}\")\nprint(f\"HitRate@3: {curves['All groups (>10)'][2]:.3f}\")\nprint(f\"HitRate@5: {curves['All groups (>10)'][4]:.3f}\")\nprint(f\"HitRate@10: {curves['All groups (>10)'][9]:.3f}\")","metadata":{"_uuid":"5447312e-dced-4bb5-b2be-9d416b6377e3","_cell_guid":"a2b14618-cdda-4746-a1f3-75b695ff10eb","trusted":true,"collapsed":false,"id":"N93NlvyeDxtP","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T01:20:30.219762Z","iopub.execute_input":"2025-12-04T01:20:30.220027Z","iopub.status.idle":"2025-12-04T01:20:30.228034Z","shell.execute_reply.started":"2025-12-04T01:20:30.220009Z","shell.execute_reply":"2025-12-04T01:20:30.225058Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Submission","metadata":{"_uuid":"2573e964-71d1-49b0-8114-9973f0f51590","_cell_guid":"b56fce37-5aea-4d71-91b2-d6e7dbcdd3b4","trusted":true,"collapsed":false,"id":"1zfnXgFxDxtP","jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as xla_pl\nimport polars as pl\nfrom torch.utils.data import DataLoader\n\n# Get TPU device\ndevice = xm.xla_device()\n\n# DataLoader with bucketed collate function\nBATCH_SIZE = 32  # Adjust based on TPU memory and your data\ntest_loader = DataLoader(\n    test_ds, \n    collate_fn=bucketed_test_collate_fn,\n    batch_size=BATCH_SIZE, \n    shuffle=False\n)\n\n# Wrap the DataLoader with ParallelLoader for TPU optimization\ntest_loader = xla_pl.MpDeviceLoader(test_loader, device)\n\n# ============================================================\n# Inference on test_ds → build test_df\n# ============================================================\nmodel.eval()\nmodel = model.to(device)  # Ensure model is on TPU\n\nall_ids = []\nall_ranker_ids = []\nall_pred_scores = []\nall_selected = []\n\nwith torch.no_grad():\n    for batch_idx, (padded_features, lengths, attention_masks, pos_idx_list, ids_list, rank_ids_list) in enumerate(test_loader):\n        # All tensors already on TPU device via MpDeviceLoader\n        # padded_features: (batch_size, padded_len, feature_dim)\n        # lengths: (batch_size,)\n        # attention_masks: (batch_size, padded_len)\n        \n        batch_size = padded_features.size(0)\n        \n        # Forward pass with attention mask (if your model supports it)\n        preds_batch = model(padded_features)\n        \n        # Mark step every few batches for optimal TPU performance\n        if batch_idx % 5 == 0:\n            xm.mark_step()\n        \n        # Move to CPU for numpy conversion\n        preds_cpu = preds_batch.cpu()\n        lengths_cpu = lengths.cpu()\n        xm.mark_step()  # Ensure transfer completes\n        \n        preds_np = preds_cpu.numpy()\n        lengths_np = lengths_cpu.numpy()\n        \n        # Process each item in the batch\n        for i in range(batch_size):\n            L = lengths_np[i]  # Actual length (before padding)\n            preds_item = preds_np[i, :L]  # Only take non-padded predictions\n            id_group = ids_list[i]\n            rid = rank_ids_list[i]\n            \n            # Build labels based on MAX SCORE (Argmax)\n            selected = np.zeros(L, dtype=int)\n            \n            # Find the index where the model output is highest\n            best_idx = np.argmax(preds_item)\n            \n            # Set that index to 1\n            selected[best_idx] = 1\n            \n            # Store - handle if id_group is a list or single value\n            if isinstance(id_group, list):\n                all_ids.extend(id_group)\n            else:\n                all_ids.extend([id_group] * L)\n            \n            all_ranker_ids.extend([rid] * L)\n            all_pred_scores.extend(preds_item.tolist())\n            all_selected.extend(selected.tolist())\n\n# Final mark_step to complete all pending operations\nxm.mark_step()\n\n# Convert to Polars DataFrame\ntest_df = pl.DataFrame({\n    \"Id\": all_ids,\n    \"ranker_id\": all_ranker_ids,\n    \"pred_score\": all_pred_scores,\n    \"selected\": all_selected  # Now contains 1 for the highest score, 0 for others\n})\n\n# Add group_size\ntest_df = test_df.join(\n    test_df.group_by(\"ranker_id\").agg(pl.len().alias(\"group_size\")),\n    on=\"ranker_id\"\n)\n\n# Filter to show only selected items\nfiltered_df = test_df.filter(pl.col(\"selected\") == 1)\nprint(filtered_df)\nprint(f\"Total groups: {test_df['ranker_id'].n_unique()}\")\nprint(f\"Selected items: {len(filtered_df)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T01:23:33.069116Z","iopub.execute_input":"2025-12-04T01:23:33.069449Z","iopub.status.idle":"2025-12-04T01:23:59.127864Z","shell.execute_reply.started":"2025-12-04T01:23:33.069430Z","shell.execute_reply":"2025-12-04T01:23:59.122144Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Build submission from test_df + model predictions\n# ============================================================\n\n# test_df must contain:\n#   - Id\n#   - ranker_id\n#   - pred_score  (already added when building test_df)\n\nsubmission_df = (\n    test_df\n    .select([\"Id\", \"ranker_id\", \"pred_score\"])\n    .with_columns(\n        # Rank pred_score within each ranker_id group\n        # Higher pred_score → Lower rank number (rank 1 = best)\n        pl.col(\"pred_score\")\n        .rank(method=\"ordinal\", descending=True)\n        .over(\"ranker_id\")\n        .cast(pl.Int32)\n        .alias(\"selected\")\n    )\n    .select([\"Id\", \"ranker_id\", \"selected\"])\n    # CRITICAL: Sort by Id to preserve original test.csv row order\n    .sort(\"Id\")\n)\n\n# Save to CSV\nsubmission_df.write_csv(\"submission.csv\")\n\nprint(\"Submission saved to submission.csv\")\nprint(f\"Total rows: {len(submission_df)}\")\n\n# Validation checks\nprint(\"\\n=== Validation Checks ===\")\nvalidation = submission_df.group_by(\"ranker_id\").agg([\n    pl.col(\"selected\").min().alias(\"min_rank\"),\n    pl.col(\"selected\").max().alias(\"max_rank\"),\n    pl.col(\"selected\").n_unique().alias(\"unique_ranks\"),\n    pl.col(\"selected\").count().alias(\"n_flights\")\n])\n\n# Check if ranks form valid permutations (1, 2, 3, ..., N)\ninvalid = validation.filter(\n    (pl.col(\"min_rank\") != 1) |\n    (pl.col(\"max_rank\") != pl.col(\"n_flights\")) |\n    (pl.col(\"unique_ranks\") != pl.col(\"n_flights\"))\n)\n\nif len(invalid) > 0:\n    print(f\"⚠️  WARNING: {len(invalid)} ranker_ids have invalid rank permutations!\")\n    print(invalid)\nelse:\n    print(\"✓ All ranker_ids have valid rank permutations (1, 2, 3, ..., N)\")\n\nprint(f\"✓ Total unique ranker_ids: {submission_df['ranker_id'].n_unique()}\")","metadata":{"_uuid":"9ce31c0a-ea23-4271-bb72-0c5c82d207b5","_cell_guid":"637399ec-5eb4-489e-b2de-bf0a6613ee38","trusted":true,"collapsed":false,"id":"9e54gTPiDxtP","jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-04T01:23:59.130424Z","iopub.execute_input":"2025-12-04T01:23:59.130653Z","iopub.status.idle":"2025-12-04T01:24:00.061382Z","shell.execute_reply.started":"2025-12-04T01:23:59.130632Z","shell.execute_reply":"2025-12-04T01:24:00.055509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!kaggle competitions submit -c aeroclub-recsys-2025 -f submission.csv -m \"Deep Learning Ranking Model\"","metadata":{"_uuid":"5326090c-7a78-4ade-bb4b-7ab2dfedcd9b","_cell_guid":"399ed7cc-11de-4f73-90d9-5edd2ee5319a","trusted":true,"collapsed":false,"id":"bPDoRFGBn8Z6","jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}