{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":128567,"databundleVersionId":15597029}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":9693.326236,"end_time":"2026-02-09T10:22:37.194248","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-02-09T07:41:03.868012","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Rocket League Goal Predictor model\n\nThis notebook creates the goal prediction model used in the [Goal Predictor BakkesMod plugin](https://github.com/dster2/rocket-league-goal-predictor).  We:\n1. Prep training and validation datasets\n2. Define our ML model using keras\n3. Train the model using jax backend with multiple GPUs if available\n4. Run inferences on the competition test set and produces submission file\n5. Export ONNX model for use in the BakkesMod plugin\n\nTechnically this notebook does not produce *exactly* the model which is currently included in the plugin.  Besides non-determinism, this notebook has tweaked a few hyperparameters (`D_MODEL`, `FRACTION_VAL`, learning rate schedule, etc) to run more comfortably in the Kaggle environment with 12 hour time limit, vs the plugin model which I trained on my desktop, but it's not a major difference.","metadata":{}},{"cell_type":"code","source":"# Better plotting in polars\n!pip install -q hvplot\n# For exporting to ONNX, and take from source which has critical fixes not in pip release\n!pip install -q -U git+https://github.com/onnx/tensorflow-onnx\n# Fix hidden error in `import keras` via TensorFlow depending on older `protobuf`\n!pip install \"protobuf<4.21.0\" --force-reinstall --no-deps -q","metadata":{"_cell_guid":"a54a0e17-8b82-40c3-bf43-fbadc9d6ddb4","_kg_hide-input":true,"_kg_hide-output":true,"_uuid":"e60e9aac-35cd-4239-bf9f-3f80db5d04e1","collapsed":false,"execution":{"iopub.status.busy":"2026-02-22T15:27:54.525415Z","iopub.execute_input":"2026-02-22T15:27:54.525698Z","iopub.status.idle":"2026-02-22T15:28:48.556062Z","shell.execute_reply.started":"2026-02-22T15:27:54.525676Z","shell.execute_reply":"2026-02-22T15:28:48.555309Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":31.57762,"end_time":"2026-02-09T07:41:38.132427","exception":false,"start_time":"2026-02-09T07:41:06.554807","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nos.environ['KERAS_BACKEND'] = 'jax'\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3' # fatal\n\nINTERACTIVE_SESSION = os.getenv('KAGGLE_KERNEL_RUN_TYPE') == 'Interactive'\n# INTERACTIVE_SESSION = True\nprint(f'Interactive: {INTERACTIVE_SESSION}')","metadata":{"_cell_guid":"9ba54c09-75a3-48ef-9737-894a3c96be6d","_uuid":"add1b5bb-9f7e-413d-9529-108cc40ee733","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:41:38.148164Z","iopub.status.busy":"2026-02-09T07:41:38.147731Z","iopub.status.idle":"2026-02-09T07:41:38.152543Z","shell.execute_reply":"2026-02-09T07:41:38.151843Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.014121,"end_time":"2026-02-09T07:41:38.154009","exception":false,"start_time":"2026-02-09T07:41:38.139888","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"preloaded_model_path = os.getenv('KAGGLE_PRELOAD_MODEL_PATH')\nuse_preloaded_model = not not preloaded_model_path\nprint(f'preloaded: {preloaded_model_path} {use_preloaded_model}')","metadata":{"execution":{"iopub.execute_input":"2026-02-09T07:41:38.168470Z","iopub.status.busy":"2026-02-09T07:41:38.168245Z","iopub.status.idle":"2026-02-09T07:41:38.171658Z","shell.execute_reply":"2026-02-09T07:41:38.171052Z"},"papermill":{"duration":0.012433,"end_time":"2026-02-09T07:41:38.173211","exception":false,"start_time":"2026-02-09T07:41:38.160778","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import atexit\nimport collections\nimport concurrent\nfrom datetime import datetime\nimport hvplot.polars\nimport jax\nimport jax.numpy as jnp\nfrom jax import nn as jnn\nimport kagglehub\nimport keras\nimport numpy as np\nimport pathlib\nimport polars as pl\nimport random\nimport scipy.stats\nimport tempfile\nimport time\n\nseed = 2**16\nif preloaded_model_path:\n    seed = hash(preloaded_model_path) % 2**16\njax.random.key(seed)\nkeras.utils.set_random_seed(seed)\nnp.random.seed(seed)\npl.set_random_seed(seed)\nrandom.seed(seed)","metadata":{"_cell_guid":"9fda321f-054c-475a-b203-737c4dd5e95c","_uuid":"02554455-cebe-42c7-98a3-04e91ec41960","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:41:38.188838Z","iopub.status.busy":"2026-02-09T07:41:38.188396Z","iopub.status.idle":"2026-02-09T07:42:04.721022Z","shell.execute_reply":"2026-02-09T07:42:04.720392Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":26.541916,"end_time":"2026-02-09T07:42:04.722794","exception":false,"start_time":"2026-02-09T07:41:38.180878","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"devices = jax.local_devices()\nNUM_DEVICES = len(jax.local_devices())\nprint(f'Running on {NUM_DEVICES} devices: {jax.local_devices()}')","metadata":{"_cell_guid":"6394f9b0-10fa-4c53-b9f5-e712e3bc406f","_uuid":"728e6bb4-ba29-498d-9ed8-17faa28ab5ac","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:42:04.741040Z","iopub.status.busy":"2026-02-09T07:42:04.740168Z","iopub.status.idle":"2026-02-09T07:42:04.744384Z","shell.execute_reply":"2026-02-09T07:42:04.743771Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.014429,"end_time":"2026-02-09T07:42:04.746030","exception":false,"start_time":"2026-02-09T07:42:04.731601","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset_path = pathlib.Path(kagglehub.competition_download('rocket-league-rlcs-goal-prediction-2025'))\nprint(f'data at: {dataset_path}')","metadata":{"_cell_guid":"4cee4d43-e5c6-479d-a98a-4f8a4c505d8e","_uuid":"e5effd00-cb41-4181-99bc-466b3c1f7f26","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:42:04.763441Z","iopub.status.busy":"2026-02-09T07:42:04.762845Z","iopub.status.idle":"2026-02-09T07:42:04.920138Z","shell.execute_reply":"2026-02-09T07:42:04.919324Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.167362,"end_time":"2026-02-09T07:42:04.921604","exception":false,"start_time":"2026-02-09T07:42:04.754242","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pl.scan_parquet(dataset_path / 'train.parquet')\ncolumns = train.collect_schema().names()\nif INTERACTIVE_SESSION:\n    train = train[:1_000_000]","metadata":{"_cell_guid":"0cb83994-beb7-40cd-836e-f3fd10586b44","_uuid":"f76fee83-7e8e-4b67-93f4-f9e4ba7a082d","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:42:04.939043Z","iopub.status.busy":"2026-02-09T07:42:04.938809Z","iopub.status.idle":"2026-02-09T07:42:05.052151Z","shell.execute_reply":"2026-02-09T07:42:05.051618Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.12359,"end_time":"2026-02-09T07:42:05.053515","exception":false,"start_time":"2026-02-09T07:42:04.929925","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create a mask to normalize values to reasonable ranges\nNORM_POS = 5000\nNORM_VEL = 2500\nNORM_ANGVEL = 10\nNORM_BOOST = 100\nNORM_TIMER = 10\n\ndef normalize_df(df, prefix=''):\n    return df.with_columns(\n        pl.col(rf'^{prefix}.*_pos_[xyz]$') / NORM_POS,\n        pl.col(rf'^{prefix}.*_vel_[xyz]$') / NORM_VEL,\n        pl.col(rf'^{prefix}.*_angvel_[xyz]$') / NORM_ANGVEL,\n        # rot and up are already normalized\n        pl.col(rf'^{prefix}.*p\\d_boost$') / NORM_BOOST,\n        pl.col(rf'^{prefix}.*timer$') / NORM_TIMER,\n    )","metadata":{"execution":{"iopub.execute_input":"2026-02-09T07:42:05.070538Z","iopub.status.busy":"2026-02-09T07:42:05.070312Z","iopub.status.idle":"2026-02-09T07:42:05.074771Z","shell.execute_reply":"2026-02-09T07:42:05.074027Z"},"papermill":{"duration":0.014359,"end_time":"2026-02-09T07:42:05.076159","exception":false,"start_time":"2026-02-09T07:42:05.061800","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def trim_columns(lf, columns):\n    columns = [col for col in columns if not col.endswith('10sec') and col not in ('player_scoring_next')]\n    return lf.select(*columns), columns\ntrain, columns = trim_columns(train, columns)\nNUM_PLAYER_COLS = len([col for col in columns if col.startswith('p0_')])\nprint(f'taking {NUM_PLAYER_COLS} player cols')","metadata":{"_cell_guid":"2073c0fe-fa9f-4973-aa8a-7376e3098968","_uuid":"464be3c3-8847-4f32-ad58-6922f305e22a","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:42:05.093429Z","iopub.status.busy":"2026-02-09T07:42:05.093080Z","iopub.status.idle":"2026-02-09T07:42:05.100265Z","shell.execute_reply":"2026-02-09T07:42:05.099558Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.017535,"end_time":"2026-02-09T07:42:05.101556","exception":false,"start_time":"2026-02-09T07:42:05.084021","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def append_modifiers(lf, key, modifiers):\n    for name, size in modifiers.items():\n        lf = lf.with_columns((pl.col(key) % size).alias(name), pl.col(key) // size)\n    return lf\n\n# Sets values for each column in col_list0 to the values from each corresponding column in col_list1, for rows where cond is True\ndef swap(lf, cond, col_list0, col_list1):\n    return (\n        # First \"write\" to a temp column since we often swap x -> y and y -> x together which we can't do simultaneously\n        lf.with_columns(pl.when(cond).then(pl.col(newcol)).otherwise(pl.col(oldcol)).alias(f'__temp_{oldcol}')\n                        for oldcol, newcol in zip(col_list0, col_list1))\n          .with_columns(pl.col(f'__temp_{col}').alias(col) for col in col_list0)\n          .drop('^__temp_.*$')\n    )\n\n# Sets values for each player in p_list0 to the values from each corresponding player in p_list1\ndef swap_players(lf, cond, p_list0, p_list1, columns):\n    def get_cols(p_list):\n        return [col for p in p_list for col in columns if col.startswith(f'p{p}_') or f'_p{p}_' in col]\n    return swap(lf, cond, get_cols(p_list0), get_cols(p_list1))\n\ndef swap_boosts(lf, cond, boost_i_list0, boost_i_list1):\n    def get_cols(boost_i_list):\n        return [f'boost{i}_timer' for i in boost_i_list]\n    return swap(lf, cond, get_cols(boost_i_list0), get_cols(boost_i_list1))\n\ndef augment_data(lf, seed=None, shuffle=False, append_flip_y=False):\n    L = lf.select(pl.len()).collect().item()\n    columns = lf.collect_schema().names()\n    modifiers = {'x_flip': 2, 'y_flip': 2}\n    lf = lf.with_columns(pl.Series(\n        'randoms',\n        np.full(L, seed) if seed is not None else np.random.randint(0, 2**16, size=L),\n        pl.Int32\n    ))\n    if shuffle:\n        lf = lf.sort('randoms')\n    lf = append_modifiers(lf, 'randoms', modifiers)\n    \n    # flip_x\n    flip_x_filter = pl.col('x_flip') == 1\n    lf = lf.with_columns(pl.col(r'^.*_x$').exclude(r'^.*angvel_x$') * pl.when(flip_x_filter).then(-1).otherwise(1))\n    lf = lf.with_columns(pl.col(r'^.*angvel_[yz]$') * pl.when(flip_x_filter).then(-1).otherwise(1))\n    if 'boost0_timer' in columns:\n        # Swap 0 <-> 1, 2 <-> 3, 4 <-> 5\n        lf = swap_boosts(lf, flip_x_filter, list(range(6)), [1, 0, 3, 2, 5, 4])\n    \n    # flip_y -- Also swap teams and which team scored\n    flip_y_filter = pl.col('y_flip') == 1\n    lf = lf.with_columns(pl.col('^.*_y$').exclude(r'^.*angvel_y$') * pl.when(flip_y_filter).then(-1).otherwise(1))\n    lf = lf.with_columns(pl.col(r'^.*angvel_[xz]$') * pl.when(flip_y_filter).then(-1).otherwise(1))\n    lf = swap_players(lf, flip_y_filter, list(range(6)), [3, 4, 5, 0, 1, 2], columns)\n    if 'team_scoring_next' in columns:\n        lf = lf.with_columns(\n            pl.when(flip_y_filter)\n              .then(pl.col('team_scoring_next').neg() + 1)\n              .otherwise(pl.col('team_scoring_next')))\n    if 'player_scoring_next' in columns:\n        lf = lf.with_columns(\n            pl.when(flip_y_filter)\n              .then((pl.col('player_scoring_next') + 3).mod(6))\n              .otherwise(pl.col('player_scoring_next')))\n    team0_targets = [col for col in columns if col.startswith('target_team0_')]\n    team1_targets = [col for col in columns if col.startswith('target_team1_')]\n    if len(team0_targets) > 0:\n        lf = swap(lf, flip_y_filter, team0_targets + team1_targets, team1_targets + team0_targets)\n    if 'boost0_timer' in columns:\n        # Swap 0 <-> 4, 1 <-> 5\n        lf = swap_boosts(lf, flip_y_filter, [0, 1, 4, 5], [4, 5, 0, 1])\n    \n    if append_flip_y:\n        lf = lf.with_columns(flip_y_filter.alias('flip_y'))\n    \n    lf = lf.drop('randoms', *modifiers.keys())\n    \n    return lf","metadata":{"_cell_guid":"7fb1beb0-38f6-4bdd-9b74-054f1a5e46ba","_uuid":"aef39c7d-15c3-43d0-8b82-3a92ffd6ec4f","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:42:05.118862Z","iopub.status.busy":"2026-02-09T07:42:05.118510Z","iopub.status.idle":"2026-02-09T07:42:05.131275Z","shell.execute_reply":"2026-02-09T07:42:05.130531Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.023111,"end_time":"2026-02-09T07:42:05.132702","exception":false,"start_time":"2026-02-09T07:42:05.109591","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Split 10 second target window into buckets for each team, plus a final neither bucket.\n# Apply gaussian smoothing to account for noisiness at longer horizons.\nTARGETS_PER_TEAM = 10\nTOTAL_TARGETS = 2 * TARGETS_PER_TEAM + 1\nMAX_TIME = 10.0\nBUCKET_WIDTH = MAX_TIME / TARGETS_PER_TEAM\n\n# Variable Sigma Logic: Uncertainty grows with time\nSIGMA_BASE = 1e-5\nSIGMA_GROWTH = 0.15\n\ndef generate_geometric_edges(width_ratio):\n    widths = np.array([width_ratio**i for i in range(TARGETS_PER_TEAM)])\n    widths *= MAX_TIME / widths.sum()\n    return np.concatenate([[0], np.cumsum(widths)])\n\n# edges = np.linspace(0, MAX_TIME, TARGETS_PER_TEAM + 1)\n# edges = np.concatenate([np.linspace(0.0, 4.0, 9), np.array([6.0, 8.0, 10.0, 15.0])])\nedges = generate_geometric_edges(width_ratio=1.25)\nassert len(edges) == TARGETS_PER_TEAM + 1\n\nprint(f'using {TARGETS_PER_TEAM} targets per team with {SIGMA_BASE}->{SIGMA_GROWTH} sigma')\nprint(f'bucket edges: {edges}')\n\nTARGET_EPSILON = .1\n\ndef generate_gaussian_targets(time_col, team_col):\n    n_samples = len(time_col)\n    \n    targets = np.zeros((n_samples, TOTAL_TARGETS), dtype=np.float32)\n    targets[:, -1] = 1.0\n    \n    valid_mask = np.isfinite(team_col)\n    t_goal = -time_col[valid_mask]\n    scoring_team = team_col[valid_mask].astype(int)\n    \n    sigma = SIGMA_BASE + (SIGMA_GROWTH * t_goal)\n    edges_bc = edges[None, :]       # (1, T + 1)\n    t_goal_bc = t_goal[:, None]     # (M, 1)\n    sigma_bc = sigma[:, None]       # (M, 1)\n    z_scores = (edges_bc - t_goal_bc) / sigma_bc # (M, T + 1)\n    cdfs = scipy.stats.norm.cdf(z_scores)\n    \n    # Calculate bucket masses and neither mass\n    bucket_probs = cdfs[:, 1:] - cdfs[:, :-1] # (M, T)\n    prob_neither = 1.0 - cdfs[:, -1]          # (M, )\n\n    # Write the probabilities to a new matrix (just for valid rows)\n    valid_targets = np.zeros((len(t_goal), TOTAL_TARGETS), dtype=np.float32)\n    t0_mask = scoring_team == 0\n    valid_targets[t0_mask, :TARGETS_PER_TEAM] = bucket_probs[t0_mask]\n    valid_targets[t0_mask, -1] = prob_neither[t0_mask]\n    t1_mask = scoring_team == 1\n    valid_targets[t1_mask, TARGETS_PER_TEAM:2*TARGETS_PER_TEAM] = bucket_probs[t1_mask]\n    valid_targets[t1_mask, -1] = prob_neither[t1_mask]\n    \n    # Set anything smaller than epsilon to 0.0 and renormalize to sum 1.0\n    valid_targets[valid_targets < TARGET_EPSILON] = 0.0\n    valid_targets = valid_targets / valid_targets.sum(axis=1, keepdims=True)\n    \n    # Place valid rows' targets back into the main targets array\n    targets[valid_mask] = valid_targets\n    target_names = (\n        [f'target_team0_{i:02d}' for i in range(TARGETS_PER_TEAM)] +\n        [f'target_team1_{i:02d}' for i in range(TARGETS_PER_TEAM)] +\n        ['target_neither']\n    )\n\n    return targets, target_names\n\ndef get_train_targets(df):\n    time_col = df['event_time'].to_numpy()\n    team_col = df['team_scoring_next'].to_numpy()\n\n    return generate_gaussian_targets(time_col, team_col)","metadata":{"_cell_guid":"d6b2cdd7-cbe2-49e6-9d81-8bf02a00eccb","_uuid":"43fa8101-0e87-4e8f-89da-373c5f6ed7ea","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:42:05.149808Z","iopub.status.busy":"2026-02-09T07:42:05.149545Z","iopub.status.idle":"2026-02-09T07:42:05.160285Z","shell.execute_reply":"2026-02-09T07:42:05.159513Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.02117,"end_time":"2026-02-09T07:42:05.161768","exception":false,"start_time":"2026-02-09T07:42:05.140598","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"FPS = 10\ndef append_physics_targets(df, millis=500):\n    shift = millis * FPS // 1000\n    return df.with_columns(\n        pl.col(\"^ball_.*$\")\n          .shift(-shift)\n          .over('event_id')\n          .name.prefix(f'target_future{millis}ms_')\n    )\n\ndef append_player_physics_targets(df, millis=500):\n    shift = millis * FPS // 1000\n    \n    new_cols = []\n    for p in range(6):\n        # Determine if player is CURRENTLY demoed (mask input) or FUTURE demoed (target)\n        # Assuming 'p{p}_demoed' exists (0=Alive, 1=Dead)\n        curr_demo =  pl.col(f'p{p}_pos_x').is_nan()\n        fut_demo = curr_demo.shift(-shift).over('event_id')\n        \n        # 1. Movement Targets (6 cols)\n        # If currently dead OR future dead, set movement target to NULL (don't train regression)\n        # We only want to predict movement if they survive.\n        mov_cols = ['pos_x', 'pos_y', 'pos_z', 'vel_x', 'vel_y', 'vel_z']\n        for feat in mov_cols:\n            raw_val = pl.col(f'p{p}_{feat}').shift(-shift).over('event_id')\n            masked_val = pl.when(curr_demo | fut_demo).then(None).otherwise(raw_val)\n            new_cols.append(masked_val.alias(f'target_future{millis}ms_p{p}_{feat}'))\n            \n        # 2. Survival Target (1 col)\n        # This trains the model to predict \"Will I die / respawn?\"\n        new_cols.append(fut_demo.cast(pl.Float32).alias(f'target_future{millis}ms_p{p}_is_demoed'))\n\n    return df.with_columns(*new_cols)","metadata":{"execution":{"iopub.execute_input":"2026-02-09T07:42:05.179354Z","iopub.status.busy":"2026-02-09T07:42:05.179154Z","iopub.status.idle":"2026-02-09T07:42:05.185042Z","shell.execute_reply":"2026-02-09T07:42:05.184460Z"},"papermill":{"duration":0.015889,"end_time":"2026-02-09T07:42:05.186348","exception":false,"start_time":"2026-02-09T07:42:05.170459","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 16384  # The Effective Batch Size (what the optimizer sees)\nDEVICE_MICRO_BATCH_SIZE = 1024 # Size of each micro-batch on each device\nNUM_DEVICES = len(jax.local_devices())\n\nassert BATCH_SIZE % (DEVICE_MICRO_BATCH_SIZE * NUM_DEVICES) == 0\nACCUM_STEPS = BATCH_SIZE // (DEVICE_MICRO_BATCH_SIZE * NUM_DEVICES) # How many micro-batches to sum before updating\n\nprint(f'Effective Batch Size: {BATCH_SIZE}')\nprint(f'Micro Batch Size (Per GPU): {DEVICE_MICRO_BATCH_SIZE}')\nprint(f'Accum Steps: {ACCUM_STEPS}')","metadata":{"_cell_guid":"09a5e374-4de8-4d38-bd48-ede926deab49","_uuid":"690fda62-3f2d-456a-a65d-6e723075702b","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:42:05.203154Z","iopub.status.busy":"2026-02-09T07:42:05.202727Z","iopub.status.idle":"2026-02-09T07:42:05.207164Z","shell.execute_reply":"2026-02-09T07:42:05.206408Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.014354,"end_time":"2026-02-09T07:42:05.208542","exception":false,"start_time":"2026-02-09T07:42:05.194188","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\nFRACTION_VAL = 0.05\n\nfull_train_df = train.collect()\ngame_nums = sorted(full_train_df['game_num'].unique())\nval_game_nums = set(game_nums[::len(game_nums) // int(FRACTION_VAL * len(game_nums))])\nval_game_expr = pl.col('game_num').is_in(val_game_nums)\ntrain_df = full_train_df.filter(~val_game_expr).drop('game_num')\nval_df = full_train_df.filter(val_game_expr).drop('game_num', 'event_id')\ncolumns.remove('game_num')\ndel full_train_df","metadata":{"_cell_guid":"123f69e1-0520-4a1f-ab41-91ac98e67231","_uuid":"8db26152-2004-4e63-8478-91403601feb2","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:42:05.225429Z","iopub.status.busy":"2026-02-09T07:42:05.225205Z","iopub.status.idle":"2026-02-09T07:42:46.942471Z","shell.execute_reply":"2026-02-09T07:42:46.941649Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":41.735549,"end_time":"2026-02-09T07:42:46.952063","exception":false,"start_time":"2026-02-09T07:42:05.216514","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def append_train_targets(df, columns):\n    targets, target_names = get_train_targets(df)\n    return df.with_columns(\n        *[pl.Series(name, target) for target, name in zip(targets.T, target_names)]\n    )\n\nTRAIN_STRIDE = 8\ndef slice_train_stride(df, epoch):\n    epoch = epoch % TRAIN_STRIDE\n    return df.gather_every(n=TRAIN_STRIDE, offset=epoch)\n\ntarget_filter = pl.col('^target_.*$')\ndef get_train_xy(df, batch_size):\n    df = df[:len(df) // batch_size * batch_size]\n    \n    features = df.drop(target_filter).to_numpy()\n    targets = df.select(target_filter).to_numpy()\n\n    num_batches = len(df) // batch_size\n    \n    def shard(data):\n        return data.reshape(\n            num_batches,\n            NUM_DEVICES,\n            ACCUM_STEPS,\n            DEVICE_MICRO_BATCH_SIZE,\n            *data.shape[1:]\n        )\n    \n    return shard(features), shard(targets)","metadata":{"_cell_guid":"b9f6341c-3647-4f03-8c9b-10dabc00d2b3","_uuid":"151ae74f-26a9-4de4-b020-40d03b30b2a6","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:42:46.969573Z","iopub.status.busy":"2026-02-09T07:42:46.969062Z","iopub.status.idle":"2026-02-09T07:42:46.975286Z","shell.execute_reply":"2026-02-09T07:42:46.974567Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.016601,"end_time":"2026-02-09T07:42:46.976708","exception":false,"start_time":"2026-02-09T07:42:46.960107","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"temp_dir = tempfile.TemporaryDirectory()\natexit.register(temp_dir.cleanup)\ntemp_dir_path = pathlib.Path(temp_dir.name)\nprint(f'using tempdir {str(temp_dir_path)}')","metadata":{"_cell_guid":"68c678af-676d-4527-bce7-359f0c91cbf6","_uuid":"885aba48-b5c4-4adb-9459-5eadfdd0c47f","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:42:46.993524Z","iopub.status.busy":"2026-02-09T07:42:46.993284Z","iopub.status.idle":"2026-02-09T07:42:46.997750Z","shell.execute_reply":"2026-02-09T07:42:46.996983Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.014478,"end_time":"2026-02-09T07:42:46.999132","exception":false,"start_time":"2026-02-09T07:42:46.984654","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = append_train_targets(train_df, columns)\ntrain = train_df.lazy()\ntrain = append_physics_targets(train, millis=200)\ntrain = append_physics_targets(train, millis=1500)\ntrain = append_player_physics_targets(train, millis=1500)\ntrain = normalize_df(train, prefix='target_future')\n\ncols_to_drop = ['event_time', 'team_scoring_next', 'event_id']\ntrain = train.drop(*cols_to_drop)\ncolumns = [col for col in columns if col not in cols_to_drop]\n\ntrain_path = temp_dir_path / 'train.parquet'\ntrain.sink_parquet(train_path, compression='lz4')\nprint('wrote train')\n\ntrain_facade = train.clear(1).collect()\n\ndel train_df\n\nNUM_BALL_PHYSICS_TARGETS = len([col for col in train_facade.columns if col.startswith('target_future') and 'ball' in col])\nNUM_PLAYER_PHYSICS_TARGETS = len([col for col in train_facade.columns if col.startswith('target_future') and 'p0_' in col])\nprint(f'got {NUM_BALL_PHYSICS_TARGETS} physics targets {NUM_PLAYER_PHYSICS_TARGETS} players')\n\nnormalizer_df = train_facade.drop(target_filter).select(pl.all().fill_null(pl.lit(1.0, dtype=pl.Float32)))\nnormalizer_df = normalize_df(normalizer_df)\ninput_normalizer = normalizer_df.to_numpy()[0, :]\n\ntrain = pl.scan_parquet(train_path)\ndef get_augmented_train_xy(epoch):\n    strided = slice_train_stride(train, epoch)\n    train_df = augment_data(strided, shuffle=True).collect()\n    features, targets = get_train_xy(train_df, BATCH_SIZE)\n    del train_df\n    return features, targets\n\ntrain_features, train_targets = get_augmented_train_xy(0)\nprint(f'Train features: {train_features.shape} train targets {train_targets.shape}')","metadata":{"_cell_guid":"93a46f96-5ec3-424e-b2ee-8455dbdf7c48","_uuid":"21d6348e-d0c8-41f4-bbda-852942f0daac","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:42:47.016185Z","iopub.status.busy":"2026-02-09T07:42:47.015942Z","iopub.status.idle":"2026-02-09T07:46:33.558100Z","shell.execute_reply":"2026-02-09T07:46:33.555247Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":226.563963,"end_time":"2026-02-09T07:46:33.571183","exception":false,"start_time":"2026-02-09T07:42:47.007220","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_val_xy(df):\n    features = df.drop('event_time', 'team_scoring_next').to_numpy()\n    targets_sparse = df.select(\n        pl.when(pl.col('team_scoring_next').is_null() | (pl.col('event_time') < -10))\n                 .then(2)\n                 .otherwise(pl.col('team_scoring_next'))).to_series().to_numpy()\n    targets = np.array(jnn.one_hot(targets_sparse, 3))\n\n    return features, targets","metadata":{"_cell_guid":"1a760927-bc96-43be-9c3d-787595ed795d","_uuid":"6b77abb0-24e5-4c63-a87e-7e838bdcaa82","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:46:33.589155Z","iopub.status.busy":"2026-02-09T07:46:33.588900Z","iopub.status.idle":"2026-02-09T07:46:33.594130Z","shell.execute_reply":"2026-02-09T07:46:33.593585Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.016043,"end_time":"2026-02-09T07:46:33.595330","exception":false,"start_time":"2026-02-09T07:46:33.579287","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"_, val_targets = get_val_xy(val_df)\nval_features_list = []\nval_flip_y_list = []\nval = val_df.lazy()\nfor i in [0, 3]:\n    val_df = augment_data(val, seed=i, shuffle=False, append_flip_y=True).collect()\n    features, _ = get_val_xy(val_df.drop('flip_y'))\n    val_features_list.append(features)\n    val_flip_y_list.append(val_df['flip_y'][:len(features)].to_numpy())\n    del val_df, features\nnum_val_rows = np.prod(val_features_list[0].shape[:-1])\nassert num_val_rows > 0\nprint(f'Validation rows: {num_val_rows}')","metadata":{"_cell_guid":"aef5174c-3679-4e53-81a4-ba9619ff4c57","_uuid":"5b30e20a-d57d-4305-93a1-4d09046d3ac3","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:46:33.612383Z","iopub.status.busy":"2026-02-09T07:46:33.612178Z","iopub.status.idle":"2026-02-09T07:46:35.804675Z","shell.execute_reply":"2026-02-09T07:46:35.803846Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":2.202614,"end_time":"2026-02-09T07:46:35.806091","exception":false,"start_time":"2026-02-09T07:46:33.603477","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BOOST_LOCS = np.array([\n    [-3072, -4096, 72], [3072, -4096, 72],\n    [-3584,     0, 72], [3584,     0, 72],\n    [-3072,  4096, 72], [3072,  4096, 72],\n], dtype=np.float32) / NORM_POS\n\nGOAL_LOCS = np.array([\n    [0, -5216, 318],\n    [0,  5216, 318],\n], dtype=np.float32) / NORM_POS","metadata":{"_cell_guid":"5a1de25d-6105-480d-9a3a-fd51cf0aa34e","_uuid":"3f72f829-0479-4078-af57-128598c73296","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:46:35.824936Z","iopub.status.busy":"2026-02-09T07:46:35.824715Z","iopub.status.idle":"2026-02-09T07:46:35.830230Z","shell.execute_reply":"2026-02-09T07:46:35.829664Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.016631,"end_time":"2026-02-09T07:46:35.831450","exception":false,"start_time":"2026-02-09T07:46:35.814819","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"role_indices = np.array([0, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 3, 3], dtype=np.int32) # Goal, Player x 6, Boost x 6, Goal x 2\nteam_indices = np.array([-1, 0, 0, 0, 1, 1, 1, -1, -1, -1, -1, -1, -1, 0, 1], dtype=np.int32) # Players and goals have team 0 or 1, else -1\n\nrole_indices_1 = np.repeat(np.expand_dims(role_indices, axis=0), 15, axis=0)\nrole_indices_2 = np.repeat(np.expand_dims(role_indices, axis=1), 15, axis=1)\nteam_indices_1 = np.repeat(np.expand_dims(team_indices, axis=0), 15, axis=0)\nteam_indices_2 = np.repeat(np.expand_dims(team_indices, axis=1), 15, axis=1)\n\n# Build a lot of one-hots to encode interesting pairings the model may want to attend to\nis_self = np.eye(15, dtype=np.int32) == 1 # (15, 15)\nhas_team = (team_indices_1 >= 0) & (team_indices_2 >= 0)\nis_friend = has_team & (team_indices_1 == team_indices_2)\nis_foe = has_team & (team_indices_1 != team_indices_2)\nball_player = (role_indices_1 == 0) & (role_indices_2 == 1)\nball_goal = (role_indices_1 == 0) & (role_indices_2 == 3)\nplayer_ball = (role_indices_1 == 1) & (role_indices_2 == 0)\nplayer_teammate = (role_indices_1 == 1) & (role_indices_2 == 1) & is_friend & ~is_self\nplayer_opponent = (role_indices_1 == 1) & (role_indices_2 == 1) & is_foe\nplayer_boost = (role_indices_1 == 1) & (role_indices_2 == 2)\nplayer_own_goal = (role_indices_1 == 1) & (role_indices_2 == 3) & is_friend\nplayer_opponent_goal = (role_indices_1 == 1) & (role_indices_2 == 3) & is_foe\nboost_player = (role_indices_1 == 2) & (role_indices_2 == 1)\ngoal_ball = (role_indices_1 == 3) & (role_indices_2 == 0)\ngoal_teammate = (role_indices_1 == 3) & (role_indices_2 == 1) & is_friend\ngoal_opponent = (role_indices_1 == 3) & (role_indices_2 == 1) & is_foe\nboth_static = (role_indices_1 >= 2) & (role_indices_2 >= 2)\n\npairwise_type_flags = np.stack([\n    role_indices_1,\n    role_indices_2,\n    team_indices_1,\n    team_indices_2,\n    is_self,\n    has_team,\n    is_friend,\n    is_foe,\n    ball_player,\n    ball_goal,\n    player_ball,\n    player_teammate,\n    player_opponent,\n    player_boost,\n    player_own_goal,\n    player_opponent_goal,\n    boost_player,\n    goal_ball,\n    goal_teammate,\n    goal_opponent,\n    both_static,\n], axis=-1)\npairwise_type_flags = pairwise_type_flags.astype(np.float32)","metadata":{"execution":{"iopub.execute_input":"2026-02-09T07:46:35.849329Z","iopub.status.busy":"2026-02-09T07:46:35.849121Z","iopub.status.idle":"2026-02-09T07:46:35.864006Z","shell.execute_reply":"2026-02-09T07:46:35.863153Z"},"papermill":{"duration":0.025564,"end_time":"2026-02-09T07:46:35.865411","exception":false,"start_time":"2026-02-09T07:46:35.839847","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"D_MODEL = 32 if INTERACTIVE_SESSION else 64\nNUM_HEADS = 2 if INTERACTIVE_SESSION else 4\nNUM_LAYERS = 2 if INTERACTIVE_SESSION else 16\nKEY_DIM = D_MODEL // NUM_HEADS\nFF_DIM = D_MODEL * 4\nNECK_DIM = 256\nPHYSICS_NECK_DIM = 64\nDROPOUT = 0.1\nMAX_DROP_PATH = 0.25\nNECK_DROPOUT = 0.4\n\n@keras.saving.register_keras_serializable()\ndef gelu_approx(x):\n    return keras.activations.gelu(x, approximate=True)\nACTIVATION = gelu_approx\n\n# Use our own manual Multi-head Attention implementation which enables the extension of pairwise-bias\n# float values across each pair of entities.\n# The default keras.layers.MultiHeadAttention layer only supports boolean attention_mask, which we're\n# using here as well (for demolished players) but also want richer pairwise information like distances.\ndef manual_multihead_attention(q, k, v, num_heads, key_dim, dropout, bias=None):\n    \"\"\"\n    Implements Scaled Dot-Product Attention using EinsumDense.\n    Preserves static sequence lengths to ensure downstream shape inference works.\n    \"\"\"\n    # Capture static sequence lengths (T for query, S for key/value)\n    # If the shape is dynamic (None), it stays None.\n    T = q.shape[1]\n    S = k.shape[1]\n    \n    # --- 1. Projections (Project + Split Heads) ---\n    # Input shape:  (Batch, SeqLen, D_Model) -> \"abc\"\n    # Kernel shape: (D_Model, NumHeads, KeyDim) -> \"cde\"\n    # Output shape: (Batch, SeqLen, NumHeads, KeyDim) -> \"abde\"\n    \n    proj_equation = \"abc,cde->abde\"\n    \n    # Query Projection\n    query = keras.layers.EinsumDense(\n        equation=proj_equation,\n        output_shape=(T, num_heads, key_dim),\n        bias_axes=\"de\"\n    )(q)\n    \n    # Key Projection\n    key = keras.layers.EinsumDense(\n        equation=proj_equation,\n        output_shape=(S, num_heads, key_dim),\n        bias_axes=\"de\"\n    )(k)\n    \n    # Value Projection\n    value = keras.layers.EinsumDense(\n        equation=proj_equation,\n        output_shape=(S, num_heads, key_dim),\n        bias_axes=\"de\"\n    )(v)\n\n    # --- 2. Transpose for Attention: (B, T, H, K) -> (B, H, T, K) ---\n    query = keras.ops.transpose(query, (0, 2, 1, 3))\n    key   = keras.ops.transpose(key,   (0, 2, 1, 3))\n    value = keras.ops.transpose(value, (0, 2, 1, 3))\n    \n    # --- 3. Dot Product Attention Scores ---\n    # (B, H, T, K) @ (B, H, K, S) -> (B, H, T, S)\n    key_T = keras.ops.transpose(key, (0, 1, 3, 2))\n    scores = keras.ops.matmul(query, key_T)\n    \n    # --- 4. Scale ---\n    scale = keras.ops.cast(keras.ops.sqrt(float(key_dim)), scores.dtype)\n    scores = scores / scale\n\n    if bias is not None:\n        scores = keras.layers.Add()([scores, bias])\n    \n    # --- 6. Softmax & Dropout ---\n    weights = keras.ops.softmax(scores, axis=-1)\n    if dropout > 0:\n        weights = keras.layers.Dropout(dropout)(weights)\n        \n    # --- 7. Weighted Sum ---\n    # (B, H, T, S) @ (B, H, S, K) -> (B, H, T, K)\n    attention = keras.ops.matmul(weights, value)\n    \n    # --- 8. Transpose back: (B, H, T, K) -> (B, T, H, K) ---\n    attention = keras.ops.transpose(attention, (0, 2, 1, 3))\n    \n    # --- 9. Output Projection (Merge Heads) ---\n    # Input:  (Batch, T, NumHeads, KeyDim) -> \"abcd\"\n    # Kernel: (NumHeads, KeyDim, D_Model) -> \"cde\"\n    # Output: (Batch, T, D_Model) -> \"abe\"\n    \n    output_dim = q.shape[-1]\n    output = keras.layers.EinsumDense(\n        equation=\"abcd,cde->abe\",\n        output_shape=(T, output_dim),\n        bias_axes=\"e\",\n        kernel_initializer=\"zeros\"\n    )(attention)\n    \n    return output\n\ndef mlp_block(hidden_units, dropout=0.0, final_layer_init_zero=False):\n    layers = []\n    for units in hidden_units[:-1]:\n        layers.append(keras.layers.Dense(units, activation=ACTIVATION))\n    layers.append(keras.layers.Dense(hidden_units[-1], activation=None,\n                                     kernel_initializer='zeros' if final_layer_init_zero else 'glorot_uniform'))\n    if dropout > 0:\n        layers.append(keras.layers.Dropout(dropout))\n\n    def apply(x):\n        for layer in layers:\n            x = layer(x)\n        return x\n\n    return apply\n\ndef geglu_block(x, ff_dim, dropout=0.0, reshape_back=True):\n    \"\"\"\n    Gated Linear Unit (GEGLU) block.\n    Splits the projection into a gating branch and a value branch.\n    x -> (xW + b) * GELU(xV + c) -> Output\n    \"\"\"\n    ff = keras.layers.Dense(ff_dim * 2)(x)\n    val, gate = keras.ops.split(ff, 2, axis=-1)\n    x_ff = keras.layers.Multiply()([val, ACTIVATION(gate)])\n\n    if reshape_back:\n        x_ff = keras.layers.Dense(x.shape[-1], kernel_initializer='zeros')(x_ff)\n    if dropout > 0:\n        x_ff = keras.layers.Dropout(dropout)(x_ff)\n        \n    return x_ff\n\ndef apply_stochastic_depth(x, drop_path_rate):\n    if drop_path_rate == 0:\n        return x\n\n    ones_mask = keras.ops.ones_like(x[:, :1, :1])\n    mask = keras.layers.Dropout(drop_path_rate)(ones_mask)\n\n    return keras.layers.Multiply()([x, mask])\n\ndef pre_norm_transformer_block(x, bias, num_heads, key_dim, ff_dim, dropout, drop_path_rate=0.0):\n    # --- Sub-Layer 1: Self-Attention ---\n    skip_1 = x\n    x_norm = keras.layers.LayerNormalization()(x)\n    attn = manual_multihead_attention(\n        x_norm, x_norm, x_norm, num_heads, key_dim, dropout, bias=bias\n    )\n    attn = keras.layers.Dropout(dropout)(attn)\n    attn = apply_stochastic_depth(attn, drop_path_rate)\n    x = keras.layers.Add()([skip_1, attn])\n    \n    # --- Sub-Layer 2: Feed Forward ---\n    skip_2 = x\n    x_norm = keras.layers.LayerNormalization()(x)\n    x_ff = geglu_block(x_norm, ff_dim, dropout)\n    x_ff = apply_stochastic_depth(x_ff, drop_path_rate)\n    x = keras.layers.Add()([skip_2, x_ff])\n    \n    return x\n\ndef norm(x, axis=-1):\n    return keras.ops.sqrt(keras.ops.sum(keras.ops.square(x), axis=axis, keepdims=True) + 1e-4)\n\ndef dot(x, y, axis=-1):\n    return keras.ops.sum(x * y, axis=axis, keepdims=True)\n\ndef cross_last_axis(x1, x2):\n    assert x1.shape[-1] == 3 and x2.shape[-1] == 3\n    x1_x, x1_y, x1_z = x1[..., 0:1], x1[..., 1:2], x1[..., 2:3]\n    x2_x, x2_y, x2_z = x2[..., 0:1], x2[..., 1:2], x2[..., 2:3]\n    cross_x = x1_y * x2_z - x1_z * x2_y\n    cross_y = x1_z * x2_x - x1_x * x2_z\n    cross_z = x1_x * x2_y - x1_y * x2_x\n    return keras.ops.concatenate([cross_x, cross_y, cross_z], axis=-1)\n\ninputs = keras.Input(shape=(len(columns),), name='inputs')\n\ninput_normalizer_tensor = keras.ops.convert_to_tensor(input_normalizer) # (1, N)\ninputs_normed = inputs * input_normalizer_tensor\n\nidx_ball_end = 6\nidx_players_end = idx_ball_end + (6 * NUM_PLAYER_COLS)\nidx_boosts_end = idx_players_end + 6\nassert idx_boosts_end == len(columns)\n\nplayer_alive = keras.ops.cast(keras.ops.isfinite(inputs_normed[:, idx_ball_end:idx_players_end:NUM_PLAYER_COLS]), dtype='float32')[:, :, None] # (Batch, 6, 1)\ninputs_normed = keras.ops.nan_to_num(inputs_normed, 0)\n\nball_raw = inputs_normed[:, :6]\nplayers_raw = keras.ops.reshape(inputs_normed[:, 6:idx_players_end], (-1, 6, NUM_PLAYER_COLS))\nboost_timers = inputs_normed[:, idx_players_end:idx_boosts_end]\n\nball_expanded = keras.ops.expand_dims(ball_raw, axis=1)\n\n# Feature engineering for player vs ball information\n# Make sure to mask out engineering features for demoed players.\nplayer_to_ball = ball_expanded - players_raw[:, :, :6]\nball_dists = norm(player_to_ball[:, :, :3], axis=2)\nplayer_to_ball_norms = player_to_ball[:, :, :3] / ball_dists\nplayer_vel_to_ball_dots = dot(player_to_ball_norms, players_raw[:, :, 3:6])\nplayer_rot_to_ball_dots = dot(player_to_ball_norms, players_raw[:, :, 6:9])\nplayer_up_to_ball_dots = dot(player_to_ball_norms, players_raw[:, :, 9:12])\n\np_fwd = players_raw[:, :, 6:9]\np_up  = players_raw[:, :, 9:12]\np_right = cross_last_axis(p_fwd, p_up) # Right vector\np_angvel = players_raw[:, :, 12:15]    # Raw World AngVel\n\n# 2. Project to Local (Roll/Pitch/Yaw)\n# This creates invariance: A front flip is always positive pitch_rate\nroll_rate  = dot(p_angvel, p_fwd)\npitch_rate = dot(p_angvel, p_right)\nyaw_rate   = dot(p_angvel, p_up)\nangvel_mag = norm(p_angvel, axis=2)\nnose_vel = cross_last_axis(p_angvel, p_fwd)\nturn_to_ball_dots = dot(nose_vel, player_to_ball_norms)\n\nplayer_features = keras.ops.concatenate([\n        players_raw,\n        player_alive,\n        player_to_ball,\n        ball_dists,\n        player_vel_to_ball_dots,\n        player_rot_to_ball_dots,\n        player_up_to_ball_dots,\n        roll_rate,\n        pitch_rate,\n        yaw_rate,\n        angvel_mag,\n        turn_to_ball_dots,\n    ], axis=2)\n\nboost_timers_expanded = keras.ops.expand_dims(boost_timers, axis=2) # (Batch, 6, 1)\nboost_zeros_batch = keras.ops.zeros_like(boost_timers_expanded) # (Batch, 6, 1)\nboost_locs_const = keras.ops.expand_dims(BOOST_LOCS, axis=0) # (1, 6, 3)\n# Broadcast add: (Batch, 6, 1) + (1, 6, 3) = (Batch, 6, 3)\nboost_pos = keras.layers.Add()([boost_zeros_batch, boost_locs_const])\n# Append z=0\nboost_active = keras.ops.cast(boost_timers_expanded > -1e-5, dtype='float32')\n\nboost_features = keras.ops.concatenate([boost_timers_expanded, boost_active], axis=2)\n\ngoal_zeros_batch = keras.ops.zeros_like(boost_timers_expanded[:, :2, :1]) # (Batch, 2, 1)\ngoal_locs_const = keras.ops.expand_dims(GOAL_LOCS, axis=0) # (1, 2, 3)\n# Broadcast add: (Batch, 2, 1) + (1, 2, 3) = (Batch, 2, 3)\ngoal_pos = keras.layers.Add()([goal_zeros_batch, goal_locs_const])\n\nball_encoder = mlp_block([D_MODEL, D_MODEL])\nplayer_encoder = mlp_block([D_MODEL, D_MODEL])\nboost_encoder = mlp_block([D_MODEL])\nrole_emb = keras.layers.Embedding(4, D_MODEL) # 0:Ball, 1:Player, 2:Boost, 3:Goal\nteam_emb = keras.layers.Embedding(2, D_MODEL) # 0:Team0, 1:Team1\nboost_emb = keras.layers.Embedding(6, D_MODEL)\n\n# Ball\nball_pos = ball_raw[:, :3]\nidx_ball = keras.ops.zeros_like(ball_pos[:, 0], dtype='int32')\nball_token = ball_encoder(ball_raw) + role_emb(idx_ball)\nball_token = keras.ops.expand_dims(ball_token, axis=1)\n\n# Players\nplayer_pos = players_raw[:, :, :3]\nidx_players = keras.ops.ones_like(players_raw[:, :, 0], dtype='int32')\nidx_t0 = keras.ops.zeros_like(player_pos[:, :3, 0], dtype='int32')\nidx_t1 = keras.ops.ones_like(player_pos[:, :3, 0], dtype='int32')\nidx_teams = keras.ops.concatenate([idx_t0, idx_t1], axis=1)\nplayer_tokens = player_encoder(player_features) + role_emb(idx_players) + team_emb(idx_teams)\n\n# Boosts\nidx_boost = keras.ops.ones_like(boost_pos[:, :, 0], dtype='int32') * 2\nidx_boost_team = keras.ops.zeros_like(boost_pos[:, :, 0], dtype='int32') * 3\nboost_indices = keras.ops.arange(6)\nboost_tokens = boost_encoder(boost_features) + role_emb(idx_boost) + boost_emb(boost_indices)\n\n# Goals\nidx_goal = keras.ops.ones_like(goal_pos[:, :, 0], dtype='int32') * 3\nidx_goal0 = keras.ops.zeros_like(goal_pos[:, :1, 0], dtype='int32')\nidx_goal1 = keras.ops.ones_like(goal_pos[:, :1, 0], dtype='int32')\nidx_goal_team = keras.ops.concatenate([idx_goal0, idx_goal1], axis=1)\ngoal_tokens = role_emb(idx_goal) + team_emb(idx_goal_team)\n\nx = keras.ops.concatenate([ball_token, player_tokens, boost_tokens, goal_tokens], axis=1) # (Batch, 15, D)\nx = keras.layers.LayerNormalization()(x) # (Batch, 15, D)\n\n# (Batch, 1)\nball_alive = keras.ops.ones_like(ball_raw[:, :1], dtype='bool')\n# (Batch, 6)\nplayer_alive = keras.ops.cast(player_alive[:, :, 0], 'bool')\n# player_alive = keras.ops.ones_like(player_alive[:, :, 0], 'bool')\n# (Batch, 6)\nboost_alive = keras.ops.ones_like(ball_raw[:, :6], dtype='bool')\n# (Batch, 2)\ngoal_alive = keras.ops.ones_like(ball_raw[:, :2], dtype='bool')\n# Concatenate validity: [Ball, P0..P5, B0..B5] -> (Batch, 15)\nvalid_tokens = keras.ops.concatenate([ball_alive, player_alive, boost_alive, goal_alive], axis=1)\n\nall_pos = keras.ops.concatenate([ball_expanded[:, :, :3], player_pos, boost_pos, goal_pos], axis=1) # (Batch, 15, 3)\nall_pos_diffs = keras.ops.expand_dims(all_pos, axis=2) - keras.ops.expand_dims(all_pos, axis=1) # (Batch, 15, 15, 3)\nall_pos_dists = norm(all_pos_diffs, axis=3) # (Batch, 15, 15, 1)\nall_pos_diffs_normed = all_pos_diffs / all_pos_dists # (Batch, 15, 15, 3)\nall_pos_diffs_normed_t = -keras.ops.transpose(all_pos_diffs_normed, (0, 2, 1, 3)) # (Batch, 15, 15, 3)\n\nall_vel = keras.ops.concatenate([ball_expanded[:, :, 3:6], players_raw[:, :, 3:6], keras.ops.zeros_like(boost_pos), keras.ops.zeros_like(goal_pos)], axis=1) # (Batch, 15, 3)\nall_vel_1 = all_vel[:, None, :, :] # (Batch, 1, 15, 3)\nall_vel_2 = all_vel[:, :, None, :] # (Batch, 15, 1, 3)\nvel_dots_1 = dot(all_vel_1, all_pos_diffs_normed) # (Batch, 15, 15, 1)\nvel_dots_2 = dot(all_vel_2, all_pos_diffs_normed_t) # (Batch, 15, 15, 1)\nall_vel_diffs = all_vel_1 - all_vel_2 # (Batch, 15, 15, 3)\nclosing_speed = dot(all_vel_diffs, all_pos_diffs_normed) # (Batch, 15, 15, 1)\n\nball_rot = ball_expanded[:, :, 3:6] / norm(ball_expanded[:, :, 3:6], axis=2)\nall_rot = keras.ops.concatenate([ball_rot, players_raw[:, :, 6:9], keras.ops.zeros_like(boost_pos), keras.ops.zeros_like(goal_pos)], axis=1) # (Batch, 15, 3)\nall_rot_1 = all_rot[:, None, :, :] # (Batch, 1, 15, 3)\nall_rot_2 = all_rot[:, :, None, :] # (Batch, 15, 1, 3)\nrot_pos_diff_dot_1 = dot(all_rot_1, all_pos_diffs_normed) # (Batch, 15, 15, 1)\nrot_pos_diff_dot_2 = dot(all_rot_2, all_pos_diffs_normed_t) # (Batch, 15, 15, 1)\n\nall_up = keras.ops.concatenate([keras.ops.zeros_like(ball_rot), players_raw[:, :, 6:9], keras.ops.zeros_like(boost_pos), keras.ops.zeros_like(goal_pos)], axis=1) # (Batch, 15, 3)\nall_up_1 = all_up[:, None, :, :] # (Batch, 1, 15, 3)\nall_up_2 = all_up[:, :, None, :] # (Batch, 15, 1, 3)\nup_pos_diff_dot_1 = dot(all_up_1, all_pos_diffs_normed) # (Batch, 15, 15, 1)\nup_pos_diff_dot_2 = dot(all_up_2, all_pos_diffs_normed_t) # (Batch, 15, 15, 1)\nbatch_zero_anchor = keras.ops.zeros_like(all_pos_diffs_normed[:, :, :, :1])\nall_up_z_1 = batch_zero_anchor + all_up_1[:, :, :, 2:3]\nall_up_z_2 = batch_zero_anchor + all_up_2[:, :, :, 2:3]\n\nall_alive = valid_tokens[:, :, None] # (Batch, 15, 1)\nbatch_zero_anchor = keras.ops.zeros_like(all_pos_diffs_normed[:, :, :, :1])\nall_alive_1 = all_alive[:, None, :, :]\nall_alive_2 = all_alive[:, :, None, :]\nboth_alive = keras.ops.cast(keras.ops.logical_and(all_alive_1, all_alive_2), 'float32')\n\nall_angvel = keras.ops.concatenate([keras.ops.zeros_like(ball_expanded[:, :, :3]), p_angvel, keras.ops.zeros_like(boost_pos), keras.ops.zeros_like(goal_pos)], axis=1)\n\n# 2. Calculate \"Nose Velocity\" due to rotation\n# V_nose = Omega x Forward\n# This vector points in the direction the nose is currently moving\nnose_vel_due_to_rot = cross_last_axis(all_angvel, all_rot) # (Batch, 15, 3)\n\n# 3. Pair and Dot\n# We want to know: Is nose_vel aligned with the direction to the other entity?\nnose_vel_paired = keras.ops.expand_dims(nose_vel_due_to_rot, axis=1) # (Batch, 1, 15, 3)\n\n# \"Is Entity A turning such that its nose is moving toward Entity B?\"\nturn_towards_dot_1 = dot(nose_vel_paired, all_pos_diffs_normed)\n\n# \"Is Entity B turning such that its nose is moving toward Entity A?\"\n# Note: We use all_pos_diffs_normed_t here, which points B -> A\n# We re-use nose_vel_paired but implicitly broadcast against the transposed dimension\nturn_towards_dot_2 = dot(nose_vel_paired, all_pos_diffs_normed_t)\n\npairwise_type_flags_tensor = keras.ops.convert_to_tensor(pairwise_type_flags)[None, :, :, :] # (1, 15, 15, F)\nall_zeros = keras.ops.zeros_like(all_pos_dists) # (Batch, 15, 15, 1)\npairwise_type_flags_tensor = all_zeros + pairwise_type_flags_tensor # (Batch, 15, 15, F)\n\npairwise_features = keras.ops.concatenate([\n    all_pos_dists, \n    vel_dots_1,\n    vel_dots_2,\n    closing_speed,\n    rot_pos_diff_dot_1,\n    rot_pos_diff_dot_2,\n    up_pos_diff_dot_1,\n    up_pos_diff_dot_2,\n    all_up_z_1,\n    all_up_z_2,\n    both_alive,\n    turn_towards_dot_1,\n    turn_towards_dot_2,\n    pairwise_type_flags_tensor,\n], axis=-1) # (Batch, 15, 15, X)\n\ndef pairwise_bias_projection(hidden_width=12):\n    bias = mlp_block([hidden_width, NUM_HEADS], final_layer_init_zero=True)(pairwise_features) # (Batch, 15, 15, H)\n    \n    return keras.ops.transpose(bias, (0, 3, 1, 2)) # (Batch, H, 15, 15)\n\nfor i in range(NUM_LAYERS):\n    pairwise_bias = pairwise_bias_projection()\n    drop_path_rate = i * MAX_DROP_PATH / (NUM_LAYERS - 1)\n    x = pre_norm_transformer_block(x, pairwise_bias, NUM_HEADS, KEY_DIM, FF_DIM, DROPOUT, drop_path_rate)\n\nx_norm = keras.layers.LayerNormalization()(x)\n\n# Create (Batch, 1) indices of zeros for the learnable query\nquery_idx = keras.ops.zeros_like(ball_raw[:, :1], 'int32')\nsummary_query = keras.layers.Embedding(1, D_MODEL)(query_idx)\n\n# Cross Attention\npool_attn = manual_multihead_attention(\n    summary_query, x_norm, x_norm, NUM_HEADS, KEY_DIM, 0.0\n)\ngame_state = keras.layers.LayerNormalization()(pool_attn)\ngame_state = keras.layers.Flatten()(game_state)\n\nneck = geglu_block(game_state, NECK_DIM, NECK_DROPOUT, reshape_back=False)\nneck = geglu_block(neck, NECK_DIM // 2, dropout=0, reshape_back=False)\n\ntrain_target_means = np.mean(train_targets, axis=tuple(range(train_targets.ndim - 1)))[:TOTAL_TARGETS]\nlogit_bias_init = keras.initializers.Constant(np.log(train_target_means))\nlogits = keras.layers.Dense(TOTAL_TARGETS, kernel_initializer='zeros', bias_initializer=logit_bias_init)(neck)\n\nball_token_raw = x[:, 0, :] \nball_token_norm = keras.layers.LayerNormalization()(ball_token_raw)\nball_physics_inputs = keras.ops.concatenate([ball_token_norm, game_state], axis=1)\nball_neck = geglu_block(ball_physics_inputs, PHYSICS_NECK_DIM, dropout=0, reshape_back=False)\nball_preds = keras.layers.Dense(NUM_BALL_PHYSICS_TARGETS, kernel_initializer='zeros')(ball_neck)\n\nplayer_raw_tokens = x[:, 1:7, :] # (Batch, 6, D)\nplayer_tokens_norm = keras.layers.LayerNormalization()(player_raw_tokens)\nbatch_zero_anchor = keras.ops.zeros_like(player_tokens_norm)\ngame_state_repeated = batch_zero_anchor + game_state[:, None, :]\nplayer_combined = keras.ops.concatenate([player_tokens_norm, game_state_repeated], axis=2)\n\n# 5. Shared MLP for all players\nplayer_neck = geglu_block(player_combined, PHYSICS_NECK_DIM, dropout=0, reshape_back=False)\nplayer_preds = keras.layers.Dense(NUM_PLAYER_PHYSICS_TARGETS, kernel_initializer='zeros')(player_neck) \nplayer_preds_flat = keras.layers.Flatten()(player_preds) # (Batch, 42)\n\nfinal_output = keras.ops.concatenate([logits, ball_preds, player_preds_flat], axis=1)\n\nmodel = keras.Model(inputs=inputs, outputs=final_output)\n\nif use_preloaded_model:\n    print(f'loading preloaded model from {preloaded_model_path}')\n    model = keras.models.load_model(preloaded_model_path)\n\n# We'll set learning_rate in the train loop below\n# For some reason Adam is better than AdamW throughout the entire learning rate schedule\n# optimizer = keras.optimizers.AdamW(learning_rate=0.0, weight_decay=1e-2, clipnorm=None)\noptimizer = keras.optimizers.Adam(learning_rate=0.0, global_clipnorm=2.0)\n\nmodel.compile(optimizer=optimizer)\n\nprint(f'Model with {len(model.layers)} layers and {model.count_params()} params')\nwith open('summary.txt', 'w') as f:\n    model.summary(print_fn=lambda l: f.write(l + '\\n'))","metadata":{"_cell_guid":"50301a11-99c8-48da-b9bc-621bc8de6f78","_uuid":"b2de7a76-ae15-4917-95b6-3b29660ea281","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:46:35.883547Z","iopub.status.busy":"2026-02-09T07:46:35.883298Z","iopub.status.idle":"2026-02-09T07:46:40.657963Z","shell.execute_reply":"2026-02-09T07:46:40.657189Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":4.785856,"end_time":"2026-02-09T07:46:40.659426","exception":false,"start_time":"2026-02-09T07:46:35.873570","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = model.optimizer\n\nx_shape = train_features[0][0][0].shape\nprint(f'train_features shape: {x_shape}')\nmodel.build(x_shape)\noptimizer.build(model.trainable_variables)\n\ny_shape = train_targets[0][0][0].shape\nprint(f'train_targets shape: {y_shape}')","metadata":{"_cell_guid":"9507bbd4-eb76-4cef-a482-c5ef0db905b6","_uuid":"14ddfc96-6b0d-4faf-963d-91a3f4ed9cfd","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:46:40.678074Z","iopub.status.busy":"2026-02-09T07:46:40.677821Z","iopub.status.idle":"2026-02-09T07:46:41.494358Z","shell.execute_reply":"2026-02-09T07:46:41.493486Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.827676,"end_time":"2026-02-09T07:46:41.495939","exception":false,"start_time":"2026-02-09T07:46:40.668263","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IDX_GOALS_END = TOTAL_TARGETS\nIDX_BALL_END = IDX_GOALS_END + NUM_BALL_PHYSICS_TARGETS\nIDX_PLAYERS_END = IDX_BALL_END + 6 * NUM_PLAYER_PHYSICS_TARGETS\nPHYSICS_WEIGHT = 1.0 # scale for all physics preds relative to goal preds\nPLAYER_WEIGHT = 0.4 # relative weight of each player vector RMSE vs ball's\nDEMO_WEIGHT = 0.1 # relative weight of each player's demo BCE loss to their RMSE loss\n\ndef jax_loss(target, outputs):\n    # goal loss cross-entropy\n    log_prob = jax.nn.log_softmax(outputs[:, :IDX_GOALS_END], axis=-1)\n    goal_loss = -jnp.mean(jnp.sum(target[:, :IDX_GOALS_END] * log_prob, axis=-1), axis=0)\n\n    # ball loss RMSE\n    ball_target = target[:, IDX_GOALS_END:IDX_BALL_END]\n    ball_outputs = outputs[:, IDX_GOALS_END:IDX_BALL_END]\n    ball_valid = ~jnp.isnan(ball_target)\n    \n    ball_sq_errors = jnp.square(ball_outputs - jnp.where(ball_valid, ball_target, 0.0))\n    ball_masked_sq_errors = jnp.where(ball_valid, ball_sq_errors, 0.0)\n    ball_loss = jnp.sqrt(jnp.sum(ball_masked_sq_errors) / jnp.sum(ball_valid) + 1e-4)\n\n    # --- 3. Player Loss (Generalized for N sets) ---\n    # Slice from ball end to the very end of the tensor (or specific index if data follows)\n    p_targets_flat = target[:, IDX_BALL_END:]\n    p_preds_flat = outputs[:, IDX_BALL_END:]\n    \n    # Reshape: (Batch, Num_Sets, Players, Features)\n    # The '-1' allows JAX to automatically infer the number of sets based on input size\n    p_targets = p_targets_flat.reshape(target.shape[0], -1, 6, 7)\n    p_preds = p_preds_flat.reshape(target.shape[0], 6, -1, 7)\n    p_preds = jnp.swapaxes(p_preds, 1, 2)\n\n    # Split features: first 6 are movement, last 1 is demo logit\n    # Slicing the last dimension works regardless of how many 'Sets' exist in dim 1\n    p_mov_target = p_targets[..., :6]\n    p_mov_pred = p_preds[..., :6]\n    p_demo_target = p_targets[..., 6:7]\n    p_demo_logit = p_preds[..., 6:7]\n\n    # A. Player Movement MSE\n    # Operations are element-wise, so they handle the extra 'Sets' dimension automatically\n    mov_valid = ~jnp.isnan(p_mov_target)\n    mov_sq_err = jnp.square(p_mov_pred - jnp.where(mov_valid, p_mov_target, 0.0))\n    mov_sq_err_masked = jnp.where(mov_valid, mov_sq_err, 0.0)\n    # Global RMSE over all sets, players, and coordinates\n    mov_loss = jnp.sqrt(jnp.sum(mov_sq_err_masked) / (jnp.sum(mov_valid) + 1e-4) + 1e-4)\n\n    # B. Player Demo Classification (BCEWithLogits)\n    demo_valid = ~jnp.isnan(p_demo_target)\n    safe_demo_target = jnp.where(demo_valid, p_demo_target, 0.0)\n    \n    bce = (jnp.maximum(p_demo_logit, 0.0) \n           - p_demo_logit * safe_demo_target \n           + jnp.log(1.0 + jnp.exp(-jnp.abs(p_demo_logit))))\n           \n    # Mean BCE over all sets and players\n    demo_loss = jnp.sum(jnp.where(demo_valid, bce, 0.0)) / (jnp.sum(demo_valid) + 1e-4)\n\n    physics_loss = NUM_BALL_PHYSICS_TARGETS / 6 * ball_loss + PLAYER_WEIGHT * 6 * NUM_PLAYER_PHYSICS_TARGETS / 7 * (mov_loss + DEMO_WEIGHT * demo_loss)\n\n    return goal_loss + PHYSICS_WEIGHT * physics_loss","metadata":{"_cell_guid":"d7ca8569-c36e-4efe-a7c1-85c98ef0004d","_uuid":"656a4031-3eff-42dc-87c0-f957796e1f32","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:46:41.515015Z","iopub.status.busy":"2026-02-09T07:46:41.514744Z","iopub.status.idle":"2026-02-09T07:46:41.524063Z","shell.execute_reply":"2026-02-09T07:46:41.523437Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.020426,"end_time":"2026-02-09T07:46:41.525410","exception":false,"start_time":"2026-02-09T07:46:41.504984","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def jax_put(variables):\n    return jax.device_put_replicated([v.value for v in variables], devices)\n\ntrainable_vars = jax_put(model.trainable_variables)\nnon_trainable_vars = jax_put(model.non_trainable_variables)\noptimizer_vars = jax_put(optimizer.variables)","metadata":{"_cell_guid":"af6b693d-3e8c-408c-af78-756088346915","_uuid":"104e0738-0bc4-45d8-bb70-9e10046bd0a9","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:46:41.544107Z","iopub.status.busy":"2026-02-09T07:46:41.543513Z","iopub.status.idle":"2026-02-09T07:46:42.924150Z","shell.execute_reply":"2026-02-09T07:46:42.923483Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":1.391891,"end_time":"2026-02-09T07:46:42.925826","exception":false,"start_time":"2026-02-09T07:46:41.533935","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def update_model(trainable_vars, non_trainable_vars):\n    for var, value in zip(model.trainable_variables, trainable_vars):\n        var.assign(value[0])\n    for var, value in zip(model.non_trainable_variables, non_trainable_vars):\n        var.assign(value[0])\n\ndef update_optimizer(optimizer_vars):\n    for var, value in zip(optimizer.variables, optimizer_vars):\n        var.assign(value[0])","metadata":{"_cell_guid":"4476cf54-566b-42af-a22e-34ea1d2be7e2","_uuid":"442b7a0c-fbfd-404a-bf62-1d0c58f11fd4","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:46:42.944180Z","iopub.status.busy":"2026-02-09T07:46:42.943946Z","iopub.status.idle":"2026-02-09T07:46:42.948437Z","shell.execute_reply":"2026-02-09T07:46:42.947901Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.015011,"end_time":"2026-02-09T07:46:42.949754","exception":false,"start_time":"2026-02-09T07:46:42.934743","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_step(trainable_vars, non_trainable_vars, optimizer_vars, x_stack, y_stack):\n    \"\"\"\n    x_stack shape: (ACCUM_STEPS, Device_Batch, ...)\n    y_stack shape: (ACCUM_STEPS, Device_Batch, ...)\n    \"\"\"\n    \n    # 1. Define the function for a single micro-batch\n    def micro_batch_step(carry, batch):\n        # Unpack carry: non_trainable_vars updates (e.g. RNG state) must flow through the loop\n        current_non_trainable_vars = carry\n        x, y = batch\n        \n        def compute_loss(trainable, non_trainable, x_in, y_in):\n            y_pred, new_non_trainable = model.stateless_call(\n                trainable, non_trainable, x_in, training=True\n            )\n            loss = jax_loss(y_in, y_pred)\n            return loss, new_non_trainable\n\n        # Compute gradients for this micro-batch\n        grad_fn = jax.value_and_grad(compute_loss, has_aux=True)\n        (loss, new_non_trainable_vars), grads = grad_fn(\n            trainable_vars, current_non_trainable_vars, x, y\n        )\n        new_non_trainable_vars = jax.tree.map(\n            lambda x, y: x.astype(y.dtype), \n            new_non_trainable_vars, \n            current_non_trainable_vars\n        )\n        \n        # Return new state and (outputs to stack)\n        return new_non_trainable_vars, (grads, loss)\n\n    # 2. Run the loop over ACCUM_STEPS using jax.lax.scan\n    # This executes purely on device (no python overhead)\n    final_non_trainable_vars, (grads_stack, loss_stack) = jax.lax.scan(\n        micro_batch_step, \n        non_trainable_vars, \n        (x_stack, y_stack)\n    )\n    \n    # 3. Aggregate results\n    # Average the gradients across the accumulation steps\n    grads = jax.tree.map(lambda g: jnp.mean(g, axis=0), grads_stack)\n    # Average the loss\n    loss = jnp.mean(loss_stack)\n\n    # 4. Sync across devices (Data Parallelism)\n    grads = jax.lax.pmean(grads, axis_name='batch')\n    loss = jax.lax.pmean(loss, axis_name='batch')\n    final_non_trainable_vars = jax.lax.pmean(final_non_trainable_vars, axis_name='batch')\n\n    # 5. Apply Optimizer Update (Once per effective batch)\n    new_trainable_vars, new_optimizer_vars = optimizer.stateless_apply(\n        optimizer_vars, grads, trainable_vars\n    )\n\n    return new_trainable_vars, final_non_trainable_vars, new_optimizer_vars, loss\n\n# Re-create the pmap\ntrain_step_pmap = jax.pmap(train_step, axis_name='batch')\n\ndef model_train(trainable_vars, non_trainable_vars, optimizer_vars, features, targets):\n    total_loss = jnp.zeros(())\n    for batch_x, batch_y in zip(features, targets):\n        trainable_vars, non_trainable_vars, optimizer_vars, loss = train_step_pmap(\n            trainable_vars, non_trainable_vars, optimizer_vars, batch_x, batch_y\n        )\n        # Loss is already averaged, just need to get it from one device\n        total_loss += loss[0]\n    jax.block_until_ready((trainable_vars, non_trainable_vars, optimizer_vars, total_loss))\n    loss = jax.device_get(total_loss) / len(features)\n\n    return trainable_vars, non_trainable_vars, optimizer_vars, loss","metadata":{"_cell_guid":"af16dda5-634f-49df-bb6c-57c054635c0a","_uuid":"30216c4f-11f3-4916-8774-5c553f988bd7","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:46:42.968045Z","iopub.status.busy":"2026-02-09T07:46:42.967450Z","iopub.status.idle":"2026-02-09T07:46:42.977639Z","shell.execute_reply":"2026-02-09T07:46:42.976915Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.020715,"end_time":"2026-02-09T07:46:42.978937","exception":false,"start_time":"2026-02-09T07:46:42.958222","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_softmax3_step(trainable_vars, non_trainable_vars, x):\n    outputs, _ = model.stateless_call(trainable_vars, non_trainable_vars, x, training=False)\n    preds = jnn.softmax(outputs[:, :TOTAL_TARGETS], axis=-1)\n\n    # Sum up the buckets for each team, then concat back together into (N, 3)\n    return jnp.concatenate([\n            jnp.sum(preds[:, :TARGETS_PER_TEAM], axis=-1, keepdims=True),\n            jnp.sum(preds[:, TARGETS_PER_TEAM:2 * TARGETS_PER_TEAM], axis=-1, keepdims=True),\n            preds[:, -1:],\n        ], axis=-1)\n\npredict_softmax3_step_pmap = jax.pmap(predict_softmax3_step, axis_name='batch')\n\n# Gets model predictions + applies softmax to convert from logits to probabilities + returns only 10sec channel\ndef model_predict_softmax3(trainable_vars, non_trainable_vars, features, batch_size=DEVICE_MICRO_BATCH_SIZE):\n    num_samples = len(features)\n    \n    # Calculate padding needed to make samples divisible by batch_size, will truncate later\n    remainder = num_samples % batch_size\n    if remainder != 0:\n        padding = np.zeros((batch_size - remainder, *features.shape[1:]), dtype=features.dtype)\n        features = np.concatenate([features, padding], axis=0)\n    \n    num_batches = len(features) // batch_size\n    device_batch_size = batch_size // NUM_DEVICES\n\n    all_results = []\n    for i in range(num_batches):\n        batch_x = features[i * batch_size:(i+1) * batch_size]\n        batch_x = batch_x.reshape(NUM_DEVICES, device_batch_size, *batch_x.shape[1:])\n        \n        batch_y_pred = predict_softmax3_step_pmap(trainable_vars, non_trainable_vars, batch_x)\n        \n        batch_preds_np = np.array(jax.device_get(batch_y_pred), dtype=np.float32).reshape(batch_size, -1)\n        all_results.append(batch_preds_np)\n        \n    predictions = np.concatenate(all_results, axis=0)[:num_samples]\n    \n    return predictions\n\ndef model_validate(trainable_vars, non_trainable_vars, batch_size=DEVICE_MICRO_BATCH_SIZE):\n    all_predictions = []\n    for features, flip_y in zip(val_features_list, val_flip_y_list):\n        raw_predictions = model_predict_softmax3(\n            trainable_vars, non_trainable_vars, features, batch_size=batch_size)\n        predictions = pl.LazyFrame({\n            'team0_scoring_within_10sec': raw_predictions[:, 0],\n            'team1_scoring_within_10sec': raw_predictions[:, 1],\n            'neither_scoring_within_10sec': raw_predictions[:, 2],\n            'flip_y': flip_y,\n        })\n        predictions = swap(\n            predictions,\n            pl.col('flip_y'),\n            ['team0_scoring_within_10sec', 'team1_scoring_within_10sec'],\n            ['team1_scoring_within_10sec', 'team0_scoring_within_10sec']\n        )\n        all_predictions.append(predictions.drop('flip_y').collect())\n\n    final_predictions = (sum(all_predictions) / len(all_predictions)).to_numpy()\n    return -np.mean(np.sum(val_targets * np.log(np.clip(final_predictions, 1e-10, None)), axis=-1))","metadata":{"_cell_guid":"edb7e7c2-a3f4-469d-a205-dff6a9ffa9dc","_uuid":"c665c744-ce65-4ff0-9c6c-4fa885c588e9","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:46:42.996799Z","iopub.status.busy":"2026-02-09T07:46:42.996289Z","iopub.status.idle":"2026-02-09T07:46:43.006488Z","shell.execute_reply":"2026-02-09T07:46:43.005763Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.020699,"end_time":"2026-02-09T07:46:43.007893","exception":false,"start_time":"2026-02-09T07:46:42.987194","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n# preload\ntrainable_vars, non_trainable_vars, optimizer_vars, loss = train_step_pmap(\n    trainable_vars, non_trainable_vars, optimizer_vars, train_features[0], train_targets[0]\n)\nprint(loss[0])","metadata":{"_cell_guid":"2c15701a-821e-4493-bc70-cb4206131d4c","_uuid":"c03d01e8-eb11-48f9-abc7-54dbcdf02c69","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:46:43.026269Z","iopub.status.busy":"2026-02-09T07:46:43.025799Z","iopub.status.idle":"2026-02-09T07:47:45.926457Z","shell.execute_reply":"2026-02-09T07:47:45.925651Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":62.919605,"end_time":"2026-02-09T07:47:45.935950","exception":false,"start_time":"2026-02-09T07:46:43.016345","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"models_path = pathlib.Path('models')\nmodels_path.mkdir(parents=True, exist_ok=True)\ndef save_model(epoch):\n    update_model(trainable_vars, non_trainable_vars)\n    model_path = models_path / f'model_{epoch:03d}.keras'\n    model.save(str(model_path))","metadata":{"_cell_guid":"4e37baac-ba99-4d35-b40c-72c9ad4bd2b9","_uuid":"a203689e-6b78-47ff-af7d-006b11a8cf75","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:47:45.953739Z","iopub.status.busy":"2026-02-09T07:47:45.953349Z","iopub.status.idle":"2026-02-09T07:47:45.957708Z","shell.execute_reply":"2026-02-09T07:47:45.957076Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.01486,"end_time":"2026-02-09T07:47:45.959108","exception":false,"start_time":"2026-02-09T07:47:45.944248","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def constant_lr(n_epochs, target_lr=None):\n    return n_epochs, lambda epoch, start_lr: target_lr or start_lr\ndef linear_lr(n_epochs, target_lr):\n    return n_epochs, lambda epoch, start_lr: start_lr + (target_lr - start_lr) * epoch / n_epochs\ndef cosine_lr(n_epochs, target_lr):\n    return n_epochs, lambda epoch, start_lr: target_lr + (start_lr - target_lr) * 0.5 * (1 + np.cos(np.pi * epoch / n_epochs))\n\nschedule = [\n    constant_lr(0, 1e-5),\n    linear_lr(1, 1.5e-3),\n    cosine_lr(17, 1e-4),\n    linear_lr(4, 1e-6),\n]\n\ndef learning_rate(epoch):\n    current_lr = 0\n    for n_epochs, getter in schedule:\n        if epoch < n_epochs:\n            return getter(epoch, current_lr)\n        epoch -= n_epochs\n        current_lr = getter(n_epochs, current_lr)\n    return current_lr","metadata":{"execution":{"iopub.execute_input":"2026-02-09T07:47:45.977016Z","iopub.status.busy":"2026-02-09T07:47:45.976791Z","iopub.status.idle":"2026-02-09T07:47:45.982694Z","shell.execute_reply":"2026-02-09T07:47:45.982097Z"},"papermill":{"duration":0.016417,"end_time":"2026-02-09T07:47:45.984006","exception":false,"start_time":"2026-02-09T07:47:45.967589","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"time_limit_hours = .2 if INTERACTIVE_SESSION else 11\nmax_epochs = 2 if INTERACTIVE_SESSION else sum([s[0] for s in schedule])\nlosses = collections.defaultdict(list)\nt_0 = time.time()\nepoch = 19 if use_preloaded_model else 0\nstride_i = 0\nlr_index = [v.path for v in optimizer.variables].index(optimizer.learning_rate.path)\nprint('starting train loop')\nwith concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor:\n    new_train_future = executor.submit(get_augmented_train_xy, stride_i + 1)\n    t_start = time.time()\n    total_train_loss = 0\n    while True:\n        lr = learning_rate(epoch + stride_i / TRAIN_STRIDE)\n        optimizer.learning_rate = lr\n        lr_tensor = jax.device_put_replicated(jnp.array(lr, dtype=\"float32\"), devices)\n        optimizer_vars[lr_index] = lr_tensor\n\n        trainable_vars, non_trainable_vars, optimizer_vars, train_loss = model_train(\n            trainable_vars, non_trainable_vars, optimizer_vars, train_features, train_targets\n        )\n        t_train = time.time()\n        total_train_loss += train_loss\n\n        stride_i += 1\n        train_features, train_targets = new_train_future.result()\n        new_train_future = executor.submit(get_augmented_train_xy, stride_i + 1)\n\n        if stride_i < TRAIN_STRIDE:\n            continue\n        stride_i = 0\n        avg_train_loss = total_train_loss / TRAIN_STRIDE\n        losses['loss'].append(avg_train_loss)\n        total_train_loss = 0\n\n        t_val_start = time.time()\n        val_loss_10 = model_validate(trainable_vars, non_trainable_vars)\n        t_val = time.time()\n        losses['val_loss'].append(val_loss_10)\n\n        elapsed_hours = (t_val - t_0) / 3600\n        print(f'[{datetime.now().strftime(\"%H:%M\")}] Epoch {epoch} / {max_epochs}: train loss {avg_train_loss:.5f} sec {t_train - t_start:.2f} '\n              f'val loss {val_loss_10:.5f} sec {t_val - t_val_start:.2f} | hours {elapsed_hours:.2f} lr {learning_rate(epoch):.6f} -> {lr:.6f}')\n\n        update_optimizer(optimizer_vars)\n        save_model(epoch)\n        \n        epoch += 1\n        if time_limit_hours and elapsed_hours >= time_limit_hours:\n            print('Over the time limit, stopping...')\n            break\n        if max_epochs and epoch >= max_epochs:\n            print('Over the epoch limit, stopping...')\n            break\n\n        t_start = time.time()","metadata":{"_cell_guid":"49c8a20e-6ac0-4ddb-9068-d9c81f06fb52","_uuid":"25e9e130-a1ce-48d4-a14b-a24e626e5a63","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T07:47:46.001882Z","iopub.status.busy":"2026-02-09T07:47:46.001654Z","iopub.status.idle":"2026-02-09T10:20:20.294887Z","shell.execute_reply":"2026-02-09T10:20:20.293242Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":9154.304395,"end_time":"2026-02-09T10:20:20.296753","exception":false,"start_time":"2026-02-09T07:47:45.992358","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"raw_predictions = keras.ops.softmax(logits, axis=-1)\npred_team0 = keras.ops.sum(raw_predictions[:, :TARGETS_PER_TEAM], axis=-1, keepdims=True)\npred_team1 = keras.ops.sum(raw_predictions[:, TARGETS_PER_TEAM:2 * TARGETS_PER_TEAM], axis=-1, keepdims=True)\npred_neither = raw_predictions[:, 2 * TARGETS_PER_TEAM:]\npredictions = keras.ops.concatenate([pred_team0, pred_team1, pred_neither], axis=-1)\n\ninference_model = keras.Model(inputs=inputs, outputs=predictions)\n\n# Need to call once to force initialization before we can export\ndummy_data = np.zeros((1, *inference_model.input_shape[1:]), dtype=np.float32)\n_ = inference_model(dummy_data)\n\ninference_model.save('model.keras')","metadata":{"execution":{"iopub.execute_input":"2026-02-09T10:20:20.319910Z","iopub.status.busy":"2026-02-09T10:20:20.319643Z","iopub.status.idle":"2026-02-09T10:20:33.409878Z","shell.execute_reply":"2026-02-09T10:20:33.409216Z"},"papermill":{"duration":13.105374,"end_time":"2026-02-09T10:20:33.411567","exception":false,"start_time":"2026-02-09T10:20:20.306193","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test = pl.scan_parquet(dataset_path / 'test.parquet')\ntest_columns = test.collect_schema().names()\ntest, test_columns = trim_columns(test, test_columns)","metadata":{"_cell_guid":"9026b950-f3d6-4a09-a723-dc4acadadae2","_uuid":"e3c11602-c99e-4c47-ac0f-f0bc7b2dcbd2","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T10:20:33.430719Z","iopub.status.busy":"2026-02-09T10:20:33.430430Z","iopub.status.idle":"2026-02-09T10:20:33.498392Z","shell.execute_reply":"2026-02-09T10:20:33.497821Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.079071,"end_time":"2026-02-09T10:20:33.499966","exception":false,"start_time":"2026-02-09T10:20:33.420895","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_submission(model, num_samples=1):\n    ids = None\n    all_predictions = []\n    for i in range(num_samples):\n        test_df = augment_data(test, seed=i, shuffle=False, append_flip_y=True).collect()\n        ids = test_df['id']\n        features = test_df.drop('id', 'flip_y').to_numpy()\n        t_start = time.time()\n        raw_predictions = model_predict_softmax3(\n            trainable_vars, non_trainable_vars, features, batch_size=DEVICE_MICRO_BATCH_SIZE)\n        t_delta = time.time() - t_start\n        print(f'Predicted {len(test_df)} rows in {t_delta:.2f} sec')\n        predictions = pl.LazyFrame({\n            'team0_scoring_within_10sec': raw_predictions[:, 0],\n            'team1_scoring_within_10sec': raw_predictions[:, 1],\n            'neither_scoring_within_10sec': raw_predictions[:, 2],\n            'flip_y': test_df['flip_y'],\n        })\n        \n        predictions = swap(\n            predictions,\n            pl.col('flip_y'),\n            ['team0_scoring_within_10sec', 'team1_scoring_within_10sec'],\n            ['team1_scoring_within_10sec', 'team0_scoring_within_10sec']\n        )\n        all_predictions.append(predictions.drop('flip_y').collect())\n\n    return (sum(all_predictions) / len(all_predictions)).insert_column(0, ids)","metadata":{"_cell_guid":"a750aa13-2597-4080-b5d8-63e1cb729b53","_uuid":"e191d21b-e376-4914-af1e-72e3d56b6388","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T10:20:33.518932Z","iopub.status.busy":"2026-02-09T10:20:33.518682Z","iopub.status.idle":"2026-02-09T10:20:33.525174Z","shell.execute_reply":"2026-02-09T10:20:33.524605Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.017427,"end_time":"2026-02-09T10:20:33.526513","exception":false,"start_time":"2026-02-09T10:20:33.509086","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\nsubmission = get_submission(model, num_samples=2 if INTERACTIVE_SESSION else 4)","metadata":{"_cell_guid":"0dcc7edd-031b-4ae7-a871-d8bd1d6e169f","_uuid":"39e74c0a-e52f-40a8-96d0-491049796938","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T10:20:33.545459Z","iopub.status.busy":"2026-02-09T10:20:33.545225Z","iopub.status.idle":"2026-02-09T10:21:39.260530Z","shell.execute_reply":"2026-02-09T10:21:39.259662Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":65.735986,"end_time":"2026-02-09T10:21:39.271478","exception":false,"start_time":"2026-02-09T10:20:33.535492","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SUBMISSION_PATH = 'submission.parquet'\nsubmission.write_parquet(SUBMISSION_PATH)\nprint(f'wrote {SUBMISSION_PATH}')","metadata":{"_cell_guid":"4b662806-14fd-4da9-bf10-015174b6c67f","_uuid":"9fcd9f29-04f5-4012-8d31-1751e9db6d3c","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T10:21:39.291111Z","iopub.status.busy":"2026-02-09T10:21:39.290482Z","iopub.status.idle":"2026-02-09T10:21:39.506310Z","shell.execute_reply":"2026-02-09T10:21:39.505452Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.227411,"end_time":"2026-02-09T10:21:39.507827","exception":false,"start_time":"2026-02-09T10:21:39.280416","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SUBMISSION_PATH = 'submission.csv'\nsubmission.write_csv(SUBMISSION_PATH)\nprint(f'wrote {SUBMISSION_PATH}')","metadata":{"_cell_guid":"2aa425e0-3d6a-48a1-9a16-6b47cd925f6e","_uuid":"7822dd14-b077-4e05-91d7-8662f2bc8fa5","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T10:21:39.528083Z","iopub.status.busy":"2026-02-09T10:21:39.527577Z","iopub.status.idle":"2026-02-09T10:21:39.722400Z","shell.execute_reply":"2026-02-09T10:21:39.721403Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.206421,"end_time":"2026-02-09T10:21:39.724058","exception":false,"start_time":"2026-02-09T10:21:39.517637","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_losses(*loss_keys):\n    data = {key: losses[key] for key in loss_keys}\n    data['epoch'] = range(len(data[loss_keys[0]]))\n    return pl.DataFrame(data).hvplot.line(x='epoch', y=loss_keys)","metadata":{"_cell_guid":"c63802df-9c5e-4871-af9c-5bef0b1e567f","_uuid":"cafe5085-bd0c-449b-906e-afeb3eebcc3a","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T10:21:39.743841Z","iopub.status.busy":"2026-02-09T10:21:39.743542Z","iopub.status.idle":"2026-02-09T10:21:39.748030Z","shell.execute_reply":"2026-02-09T10:21:39.747304Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.015777,"end_time":"2026-02-09T10:21:39.749448","exception":false,"start_time":"2026-02-09T10:21:39.733671","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_losses('loss')","metadata":{"_cell_guid":"2812f0f4-1693-4fd4-9e81-bffe7a46cc84","_uuid":"ec4bff71-a345-4a88-bcb5-08464b15db3e","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T10:21:39.770206Z","iopub.status.busy":"2026-02-09T10:21:39.769286Z","iopub.status.idle":"2026-02-09T10:21:40.173557Z","shell.execute_reply":"2026-02-09T10:21:40.172755Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.416266,"end_time":"2026-02-09T10:21:40.175225","exception":false,"start_time":"2026-02-09T10:21:39.758959","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_losses('val_loss')","metadata":{"_cell_guid":"f1698fb3-1e2f-4732-bd09-1a46dfa224a9","_uuid":"8357c019-0936-4dd1-9cad-1c3632e79f54","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T10:21:40.199020Z","iopub.status.busy":"2026-02-09T10:21:40.198457Z","iopub.status.idle":"2026-02-09T10:21:40.275558Z","shell.execute_reply":"2026-02-09T10:21:40.274930Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.089308,"end_time":"2026-02-09T10:21:40.277004","exception":false,"start_time":"2026-02-09T10:21:40.187696","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission.drop('id').hvplot.hist(bins=np.linspace(0, 1, 100))","metadata":{"_cell_guid":"3b16b6e4-045c-4c4c-b03b-df5150e90279","_uuid":"5e7ad357-f6b0-4f46-a180-0f41ce8c088a","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T10:21:40.298008Z","iopub.status.busy":"2026-02-09T10:21:40.297347Z","iopub.status.idle":"2026-02-09T10:21:45.328909Z","shell.execute_reply":"2026-02-09T10:21:45.328193Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":5.043449,"end_time":"2026-02-09T10:21:45.330510","exception":false,"start_time":"2026-02-09T10:21:40.287061","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Exporting to ONNX requires tf2onnx and running in a new process so Keras can use the Tensorflow backend\nconvert_onnx_script_path = str(temp_dir_path / 'convert_onnx.py')","metadata":{"_cell_guid":"c413b0c2-3ef9-46d7-bfdd-d71bb1c1c871","_uuid":"2a741ae8-1a36-46f2-a65d-0d3a290e026d","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T10:21:45.353262Z","iopub.status.busy":"2026-02-09T10:21:45.352515Z","iopub.status.idle":"2026-02-09T10:21:45.356261Z","shell.execute_reply":"2026-02-09T10:21:45.355686Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.016033,"end_time":"2026-02-09T10:21:45.357552","exception":false,"start_time":"2026-02-09T10:21:45.341519","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile {convert_onnx_script_path}\nimport os\nos.environ['KERAS_BACKEND'] = 'tensorflow'\nimport keras\nimport sys\nimport tf2onnx\nimport tensorflow as tf\n\n@keras.saving.register_keras_serializable()\ndef gelu_approx(x):\n    return keras.activations.gelu(x, approximate=True)\n\nmodel = keras.models.load_model(sys.argv[1])\nspec = (tf.TensorSpec(model.input_shape, tf.float32, name='inputs'),)\noutput_path = sys.argv[2]\n\nmodel_proto, _ = tf2onnx.convert.from_keras(model, input_signature=spec, output_path=output_path)\nprint(f'Successfully exported to {output_path}')","metadata":{"_cell_guid":"4718e653-7ade-4473-b802-2338ff0bd9d2","_uuid":"da5daa97-a88c-4398-9b18-a83d78835e23","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T10:21:45.379300Z","iopub.status.busy":"2026-02-09T10:21:45.378903Z","iopub.status.idle":"2026-02-09T10:21:45.384486Z","shell.execute_reply":"2026-02-09T10:21:45.383681Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.017926,"end_time":"2026-02-09T10:21:45.385840","exception":false,"start_time":"2026-02-09T10:21:45.367914","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python {convert_onnx_script_path} model.keras model.onnx","metadata":{"_cell_guid":"5e2517fb-d219-41df-b923-5a5a943045af","_uuid":"350ae111-0b48-46d9-92a4-cfccb5142414","collapsed":false,"execution":{"iopub.execute_input":"2026-02-09T10:21:45.407546Z","iopub.status.busy":"2026-02-09T10:21:45.407098Z","iopub.status.idle":"2026-02-09T10:22:31.603493Z","shell.execute_reply":"2026-02-09T10:22:31.602566Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":46.209343,"end_time":"2026-02-09T10:22:31.605522","exception":false,"start_time":"2026-02-09T10:21:45.396179","status":"completed"},"tags":[]},"outputs":[],"execution_count":null}]}