{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'\nimport tensorflow as tf\nfrom tensorflow import keras\nimport tensorflow as tf\nimport tensorflow_addons as tfa","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2022-11-03T06:32:23.723043Z","iopub.status.busy":"2022-11-03T06:32:23.722427Z","iopub.status.idle":"2022-11-03T06:32:29.167260Z","shell.execute_reply":"2022-11-03T06:32:29.166291Z"},"papermill":{"duration":5.459752,"end_time":"2022-11-03T06:32:29.169817","exception":false,"start_time":"2022-11-03T06:32:23.710065","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\n\nseed = hash('kaggle') % 2**32\nrandom.seed(seed)\nnp.random.seed(seed)\ntf.random.set_seed(seed)","metadata":{"execution":{"iopub.execute_input":"2022-11-03T06:32:29.184358Z","iopub.status.busy":"2022-11-03T06:32:29.182516Z","iopub.status.idle":"2022-11-03T06:32:29.188305Z","shell.execute_reply":"2022-11-03T06:32:29.187371Z"},"papermill":{"duration":0.014832,"end_time":"2022-11-03T06:32:29.190373","exception":false,"start_time":"2022-11-03T06:32:29.175541","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def calculate_modifiers(df, key, mods):\n    d = {}\n    for (name, size) in mods:\n        d[name] = df[key].mod(size)\n        df[key] //= size\n    return d\n\n# Sets values for each column in col_list0 to the values from each corresponding column in col_list1\n# Does all swaps at once to avoid order-of-operations mistakes\n# Only affects rows where cond is True\ndef swap(df, cond, col_list0, col_list1):\n    df.loc[cond, col_list0] = df.loc[cond, col_list1].values\n\n# Sets values for each player in p_list0 to the values from each corresponding player in p_list1\ndef swap_players(df, cond, p_list0, p_list1):\n    def get_cols(p_list):\n        return [col for p in p_list for col in df.columns if col.startswith(f'p{p}_')]\n    swap(df, cond, get_cols(p_list0), get_cols(p_list1))\n\ndef swap_boosts(df, 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    swap(df, cond, get_cols(boost_i_list0), get_cols(boost_i_list1))\n\n# order in [0, 5] determines the ordering of the 3 players starting with base_p\n# 6 options: swap(0,1), swap(0,2), swap(1,2), rot_left(), rot_right(), keep_same()\ndef apply_reorder(df, key, order, base_p):\n    if order <= 2:\n        p0 = base_p + (order // 2)\n        p1 = base_p + ((order + 3) // 2)\n        swap_players(df, key == order, [p0, p1], [p1, p0])\n    elif order <= 4:\n        p_list = list(range(base_p, base_p + 3))\n        rot_list = [p_list[-1]] + p_list[:-1] if order == 3 else p_list[1:] + [p_list[0]]\n        swap_players(df, key == order, p_list, rot_list)\n    else:\n        # Keep default order\n        pass\n\ndef reorder_players(df, key, base_p):\n    for order in range(6):\n        apply_reorder(df, key, order, base_p)\n\ndef augment_data(df, shuffle=True, append_flip_y=False, salt='nacl'):\n    df['hash'] = [hash(f'{salt}{x}') for x in range(len(df))]\n    if shuffle:\n        df.sort_values('hash', inplace=True)\n    \n    mods = calculate_modifiers(df, 'hash', [('flip_x', 2), ('flip_y', 2), ('team1_order', 6), ('team2_order', 6)])\n    \n    # flip_x\n    flip_x = mods['flip_x'] == 1\n    df.loc[flip_x, [col for col in df.columns if col.endswith('_x')]] *= -1\n    # Swap boosts 0 <-> 1, 2 <-> 3, 4 <-> 5 if present\n    if 'boost0_timer' in df:\n        swap_boosts(df, flip_x, list(range(6)), [1, 0, 3, 2, 5, 4])\n    \n    # flip_y -- Also swap teams and which team scored\n    flip_y = mods['flip_y'] == 1\n    if append_flip_y:\n        df['flip_y'] = flip_y\n    df.loc[flip_y, [col for col in df.columns if col.endswith('_y')]] *= -1\n    for target_col in [col for col in df.columns if col.startswith('team_scoring')]:\n        df.loc[flip_y, target_col] = -1 * df[target_col][flip_y]\n    # Swapping all pairs at once incurred a large ram cost\n    for i in range(3):\n        swap_players(df, flip_y, [i, i + 3], [i + 3, i])\n    # Swap boosts 0 <-> 4, 1 <-> 5 if present\n    if 'boost0_timer' in df:\n        swap_boosts(df, flip_y, [0, 1, 4, 5], [4, 5, 0, 1])\n    \n    # team1_order\n    reorder_players(df, mods['team1_order'], 0)\n    \n    # team2_order\n    reorder_players(df, mods['team2_order'], 3)\n    \n    df.drop(['hash'], axis=1, inplace=True)","metadata":{"execution":{"iopub.execute_input":"2022-11-03T06:32:29.202994Z","iopub.status.busy":"2022-11-03T06:32:29.202402Z","iopub.status.idle":"2022-11-03T06:32:29.219917Z","shell.execute_reply":"2022-11-03T06:32:29.219024Z"},"papermill":{"duration":0.026138,"end_time":"2022-11-03T06:32:29.221958","exception":false,"start_time":"2022-11-03T06:32:29.195820","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TARGETS = list(range(1, 11))\n# TARGETS = [10]\nprint(f'using targets {TARGETS}')\n\ndef get_targets(df):\n    def get_target(t):\n        return (df.team_scoring_next * (df.event_time >= -t) + 1).to_numpy(dtype='float32')[:, np.newaxis]\n    return np.concatenate([get_target(t) for t in TARGETS], axis=1)\n\ndef prep(df, salt='nacl'):\n    augment_data(df, shuffle=True, salt=salt)\n    features = df.drop(['event_time', 'team_scoring_next'], axis=1).to_numpy(dtype='float32')\n    targets = get_targets(df)\n    return features, targets\n\n# Reshape for proper orientation for multiple outputs in our keras model\n# Needs to be done after slicing a batch\ndef reshape_targets(targets):\n    return [targets[:, i] for i in range(len(TARGETS))]","metadata":{"execution":{"iopub.execute_input":"2022-11-03T06:32:29.233925Z","iopub.status.busy":"2022-11-03T06:32:29.233633Z","iopub.status.idle":"2022-11-03T06:32:29.241680Z","shell.execute_reply":"2022-11-03T06:32:29.240779Z"},"papermill":{"duration":0.017132,"end_time":"2022-11-03T06:32:29.244636","exception":false,"start_time":"2022-11-03T06:32:29.227504","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# If not using, we add validation set to training set to have maximum data\nUSE_VAL = False\n# USE_VAL = True\nprint(f'using validation set {USE_VAL}')\n\nif USE_VAL:\n    val_df = pd.read_pickle('/kaggle/input/rocket-league-tps-preprocessing/val_df.pickle')\n    (val_features, val_targets) = prep(val_df)\n    del val_df","metadata":{"execution":{"iopub.execute_input":"2022-11-03T06:32:29.256687Z","iopub.status.busy":"2022-11-03T06:32:29.256418Z","iopub.status.idle":"2022-11-03T06:32:29.261788Z","shell.execute_reply":"2022-11-03T06:32:29.260893Z"},"papermill":{"duration":0.014641,"end_time":"2022-11-03T06:32:29.264677","exception":false,"start_time":"2022-11-03T06:32:29.250036","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Takes a set of filenames for pickles with a DataFrame, and produces one epoch with batches of given size per file\n# Not totally generalizable as written, since we also do hardcoded data augmentation / splitting via `prep` fn.\nclass PagedTrainSequence(keras.utils.Sequence):\n    def __init__(self, page_filenames, batch_size, append_df=None, downsample=None):\n        self.page_filenames = page_filenames\n        self.batch_size = batch_size\n        self.append_df = append_df\n        self.downsample = downsample\n        \n        self.page_i = 0\n        self.salt = 1\n        \n        self.features = None\n        self.targets = None\n        self._load_page()\n        # Hack to get around some buggy, or at least counterintuitive keras behavior around on_epoch_end\n        # to make sure we only load a new page once between epochs.\n        self.len_called = False\n\n    def __len__(self):\n        self.len_called = True\n        return len(self.features) // self.batch_size\n\n    def __getitem__(self, batch_i):\n        start = batch_i * self.batch_size\n        end = start + self.batch_size\n        return self.features[start:end], reshape_targets(self.targets[start:end])\n\n    def on_epoch_end(self):\n        if self.len_called:\n            self.page_i = (self.page_i + 1) % len(self.page_filenames)\n            self._load_page()\n            self.len_called = False\n    \n    def _load_page(self):\n        filename = self.page_filenames[self.page_i]\n        print(f'loading file {filename}')\n        df = pd.read_pickle(filename)\n        if self.append_df is not None:\n            L = len(self.append_df) // len(self.page_filenames)\n            df = pd.concat([df, self.append_df[self.page_i*L:(self.page_i+1)*L]])\n        if self.downsample is not None:\n            df = df[::self.downsample]\n        self.features, self.targets = prep(df, salt=f'nacl_{self.salt}')\n        self.salt += 1","metadata":{"execution":{"iopub.execute_input":"2022-11-03T06:32:29.278995Z","iopub.status.busy":"2022-11-03T06:32:29.277302Z","iopub.status.idle":"2022-11-03T06:32:29.288682Z","shell.execute_reply":"2022-11-03T06:32:29.287830Z"},"papermill":{"duration":0.019726,"end_time":"2022-11-03T06:32:29.290664","exception":false,"start_time":"2022-11-03T06:32:29.270938","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 168\nprint(f'using batch_size {BATCH_SIZE}')\npage_filenames = [f'/kaggle/input/rocket-league-tps-preprocessing/train_{i}_df.pickle' for i in range(4)]\nappend_df = None if USE_VAL else pd.read_pickle('/kaggle/input/rocket-league-tps-preprocessing/val_df.pickle')\ninput_seq = PagedTrainSequence(page_filenames, BATCH_SIZE, append_df, downsample=None)","metadata":{"execution":{"iopub.execute_input":"2022-11-03T06:32:29.303825Z","iopub.status.busy":"2022-11-03T06:32:29.302291Z","iopub.status.idle":"2022-11-03T06:34:11.048304Z","shell.execute_reply":"2022-11-03T06:34:11.047206Z"},"papermill":{"duration":101.754798,"end_time":"2022-11-03T06:34:11.050919","exception":false,"start_time":"2022-11-03T06:32:29.296121","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# activation = keras.activations.relu\nactivation = tfa.activations.mish\n\nUSE_LOGITS = False\n# USE_LOGITS = True\nprint(f'using activation {activation} with logits {USE_LOGITS}')\n\ndef concat(layers, axis):\n    return keras.layers.Concatenate(axis=axis)(layers)\n\ndef batched_layers(in_layer, widths, layer_fn):\n    for width in widths:\n        for w in [width] if '__len__' not in dir(width) else width:\n            in_layer = layer_fn(w)(in_layer)\n        in_layer = keras.layers.BatchNormalization()(in_layer)\n    return in_layer\ndef batched_denses(in_layer, widths=[]):\n    return batched_layers(in_layer, widths, lambda width: keras.layers.Dense(width, activation=activation))\ndef batched_conv1ds(in_layer, widths=[]):\n    return batched_layers(in_layer, widths, lambda width: keras.layers.Conv1D(\n        width, kernel_size=1, strides=1, activation=activation))\ndef batched_conv2ds(in_layer, widths=[]):\n    return batched_layers(in_layer, widths, lambda width: keras.layers.Conv2D(\n        width, kernel_size=(1, 1), strides=(1, 1), activation=activation))\ndef batched_conv3ds(in_layer, widths=[]):\n    return batched_layers(in_layer, widths, lambda width: keras.layers.Conv3D(\n        width, kernel_size=(1, 1, 1), strides=(1, 1, 1), activation=activation))\n\ndef convs_and_pool3d(pairs, widths=[]):\n    conv = batched_conv3ds(pairs, widths)\n    return keras.backend.squeeze(keras.layers.MaxPool3D(pool_size=(1, 1, pairs.shape[3]))(conv), axis=3)\ndef convs_and_pool2d(pairs, widths=[]):\n    conv = batched_conv2ds(pairs, widths)\n    return keras.backend.squeeze(keras.layers.MaxPool2D(pool_size=(1, pairs.shape[2]))(conv), axis=2)\n\ndef get_pairs(layer):\n    team_pairs = keras.backend.stack([tf.roll(layer, axis=2, shift=i) for i in range(3)], axis=3)\n    oppo_pairs = tf.roll(team_pairs, axis=2, shift=1)\n    players_tiled = keras.backend.tile(keras.backend.reshape(layer, (-1, 2, 3, 1, layer.shape[-1])), [1, 1, 1, 3, 1])\n    teammates = concat([team_pairs[:, :, :, 1:, :], players_tiled[:, :, :, 1:, :]], axis=4)\n    opponents = concat([oppo_pairs, players_tiled], axis=4)\n    return teammates, opponents\n\ninputs = keras.Input(shape=(54,), name='inputs')\nball = inputs[:, :6]\nplayers = keras.backend.reshape(inputs[:, 6:54], (-1, 2, 3, 8)) # batch_num, team_num, player_num, features (posX3, velX3, boost, demoed)\n\n# Do some dense layers to extract ball features\nball_preprocess = batched_denses(ball, [[48, 36], [36, 32]])\n\ndef to_players_shape(layer):\n    return keras.backend.tile(keras.backend.reshape(layer, (-1, 1, 1, layer.shape[-1])), [1, 2, 3, 1])\n\n# Append player features and ball features together, as well as team number\nteam_muse = players[:, :1, :, :1]\nteams = concat([keras.backend.zeros_like(team_muse), keras.backend.ones_like(team_muse)], axis=1)\n# Get vector diffs from each player pos/vel to ball pos/vel\nball_diffs = to_players_shape(ball) - players[:, :, :, :6]\nplayers_inflated = concat([players, teams, ball_diffs, to_players_shape(ball_preprocess)], axis=3)\n\n# Apply a network-in-network filter over each player\nplayer_filter_widths = [[81, 64], [64, 48], 36]\np_conv = batched_conv2ds(players_inflated, player_filter_widths)\n\n# For each player, find each pair of teammates and each pair of opponents, and append the features of the two players\nteammates, opponents = get_pairs(p_conv)\n# Then apply a new network-in-network filter over each pair of players, one for teammate pairs and one for opponent pairs\n# and use Pooling to compress the most important informance per source player\npair_filter_widths = player_filter_widths\nt_conv = convs_and_pool3d(teammates, pair_filter_widths)\no_conv = convs_and_pool3d(opponents, pair_filter_widths)\n\n# Lump all per-player information and apply one last network-in-network filter over each player, pooling again\np_full = concat([p_conv, t_conv, o_conv], axis=3)\nfull_filter_widths = player_filter_widths\nfull_conv = convs_and_pool2d(p_full, full_filter_widths)\n\n# Append the original ball preprocessing, and just run some dense layers into our prediction layers\nx = concat([ball_preprocess, keras.layers.Flatten()(full_conv)], axis=1)\nx = batched_denses(x, [[81, 64], [64, 48], 42])\ndef output_name(t):\n    return f'outputs_{t}'\noutputs = [keras.layers.Dense(3, activation=None if USE_LOGITS else keras.activations.softmax, name=output_name(t))(x) for t in TARGETS]\n\nmodel = keras.Model(inputs=inputs, outputs=outputs)\nmodel.compile(\n    loss=keras.losses.SparseCategoricalCrossentropy(from_logits=USE_LOGITS),\n    loss_weights={output_name(t): 1 if t == 10 else 0.3 for t in TARGETS},\n    optimizer=keras.optimizers.Adam() # We're using a LearningRateScheduler below\n)\nprint(f'Model with {len(model.layers)} layers and {model.count_params()} params')\nmodel.summary()","metadata":{"execution":{"iopub.execute_input":"2022-11-03T06:34:11.065174Z","iopub.status.busy":"2022-11-03T06:34:11.063709Z","iopub.status.idle":"2022-11-03T06:34:14.689561Z","shell.execute_reply":"2022-11-03T06:34:14.688579Z"},"papermill":{"duration":3.636247,"end_time":"2022-11-03T06:34:14.693002","exception":false,"start_time":"2022-11-03T06:34:11.056755","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"keras.utils.plot_model(model)","metadata":{"execution":{"iopub.execute_input":"2022-11-03T06:34:14.709827Z","iopub.status.busy":"2022-11-03T06:34:14.708281Z","iopub.status.idle":"2022-11-03T06:34:16.440194Z","shell.execute_reply":"2022-11-03T06:34:16.437562Z"},"papermill":{"duration":1.743388,"end_time":"2022-11-03T06:34:16.444853","exception":false,"start_time":"2022-11-03T06:34:14.701465","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time\n\nclass TimeLimit(keras.callbacks.Callback):\n    def __init__(self, max_hours=11, verbose_mod=None):\n        super(TimeLimit, self).__init__()\n        self.max_hours = max_hours\n        self.verbose_mod = verbose_mod\n        self.start_time = time.time()\n\n    def on_epoch_end(self, epoch, logs=None):\n        time_diff_hours = (time.time() - self.start_time) / 3600\n        if self.verbose_mod is not None and (epoch + 1) % self.verbose_mod == 0:\n            print(f'Hours elapsed: {time_diff_hours:.2f}')\n        if time_diff_hours > self.max_hours:\n            print(f'Over the time limit, stopping...')\n            self.stopped_epoch = epoch\n            self.model.stop_training = True","metadata":{"execution":{"iopub.execute_input":"2022-11-03T06:34:16.467376Z","iopub.status.busy":"2022-11-03T06:34:16.467058Z","iopub.status.idle":"2022-11-03T06:34:16.474591Z","shell.execute_reply":"2022-11-03T06:34:16.473671Z"},"papermill":{"duration":0.021106,"end_time":"2022-11-03T06:34:16.476590","exception":false,"start_time":"2022-11-03T06:34:16.455484","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"L = len(page_filenames)\ndef scheduler(epoch, lr):\n    if epoch < L:\n        return 20e-4\n    if epoch < 3 * L:\n        return 15e-4\n    if epoch < 5 * L:\n        return 10e-4\n    if epoch < 7 * L:\n        return 5e-4\n    if epoch < 10 * L:\n        return 2e-4\n    if epoch < 13 * L:\n        return 1e-4\n    return 5e-5","metadata":{"execution":{"iopub.execute_input":"2022-11-03T06:34:16.497807Z","iopub.status.busy":"2022-11-03T06:34:16.497549Z","iopub.status.idle":"2022-11-03T06:34:16.502809Z","shell.execute_reply":"2022-11-03T06:34:16.501921Z"},"papermill":{"duration":0.01848,"end_time":"2022-11-03T06:34:16.505011","exception":false,"start_time":"2022-11-03T06:34:16.486531","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = model.fit(\n    input_seq,\n    epochs=9999,\n    workers=2,\n    use_multiprocessing=True,\n    validation_data=(val_features, reshape_targets(val_targets)) if USE_VAL else None,\n    callbacks=[\n        TimeLimit(max_hours=11.25, verbose_mod=len(page_filenames)),\n        keras.callbacks.LearningRateScheduler(scheduler, verbose=1)\n    ],\n    verbose=2\n)","metadata":{"execution":{"iopub.execute_input":"2022-11-03T06:34:16.526499Z","iopub.status.busy":"2022-11-03T06:34:16.526235Z","iopub.status.idle":"2022-11-03T17:36:42.990886Z","shell.execute_reply":"2022-11-03T17:36:42.988844Z"},"papermill":{"duration":39746.505336,"end_time":"2022-11-03T17:36:43.020056","exception":false,"start_time":"2022-11-03T06:34:16.514720","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.save('model')\n!zip -r model.zip model","metadata":{"execution":{"iopub.execute_input":"2022-11-03T17:36:43.129718Z","iopub.status.busy":"2022-11-03T17:36:43.128895Z","iopub.status.idle":"2022-11-03T17:38:25.711419Z","shell.execute_reply":"2022-11-03T17:38:25.710169Z"},"papermill":{"duration":102.651667,"end_time":"2022-11-03T17:38:25.713766","exception":false,"start_time":"2022-11-03T17:36:43.062099","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib import pyplot as plt\n\nloss_keys = [key for key in history.history.keys() if not key.startswith('val_') and key != 'lr']\nfor key in loss_keys:\n    pd.Series(history.history[key]).plot(label=key)\n    val_key = f'val_{key}'\n    if val_key in history.history:\n        pd.Series(history.history[val_key]).plot(label=val_key)\n    plt.legend()\n    plt.show()","metadata":{"execution":{"iopub.execute_input":"2022-11-03T17:38:25.772769Z","iopub.status.busy":"2022-11-03T17:38:25.767529Z","iopub.status.idle":"2022-11-03T17:38:28.583576Z","shell.execute_reply":"2022-11-03T17:38:28.582591Z"},"papermill":{"duration":2.853731,"end_time":"2022-11-03T17:38:28.585842","exception":false,"start_time":"2022-11-03T17:38:25.732111","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if 'input_seq' in globals():\n    del input_seq\nif 'val_features' in globals():\n    del val_features, val_targets","metadata":{"execution":{"iopub.execute_input":"2022-11-03T17:38:28.638082Z","iopub.status.busy":"2022-11-03T17:38:28.637185Z","iopub.status.idle":"2022-11-03T17:38:28.643355Z","shell.execute_reply":"2022-11-03T17:38:28.642438Z"},"papermill":{"duration":0.03351,"end_time":"2022-11-03T17:38:28.645416","exception":false,"start_time":"2022-11-03T17:38:28.611906","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_pickle('/kaggle/input/rocket-league-tps-preprocessing/test_df.pickle')","metadata":{"execution":{"iopub.execute_input":"2022-11-03T17:38:28.690008Z","iopub.status.busy":"2022-11-03T17:38:28.688415Z","iopub.status.idle":"2022-11-03T17:38:30.651616Z","shell.execute_reply":"2022-11-03T17:38:30.650576Z"},"papermill":{"duration":1.988047,"end_time":"2022-11-03T17:38:30.654285","exception":false,"start_time":"2022-11-03T17:38:28.666238","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess_features(df):\n    df.drop([col for col in df.columns if col.startswith('boost')], axis=1, inplace=True)\n    \n    for p in range(6):\n        if f'p{p}_demoed' not in df:\n            df.insert(list(df.columns).index(f'p{p}_boost') + 1, f'p{p}_demoed', df[f'p{p}_boost'].isna())\n        x = -20 + 20 * (p % 3)\n        y = -20 + 40 * (p // 3)\n        cond = df[f'p{p}_boost'].isna()\n        df.loc[cond, [f'p{p}_pos_x', f'p{p}_vel_x', f'p{p}_pos_y', f'p{p}_vel_y', f'p{p}_pos_z', f'p{p}_vel_z', f'p{p}_boost']] = \\\n            [x, x, y, y, 100, 10, 0]","metadata":{"execution":{"iopub.execute_input":"2022-11-03T17:38:30.700285Z","iopub.status.busy":"2022-11-03T17:38:30.699697Z","iopub.status.idle":"2022-11-03T17:38:30.707383Z","shell.execute_reply":"2022-11-03T17:38:30.706439Z"},"papermill":{"duration":0.03199,"end_time":"2022-11-03T17:38:30.709407","exception":false,"start_time":"2022-11-03T17:38:30.677417","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\npreprocess_features(test_df)","metadata":{"execution":{"iopub.execute_input":"2022-11-03T17:38:30.757123Z","iopub.status.busy":"2022-11-03T17:38:30.755716Z","iopub.status.idle":"2022-11-03T17:38:30.927165Z","shell.execute_reply":"2022-11-03T17:38:30.925896Z"},"papermill":{"duration":0.199629,"end_time":"2022-11-03T17:38:30.930014","exception":false,"start_time":"2022-11-03T17:38:30.730385","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_test_preds(features, model):\n    model_outputs = model.predict(features, verbose=2)\n    if USE_LOGITS:\n        model_outputs = keras.activations.softmax(keras.backend.constant(model_outputs))\n    target_probs = model_outputs[TARGETS.index(10)] if len(TARGETS) > 1 else model_outputs\n    return target_probs","metadata":{"execution":{"iopub.execute_input":"2022-11-03T17:38:30.980883Z","iopub.status.busy":"2022-11-03T17:38:30.980145Z","iopub.status.idle":"2022-11-03T17:38:30.986243Z","shell.execute_reply":"2022-11-03T17:38:30.985088Z"},"papermill":{"duration":0.034175,"end_time":"2022-11-03T17:38:30.988984","exception":false,"start_time":"2022-11-03T17:38:30.954809","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_preds_df(model, num_samples=1):\n    pred_sum_df = None\n    for i in range(num_samples):\n        df = test_df.copy()\n        augment_data(df, shuffle=False, append_flip_y=True, salt=f'test_nacl_{i}')\n        preds = get_test_preds(df.drop(['id', 'flip_y'], axis=1).to_numpy('float32'), model)\n        df2 = pd.DataFrame({'A': preds[:, 0], 'N': preds[:, 1], 'B': preds[:, 2]})\n        df2.loc[df.flip_y != False, ['A', 'B']] = df2.loc[df.flip_y != False, ['B', 'A']].values\n        if pred_sum_df is None:\n            pred_sum_df = df2\n        else:\n            pred_sum_df += df2\n    preds_df = pred_sum_df / num_samples\n    return pd.DataFrame({'id': test_df.id, 'team_A_scoring_within_10sec': preds_df.A, 'team_B_scoring_within_10sec': preds_df.B})\n\ndef make_submission(model, outfile, num_samples=1):\n    df = get_preds_df(model, num_samples)\n    df.to_csv(outfile, index=False)\n    print(f'wrote {outfile}')\n    return df","metadata":{"execution":{"iopub.execute_input":"2022-11-03T17:38:31.036187Z","iopub.status.busy":"2022-11-03T17:38:31.035881Z","iopub.status.idle":"2022-11-03T17:38:31.045938Z","shell.execute_reply":"2022-11-03T17:38:31.045023Z"},"papermill":{"duration":0.036063,"end_time":"2022-11-03T17:38:31.048021","exception":false,"start_time":"2022-11-03T17:38:31.011958","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df = make_submission(model, 'submission.csv', 16)","metadata":{"execution":{"iopub.execute_input":"2022-11-03T17:38:31.094713Z","iopub.status.busy":"2022-11-03T17:38:31.093880Z","iopub.status.idle":"2022-11-03T18:15:44.788202Z","shell.execute_reply":"2022-11-03T18:15:44.787173Z"},"papermill":{"duration":2233.743334,"end_time":"2022-11-03T18:15:44.813557","exception":false,"start_time":"2022-11-03T17:38:31.070223","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.team_A_scoring_within_10sec.hist(bins=np.linspace(0, 1, 100))","metadata":{"execution":{"iopub.execute_input":"2022-11-03T18:15:44.860215Z","iopub.status.busy":"2022-11-03T18:15:44.859907Z","iopub.status.idle":"2022-11-03T18:15:45.266358Z","shell.execute_reply":"2022-11-03T18:15:45.265490Z"},"papermill":{"duration":0.431572,"end_time":"2022-11-03T18:15:45.268339","exception":false,"start_time":"2022-11-03T18:15:44.836767","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.team_B_scoring_within_10sec.hist(bins=np.linspace(0, 1, 100))","metadata":{"execution":{"iopub.execute_input":"2022-11-03T18:15:45.315020Z","iopub.status.busy":"2022-11-03T18:15:45.314206Z","iopub.status.idle":"2022-11-03T18:15:45.714052Z","shell.execute_reply":"2022-11-03T18:15:45.713171Z"},"papermill":{"duration":0.424956,"end_time":"2022-11-03T18:15:45.715942","exception":false,"start_time":"2022-11-03T18:15:45.290986","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.team_A_scoring_within_10sec[submission_df.team_A_scoring_within_10sec > 0.5].hist(bins=np.linspace(0.5, 1, 100))","metadata":{"execution":{"iopub.execute_input":"2022-11-03T18:15:45.763645Z","iopub.status.busy":"2022-11-03T18:15:45.762002Z","iopub.status.idle":"2022-11-03T18:15:46.099773Z","shell.execute_reply":"2022-11-03T18:15:46.098753Z"},"papermill":{"duration":0.363138,"end_time":"2022-11-03T18:15:46.101703","exception":false,"start_time":"2022-11-03T18:15:45.738565","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.team_B_scoring_within_10sec[submission_df.team_B_scoring_within_10sec > 0.5].hist(bins=np.linspace(0.5, 1, 100))","metadata":{"execution":{"iopub.execute_input":"2022-11-03T18:15:46.148922Z","iopub.status.busy":"2022-11-03T18:15:46.148110Z","iopub.status.idle":"2022-11-03T18:15:46.501533Z","shell.execute_reply":"2022-11-03T18:15:46.500515Z"},"papermill":{"duration":0.379339,"end_time":"2022-11-03T18:15:46.503745","exception":false,"start_time":"2022-11-03T18:15:46.124406","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}