{"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":"markdown","source":"remaining task\n- preprocess S (done)  \n- treat 2017 data (done)  \n- set same channel size (done)  \n- crps loss (done but not work well)   \n- augumentation and TTA   \n- scheduler (done but not work well)  \n- test part    ","metadata":{}},{"cell_type":"markdown","source":"## feature","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport os\nfrom collections import defaultdict\nimport torch\n# from torch import nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, TensorDataset\nfrom torch.utils.tensorboard import SummaryWriter\n\nimport jax\nimport jax.numpy as jnp\nimport flax\nfrom flax import linen as nn\nfrom flax.training import train_state, checkpoints\nimport optax\nfrom typing import Any\nfrom tqdm.auto import tqdm\n\nimport seaborn as sns\nimport glob\nimport io\nimport gc\n\nfrom sklearn.metrics import accuracy_score\nfrom sklearn.metrics import confusion_matrix\nfrom sklearn.metrics import plot_confusion_matrix\n\nimport time\npd.options.display.max_columns = 999","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:41:56.269030Z","iopub.execute_input":"2022-07-27T05:41:56.269352Z","iopub.status.idle":"2022-07-27T05:42:05.887755Z","shell.execute_reply.started":"2022-07-27T05:41:56.269280Z","shell.execute_reply":"2022-07-27T05:42:05.886775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#features\nimport numpy as np\nfrom numba import jit\n\n\ndef create_features(df):\n    xysdir_o = df[(df.IsOnOffense == True) & (df.IsRusher == False)][['X','Y','X_S','Y_S']].values\n    xysdir_rush = df[df.IsRusher == True][['X','Y','X_S','Y_S']].values\n    xysdir_d = df[df.IsOnOffense == False][['X','Y','X_S','Y_S']].values\n    \n    off_x = np.array(df[(df.IsOnOffense == True) & (df.IsRusher == False)].groupby('PlayId')['X'].apply(np.array))\n    def_x = np.array(df[(df.IsOnOffense == False) ].groupby('PlayId')['X'].apply(np.array))\n    off_y = np.array(df[(df.IsOnOffense == True) & (df.IsRusher == False)].groupby('PlayId')['Y'].apply(np.array))\n    def_y = np.array(df[(df.IsOnOffense == False) ].groupby('PlayId')['Y'].apply(np.array))\n    off_sx = np.array(df[(df.IsOnOffense == True) & (df.IsRusher == False)].groupby('PlayId')['X_S'].apply(np.array))\n    def_sx = np.array(df[(df.IsOnOffense == False) ].groupby('PlayId')['X_S'].apply(np.array))\n    off_sy = np.array(df[(df.IsOnOffense == True) & (df.IsRusher == False)].groupby('PlayId')['Y_S'].apply(np.array))\n    def_sy = np.array(df[(df.IsOnOffense == False) ].groupby('PlayId')['Y_S'].apply(np.array))\n    \n    player_vector = []\n    for play in range(len(off_x)):\n        player_feat = player_feature(off_x[play],def_x[play],off_y[play],def_y[play],off_sx[play],def_sx[play],\n                                     off_sy[play],def_sy[play],xysdir_rush[play])\n        player_vector.append(player_feat)\n    \n    return np.array(player_vector)\n\n    \ndef player_feature(off_x,def_x,off_y,def_y,off_sx,def_sx,off_sy,def_sy,xysdir_rush):\n    if(len(off_x<10)):\n        off_x = np.pad(off_x,(10-len(off_x),0), 'mean' )\n        off_y = np.pad(off_y,(10-len(off_y),0), 'mean' )\n        off_sx = np.pad(off_sx,(10-len(off_sx),0), 'mean' )\n        off_sy = np.pad(off_sy,(10-len(off_sy),0), 'mean' )\n    if(len(def_x<11)):\n        def_x = np.pad(def_x,(11-len(def_x),0), 'mean' )\n        def_y = np.pad(def_y,(11-len(def_y),0), 'mean' )\n        def_sx = np.pad(def_sx,(11-len(def_sx),0), 'mean' )\n        def_sy = np.pad(def_sy,(11-len(def_sy),0), 'mean' )\n\n    dist_def_off_x = def_x.reshape(-1,1)-off_x.reshape(1,-1)\n    dist_def_off_sx = def_sx.reshape(-1,1)-off_sx.reshape(1,-1)\n    dist_def_off_y = def_y.reshape(-1,1)-off_y.reshape(1,-1)\n    dist_def_off_sy = def_sy.reshape(-1,1)-off_sy.reshape(1,-1)\n    dist_def_rush_x = def_x.reshape(-1,1)-np.repeat(xysdir_rush[0],10).reshape(1,-1)\n    dist_def_rush_y = def_y.reshape(-1,1)-np.repeat(xysdir_rush[1],10).reshape(1,-1)\n    dist_def_rush_sx = def_sx.reshape(-1,1)-np.repeat(xysdir_rush[2],10).reshape(1,-1)\n    dist_def_rush_sy = def_sy.reshape(-1,1)-np.repeat(xysdir_rush[3],10).reshape(1,-1)\n    def_sx = np.repeat(def_sx,10).reshape(11,-1)\n    def_sy = np.repeat(def_sy,10).reshape(11,-1)\n    feats = [dist_def_off_x, dist_def_off_sx, dist_def_off_y, dist_def_off_sy, dist_def_rush_x, dist_def_rush_y,\n            dist_def_rush_sx, dist_def_rush_sy, def_sx, def_sy]\n    \n    return np.stack(feats)\n\n\ndef get_def_speed(df):\n    df_cp = df[~df.IsOnOffense].copy()\n    speed = df_cp[\"S\"].T.values\n    speed = speed.reshape(-1, 1, 1, 11) \n    speed = np.repeat(speed, 10, axis=2)\n\n    return speed\n\n\ndef get_dist(df, col1, col2, type=\"defence\"):\n    if type == \"defence\":\n        df_cp = df[~df.IsOnOffense].copy()\n    elif type == \"offence\":\n        df_cp = df[df.IsOnOffense].copy()\n    dist = np.linalg.norm(df_cp[col1].values - df_cp[col2].values, axis=1)\n    dist = dist.T\n    dist = dist.reshape(-1, 1, 1, 11)\n    dist = np.repeat(dist, 10, axis=2)\n\n    return dist\n\n\n\ndef dist_def_off(df, n_train, cols):\n    off_x = np.array(df[(df.IsOnOffense) & (~train.IsRusher)].groupby('PlayId')['X'].apply(np.array))\n    def_x = np.array(df[(~df.IsOnOffense) ].groupby('PlayId')['X'].apply(np.array))\n    off_y = np.array(df[(df.IsOnOffense) & (~train.IsRusher)].groupby('PlayId')['Y'].apply(np.array))\n    def_y = np.array(df[(~df.IsOnOffense) ].groupby('PlayId')['Y_S'].apply(np.array))\n    off_xs = np.array(df[(df.IsOnOffense) & (~train.IsRusher)].groupby('PlayId')['X_S'].apply(np.array))\n    def_xs = np.array(df[(~df.IsOnOffense) ].groupby('PlayId')['X_S'].apply(np.array))\n    off_ys = np.array(df[(df.IsOnOffense) & (~train.IsRusher)].groupby('PlayId')['Y_S'].apply(np.array))\n    def_ys = np.array(df[(~df.IsOnOffense) ].groupby('PlayId')['Y_S'].apply(np.array))\n    feats = []\n    for play in range(len(off_x)):\n        dist_x = off_x[play].reshape(-1, 1) - def_x[play].reshape(1, -1)\n        dist_y = off_y[play].reshape(-1, 1) - def_y[play].reshape(1, -1)\n        dist = np.concatenate([dist_x[:, :, np.newaxis], dist_y[:, :, np.newaxis]], axis=2)\n        dist_xy = np.linalg.norm(dist.astype(np.float64), axis=2)\n        dist_xs = off_xs[play].reshape(-1, 1) - def_xs[play].reshape(1, -1)\n        dist_ys = off_ys[play].reshape(-1, 1) - def_ys[play].reshape(1, -1)\n        dist = np.concatenate([dist_xs[:, :, np.newaxis], dist_ys[:, :, np.newaxis]], axis=2)\n        dist_xys = np.linalg.norm(dist.astype(np.float64), axis=2)\n        feats.append(np.concatenate([dist_xy[np.newaxis, :], dist_xys[np.newaxis, :]], axis=0))\n    return np.array(feats)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-27T05:42:05.889897Z","iopub.execute_input":"2022-07-27T05:42:05.890473Z","iopub.status.idle":"2022-07-27T05:42:06.608432Z","shell.execute_reply.started":"2022-07-27T05:42:05.890445Z","shell.execute_reply":"2022-07-27T05:42:06.607488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## preprocess","metadata":{}},{"cell_type":"code","source":"#preprocess\nimport numpy as np\n\n\ndef reorient(df, flip_left, aug=False):\n    df['ToLeft'] = df.PlayDirection == \"left\"\n    \n    df.loc[df.VisitorTeamAbbr == \"ARI\", 'VisitorTeamAbbr'] = \"ARZ\"\n    df.loc[df.HomeTeamAbbr == \"ARI\", 'HomeTeamAbbr'] = \"ARZ\"\n\n    df.loc[df.VisitorTeamAbbr == \"BAL\", 'VisitorTeamAbbr'] = \"BLT\"\n    df.loc[df.HomeTeamAbbr == \"BAL\", 'HomeTeamAbbr'] = \"BLT\"\n\n    df.loc[df.VisitorTeamAbbr == \"CLE\", 'VisitorTeamAbbr'] = \"CLV\"\n    df.loc[df.HomeTeamAbbr == \"CLE\", 'HomeTeamAbbr'] = \"CLV\"\n\n    df.loc[df.VisitorTeamAbbr == \"HOU\", 'VisitorTeamAbbr'] = \"HST\"\n    df.loc[df.HomeTeamAbbr == \"HOU\", 'HomeTeamAbbr'] = \"HST\"\n\n    df['TeamOnOffense'] = \"home\"\n    df.loc[df.PossessionTeam != df.HomeTeamAbbr, 'TeamOnOffense'] = \"away\"\n    df['IsOnOffense'] = df.Team == df.TeamOnOffense  # Is player on offense?\n    df['YardLine_std'] = 100 - df.YardLine\n    df.loc[df.FieldPosition.fillna('') == df.PossessionTeam, 'YardLine_std'] = \\\n        df.loc[df.FieldPosition.fillna('') == df.PossessionTeam, 'YardLine']\n    df.loc[df.ToLeft, 'X'] = 120 - df.loc[df.ToLeft, 'X']\n    df.loc[df.ToLeft, 'Y'] = 160 / 3 - df.loc[df.ToLeft, 'Y']\n    df.loc[df.ToLeft, 'Orientation'] = np.mod(180 + df.loc[df.ToLeft, 'Orientation'], 360)\n    df['Dir'] = 90 - df.Dir\n    df.loc[df.ToLeft, 'Dir'] = np.mod(180 + df.loc[df.ToLeft, 'Dir'], 360)\n    df.loc[df.IsOnOffense, 'Dir'] = df.loc[df.IsOnOffense, 'Dir'].fillna(0).values\n    df.loc[~df.IsOnOffense, 'Dir'] = df.loc[~df.IsOnOffense, 'Dir'].fillna(180).values\n\n    df['IsRusher'] = df['NflId'] == df['NflIdRusher']\n    if flip_left:\n        tmp = df[df['IsRusher']].copy()\n        # df['left'] = df.Y < 160/6\n        tmp['left'] = tmp.Dir < 0\n        df = df.merge(tmp[['PlayId', 'left']], how='left', on='PlayId')\n        df['Y'] = df.Y\n        df.loc[df[\"left\"], 'Y'] = 160 / 3 - df.loc[df[\"left\"], 'Y']\n        df['Dir'] = df.Dir\n        df.loc[df[\"left\"], 'Dir'] = np.mod(- df.loc[df[\"left\"], 'Dir'], 360)\n        df.drop('left', axis=1, inplace=True)\n\n    df[\"S\"] = df[\"Dis\"] * 10\n    df['X_dir'] = np.cos((np.pi / 180) * df.Dir)\n    df['Y_dir'] = np.sin((np.pi / 180) * df.Dir)\n    df['X_S'] = df.X_dir * df.S\n    df['Y_S'] = df.Y_dir * df.S\n    df['X_A'] = df.X_dir * df.A\n    df['Y_A'] = df.Y_dir * df.A\n    #df.loc[df['Season'] == 2017, 'S'] = (df['S'][df['Season'] == 2017] - 2.4355) / 1.2930 * 1.4551 + 2.7570\n    df['time_step'] = 0.0\n    df = df.sort_values(by=['PlayId', 'IsOnOffense', 'IsRusher', 'Y']).reset_index(drop=True)\n    \n    if aug:\n        df_aug = df.copy()\n        df_aug[\"Y\"] = 53.3 - df_aug[\"Y\"]\n        df = df.append(df_aug).reset_index()\n    \n    return df\n\n\ndef merge_rusherfeats(df):\n    rusher_feats = df[df['NflId'] == df['NflIdRusher']].drop_duplicates()\n    rusher_feats = rusher_feats[[\"PlayId\", \"X\", \"Y\", \"X_S\", \"Y_S\"]]\n    rusher_feats = rusher_feats.rename(\n        columns={\"X\": \"Rusher_X\", \"Y\": \"Rusher_Y\", \"X_S\": \"Rusher_X_S\", \"Y_S\": \"Rusher_Y_S\"})\n    df = df.merge(rusher_feats, how=\"left\", on=\"PlayId\")\n\n    return df\n\ndef scaling(feats, sctype=\"standard\"):\n    v1 = []\n    v2 = []\n    for i in range(feats.shape[1]):\n        feats_ = feats[:, i, :]\n        if sctype == \"standard\":\n            mean_ = np.mean(feats_)\n            std_ = np.std(feats_)\n            feats[:, i, :] -= mean_\n            feats[:, i, :] /= std_\n            v1.append(mean_)\n            v2.append(std_)\n        elif sctype == \"minmax\":\n            max_ = np.max(feats_)\n            min_ = np.min(feats_)\n            feats[:, i, :] = (feats_ - min_) / (max_ - min_)\n            v1.append(max_)\n            v2.append(min_)\n\n    return feats, v1, v2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-27T05:42:06.610043Z","iopub.execute_input":"2022-07-27T05:42:06.610441Z","iopub.status.idle":"2022-07-27T05:42:06.637187Z","shell.execute_reply.started":"2022-07-27T05:42:06.610406Z","shell.execute_reply":"2022-07-27T05:42:06.636080Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## metrics","metadata":{}},{"cell_type":"code","source":"\ndef crps(y_loss_val, y_pred):\n    y_true = jnp.array(jnp.clip(jnp.cumsum(y_loss_val, axis=1), 0, 1))\n    y_pred = jnp.array(jnp.clip(jnp.cumsum(y_pred, axis=1), 0, 1))\n    y_pred.at[:, :99-30].set(0.0)\n    y_pred.at[:, 50+99:].set(1.0)\n    val_s = ((y_true - y_pred) ** 2).sum(axis=1).sum(axis=0) / (199 * y_loss_val.shape[0])\n    crps = jnp.round(val_s, 6)\n    \n    return crps","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:49:42.129460Z","iopub.execute_input":"2022-07-27T05:49:42.129934Z","iopub.status.idle":"2022-07-27T05:49:42.144574Z","shell.execute_reply.started":"2022-07-27T05:49:42.129891Z","shell.execute_reply":"2022-07-27T05:49:42.143293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## model","metadata":{}},{"cell_type":"code","source":"googlenet_kernel_init = nn.initializers.kaiming_normal()\n\nclass ZooNet(nn.Module):\n    num_classes : int\n    act_fn : callable\n    # num_features : int\n\n    @nn.compact\n    def __call__(self, x, train=True):\n        \n        #First Conv block\n        x = self.act_fn(x)\n        x = nn.Conv(160, kernel_size=[1,1])(x)\n        x = self.act_fn(x)\n        x = nn.Conv(128, kernel_size=[1,1])(x)\n        x = self.act_fn(x)\n        x = nn.BatchNorm()(x, use_running_average=train)\n        x = jnp.squeeze(nn.avg_pool(x, (1,11)),2)\n\n\n        #Second Conv block\n        x = nn.Conv(160, kernel_size=(1,))(x)\n        x = self.act_fn(x)\n        x = nn.Conv(96, kernel_size=(1,))(x)\n        x = self.act_fn(x)\n        x = nn.Conv(96, kernel_size=(1,))(x)\n        x = self.act_fn(x)\n        x = nn.BatchNorm()(x, use_running_average=train)\n        x = jnp.squeeze(nn.avg_pool(x, (1,10)),1)\n\n        #Linear block\n        x = nn.Dense(96)(x)\n        x = self.act_fn(x)\n        x = nn.Dense(256)(x)\n        x = self.act_fn(x)\n        x = nn.Dense(self.num_classes)(x)\n\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:42:06.656186Z","iopub.execute_input":"2022-07-27T05:42:06.657085Z","iopub.status.idle":"2022-07-27T05:42:06.670630Z","shell.execute_reply.started":"2022-07-27T05:42:06.657047Z","shell.execute_reply":"2022-07-27T05:42:06.669620Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## utils","metadata":{}},{"cell_type":"code","source":"class TrainState(train_state.TrainState):\n    # A simple extension of TrainState to also include batch statistics\n    batch_stats: Any","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:42:06.672040Z","iopub.execute_input":"2022-07-27T05:42:06.672642Z","iopub.status.idle":"2022-07-27T05:42:06.688762Z","shell.execute_reply.started":"2022-07-27T05:42:06.672607Z","shell.execute_reply":"2022-07-27T05:42:06.687729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CHECKPOINT_PATH = \"./\"\nclass TrainerModule:\n\n    def __init__(self, \n                 model_name : str, \n                 model_class : nn.Module, \n                 model_hparams : dict, \n                 optimizer_name : str, \n                 optimizer_hparams : dict, \n                 exmp_imgs : Any, \n                 seed=42):\n        \"\"\"\n        Module for summarizing all training functionalities for classification on CIFAR10.\n        \n        Inputs:\n            model_name - String of the class name, used for logging and saving\n            model_class - Class implementing the neural network\n            model_hparams - Hyperparameters of the model, used as input to model constructor\n            optimizer_name - String of the optimizer name, supporting ['sgd', 'adam', 'adamw']\n            optimizer_hparams - Hyperparameters of the optimizer, including learning rate as 'lr'\n            exmp_imgs - Example imgs, used as input to initialize the model\n            seed - Seed to use in the model initialization\n        \"\"\"\n        super().__init__()\n        self.model_name = model_name\n        self.model_class = model_class\n        self.model_hparams = model_hparams\n        self.optimizer_name = optimizer_name\n        self.optimizer_hparams = optimizer_hparams\n        self.seed = seed\n        # Create empty model. Note: no parameters yet\n        self.model = self.model_class(**self.model_hparams)\n        # Prepare logging\n        self.log_dir = os.path.join(CHECKPOINT_PATH, self.model_name)\n        self.logger = SummaryWriter(log_dir=self.log_dir)\n        # Create jitted training and eval functions\n        self.create_functions()\n        # Initialize model\n        self.init_model(exmp_imgs)\n\n    def create_functions(self):\n        # Function to calculate the classification loss and accuracy for a model\n        def calculate_loss(params, batch_stats, batch, train):\n            imgs, labels = batch\n            labels_onehot = jax.nn.one_hot(labels, num_classes=199)\n            # Run model. During training, we need to update the BatchNorm statistics.\n            outs = self.model.apply({'params': params, 'batch_stats': batch_stats}, \n                                    imgs,\n                                    train=train,\n                                    mutable=['batch_stats'],\n                                   rngs={'dropout': jax.random.PRNGKey(2)})\n            logits, new_model_state = outs if train else (outs, None)\n            loss = optax.softmax_cross_entropy(logits, labels_onehot).mean()\n            acc = (logits.argmax(axis=-1) == labels).mean()\n            return loss, (acc, new_model_state, logits, labels)\n        # Training function\n        def train_step(state, batch):\n            loss_fn = lambda params: calculate_loss(params, state.batch_stats, batch, True)\n            # Get loss, gradients for loss, and other outputs of loss function\n            ret, grads = jax.value_and_grad(loss_fn, has_aux=True)(state.params)  \n            loss, acc, new_model_state, _, _ = ret[0], *ret[1]\n            # Update parameters and batch statistics\n            state = state.apply_gradients(grads=grads, batch_stats=new_model_state['batch_stats'])\n            return state, loss, acc\n        # Eval function\n        def eval_step(state, batch):\n            # Return the accuracy for a single batch\n            _, (acc, _, logits, labels) = calculate_loss(state.params, state.batch_stats, batch, True)\n            return acc, logits, labels\n        # jit for efficiency\n        self.train_step = jax.jit(train_step)\n        self.eval_step = jax.jit(eval_step)\n\n    def init_model(self, exmp_imgs):\n        # Initialize model\n        init_rng = {'params': jax.random.PRNGKey(0), 'dropout': jax.random.PRNGKey(1)}\n        variables = self.model.init(init_rng, exmp_imgs, True)\n        self.init_params, self.init_batch_stats = variables['params'], variables['batch_stats']\n        self.state = None\n        \n    def init_optimizer(self, num_epochs, num_steps_per_epoch):\n        # Initialize learning rate schedule and optimizer\n        if self.optimizer_name.lower() == 'adam':\n            opt_class = optax.adam\n        elif self.optimizer_name.lower() == 'adamw':\n            opt_class = optax.adamw\n        elif self.optimizer_name.lower() == 'sgd':\n            opt_class = optax.sgd\n        else:\n            assert False, f'Unknown optimizer \"{opt_class}\"'\n        # We decrease the learning rate by a factor of 0.1 after 60% and 85% of the training\n        lr_schedule = optax.piecewise_constant_schedule(\n            init_value=self.optimizer_hparams.pop('lr'),\n            boundaries_and_scales=\n                {int(num_steps_per_epoch*num_epochs*0.6): 0.1,\n                 int(num_steps_per_epoch*num_epochs*0.85): 0.1}\n        )\n        # Clip gradients at max value, and evt. apply weight decay\n        transf = [optax.clip(1.0)]\n        if opt_class == optax.sgd and 'weight_decay' in self.optimizer_hparams:  # wd is integrated in adamw\n            transf.append(optax.add_decayed_weights(self.optimizer_hparams.pop('weight_decay')))\n        optimizer = optax.chain(\n            *transf,\n            opt_class(lr_schedule, **self.optimizer_hparams)\n        )\n        # Initialize training state\n        self.state = TrainState.create(apply_fn=self.model.apply, \n                                       params=self.init_params if self.state is None else self.state.params,\n                                       batch_stats=self.init_batch_stats if self.state is None else self.state.batch_stats,\n                                       tx=optimizer)\n\n    def train_model(self, train_loader, val_loader, num_epochs=200):\n        # Train model for defined number of epochs\n        # We first need to create optimizer and the scheduler for the given number of epochs\n        self.init_optimizer(num_epochs, len(train_loader))\n        # Track best eval accuracy\n        best_eval = 0.0\n        for epoch_idx in tqdm(range(1, num_epochs+1)):\n            self.train_epoch(epoch=epoch_idx)\n            eval_acc, logits, labels = self.eval_model(val_loader)\n            \n            self.logger.add_scalar('val/acc', eval_acc, global_step=epoch_idx)\n#             print( eval_acc)\n\n#             print(np.concatenate(labels))\n#             print( flax.linen.softmax(jnp.array(np.concatenate(logits))))\n            y_crps = jax.nn.one_hot(np.concatenate(labels), num_classes=199)\n            eval_crps = crps(y_crps, flax.linen.softmax(jnp.array(np.concatenate(logits))))\n            print(f\"epoch {epoch_idx}, valid crps: {eval_crps:.6f}\")\n            if epoch_idx == 50:\n                self.save_model(step=epoch_idx)\n            # if epoch_idx % 2 == 0:\n                \n            #     self.logger.add_scalar('val/acc', eval_acc, global_step=epoch_idx)\n            #     if eval_acc >= best_eval:\n            #         best_eval = eval_acc\n            #         self.save_model(step=epoch_idx)\n            #     self.logger.flush()\n\n    def train_epoch(self, epoch):\n        # Train model for one epoch, and log avg loss and accuracy\n        metrics = defaultdict(list)\n        for batch in tqdm(train_loader, desc='Training', leave=False):\n            self.state, loss, acc = self.train_step(self.state, batch)\n            metrics['loss'].append(loss)\n            metrics['acc'].append(acc)\n        for key in metrics:\n            avg_val = np.stack(jax.device_get(metrics[key])).mean()\n            self.logger.add_scalar('train/'+key, avg_val, global_step=epoch)\n\n    def eval_model(self, data_loader):\n        # Test model on all images of a data loader and return avg loss\n        correct_class, count = 0, 0\n        logits_list = []\n        labels_list = []\n        for batch in data_loader:\n            acc, logits, labels = self.eval_step(self.state, batch)\n            logits_list.append(np.array(logits))\n            labels_list.append(np.array(labels))\n            correct_class += acc * batch[0].shape[0]\n            count += batch[0].shape[0]\n        eval_acc = (correct_class / count).item()\n        logits_list = np.array(logits_list).flatten()\n        labels_list = np.array(labels_list).flatten()\n#         eval_auc = fast_auc(labels_list, logits_list)\n        return eval_acc, logits_list, labels_list\n\n    def save_model(self, step=0):\n        # Save current model at certain training iteration\n        checkpoints.save_checkpoint(ckpt_dir=self.log_dir, \n                                    target={'params': self.state.params, \n                                            'batch_stats': self.state.batch_stats}, \n                                    step=step,\n                                   overwrite=True)\n\n    def load_model(self, pretrained=False):\n        # Load model. We use different checkpoint for pretrained models\n        if not pretrained:\n            state_dict = checkpoints.restore_checkpoint(ckpt_dir=self.log_dir, target=None)\n        else:\n            state_dict = checkpoints.restore_checkpoint(ckpt_dir=os.path.join(CHECKPOINT_PATH, f'{self.model_name}.ckpt'), target=None)\n        self.state = TrainState.create(apply_fn=self.model.apply, \n                                       params=state_dict['params'],\n                                       batch_stats=state_dict['batch_stats'],\n                                       rngs=state_dict['dropout'],\n                                       tx=self.state.tx if self.state else optax.sgd(0.1)   # Default optimizer\n                                      )\n\n    def checkpoint_exists(self):\n        # Check whether a pretrained model exist for this autoencoder\n        return os.path.isfile(os.path.join(CHECKPOINT_PATH, f'{self.model_name}.ckpt'))","metadata":{"execution":{"iopub.status.busy":"2022-07-27T06:13:31.849561Z","iopub.execute_input":"2022-07-27T06:13:31.850079Z","iopub.status.idle":"2022-07-27T06:13:31.895982Z","shell.execute_reply.started":"2022-07-27T06:13:31.850030Z","shell.execute_reply":"2022-07-27T06:13:31.894900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# #utils\n# import os\n# import random\n\n# import numpy as np\n# import torch\n\n\n# def seed_torch(seed=1029):\n#     random.seed(seed)\n#     os.environ[\"PYTHONHASHSEED\"] = str(seed)\n#     np.random.seed(seed)\n#     torch.manual_seed(seed)\n#     torch.cuda.manual_seed_all(seed)\n#     torch.backends.cudnn.deterministic = True","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-27T05:42:06.730572Z","iopub.execute_input":"2022-07-27T05:42:06.731060Z","iopub.status.idle":"2022-07-27T05:42:06.748403Z","shell.execute_reply.started":"2022-07-27T05:42:06.731025Z","shell.execute_reply":"2022-07-27T05:42:06.747501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## logger","metadata":{}},{"cell_type":"code","source":"#logger\nimport logging\nimport sys\n\nLOGGER = logging.getLogger()\nFORMATTER = logging.Formatter(\"%(asctime)s - %(levelname)s - %(message)s\")\n\n\ndef setup_logger(out_file=None, stderr=True, stderr_level=logging.INFO, file_level=logging.DEBUG):\n    LOGGER.handlers = []\n    LOGGER.setLevel(min(stderr_level, file_level))\n\n    if stderr:\n        handler = logging.StreamHandler(sys.stderr)\n        handler.setFormatter(FORMATTER)\n        handler.setLevel(stderr_level)\n        LOGGER.addHandler(handler)\n\n    if out_file is not None:\n        handler = logging.FileHandler(out_file)\n        handler.setFormatter(FORMATTER)\n        handler.setLevel(file_level)\n        LOGGER.addHandler(handler)\n\n    LOGGER.info(\"logger set up\")\n    return LOGGER","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-27T05:42:06.749550Z","iopub.execute_input":"2022-07-27T05:42:06.751493Z","iopub.status.idle":"2022-07-27T05:42:06.760764Z","shell.execute_reply.started":"2022-07-27T05:42:06.751457Z","shell.execute_reply":"2022-07-27T05:42:06.759788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## trainer","metadata":{}},{"cell_type":"code","source":"def train_classifier(*args, num_epochs=200, **kwargs):\n    # Create a trainer module with specified hyperparameters\n    trainer = TrainerModule(*args, **kwargs)\n#     if not trainer.checkpoint_exists():  # Skip training if pretrained model exists\n    trainer.train_model(train_loader, val_loader, num_epochs=num_epochs)\n#         trainer.load_model()\n#     else:\n#         trainer.load_model(pretrained=True)\n    # Test trained model\n    val_acc = trainer.eval_model(val_loader)\n    # test_acc = trainer.eval_model(test_loader)\n    return trainer, {'val': val_acc}","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:42:06.764483Z","iopub.execute_input":"2022-07-27T05:42:06.764817Z","iopub.status.idle":"2022-07-27T05:42:06.774852Z","shell.execute_reply.started":"2022-07-27T05:42:06.764787Z","shell.execute_reply":"2022-07-27T05:42:06.773833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def collate_fn(batch):\n    data_list, label_list = [], []\n    for _data, _label in batch:\n        data_list.append(np.array(_data))\n        label_list.append(np.array(_label))\n    return np.array(data_list), np.array(label_list)","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:42:06.776376Z","iopub.execute_input":"2022-07-27T05:42:06.777081Z","iopub.status.idle":"2022-07-27T05:42:06.786881Z","shell.execute_reply.started":"2022-07-27T05:42:06.777046Z","shell.execute_reply":"2022-07-27T05:42:06.785739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# loss","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport sys\nimport time\n\nimport json\nimport pandas as pd\nimport numpy as np\nfrom contextlib import contextmanager\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.preprocessing import LabelEncoder, StandardScaler\nimport torch\nfrom torch.utils.data import DataLoader, TensorDataset\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\n#from trainer import train_one_epoch, validate\n\n# ===============\n# Constants\n# ===============\nDATA_DIR = \"../input/nfl-big-data-bowl-2020-final-stage-2-dataset\"\nTRAIN_PATH = os.path.join(DATA_DIR, \"train.csv\")\nTEST_PATH = os.path.join(DATA_DIR, \"test.csv\")\nSOLUTION_PATH = os.path.join(DATA_DIR, \"solution.csv\")\nLOGGER_PATH = \"log.txt\"\nTARGET_COLUMNS = 'Yards'\nN_CLASSES = 199\n\n# ===============\n# Settings\n# ===============\nSEED = 69160\ndevice = \"cpu\"\nN_SPLITS = 5\nBATCH_SIZE = 64\nAUG = False\n\nepochs = 50\nEXP_ID = \"exp2_10feat\"\n\nsetup_logger(out_file=LOGGER_PATH)\n\n@contextmanager\ndef timer(name):\n    t0 = time.time()\n    yield\n    LOGGER.info('[{}] done in {} s'.format(name, round(time.time() - t0, 2)))\n\n\nwith timer('load data'):\n    train = pd.read_csv(TRAIN_PATH, dtype={'WindSpeed': 'object'})\n    test = pd.read_csv(TEST_PATH, dtype={'WindSpeed': 'object'})\n    solution = pd.read_csv(SOLUTION_PATH)\n    test = pd.merge(test,solution)\n    train = pd.concat([train,test])\n#     train = train[train['Season'] != 2017]\n    game_id = train[\"GameId\"][::22].values\n    season = train[\"Season\"][::22].values\n    y_mae = train[TARGET_COLUMNS][::22].values\n    y_mae = np.where(y_mae < -16, -16, y_mae)\n    y_mae = np.where(y_mae > 50, 50, y_mae)\n    y_crps = np.zeros((y_mae.shape[0], 199))\n    for idx, target in enumerate(list(y_mae)):\n        y_crps[idx][99 + target] = 1\n\n    n_train = len(train) // 22\n    n_train_2017 = len(train[train.Season == 2017]) // 22\n    n_df = len(train)\n\nwith timer('create features'):\n    train = reorient(train, flip_left=True, aug=AUG)\n    train = merge_rusherfeats(train)\n\n    #x_def_speed = get_def_speed(train)\n    #x_def_rusher_dist = get_dist(train, [\"X\", \"Y\"], [\"Rusher_X\", \"Rusher_Y\"], \"defence\")\n    #x_def_rusher_speeddist = get_dist(train, [\"X_S\", \"Y_S\"], [\"Rusher_X_S\", \"Rusher_Y_S\"], \"defence\")\n    #x_def_off_dist = dist_def_off(train, n_train*2 if AUG else n_train, [[\"X\", \"Y\"], [\"X_S\", \"Y_S\"]])\n    #x = np.concatenate([x_def_speed, x_def_rusher_dist, x_def_rusher_speeddist, x_def_off_dist], axis=1)\n    x = create_features(train)\n    x, sc_mean, sc_std = scaling(x)\n    x = np.swapaxes(x,1,3)\n\nwith timer('split data'):\n    if AUG:\n        x_aug = x[n_df:]\n        x = x[:n_df]\n        x_aug_2017 = x_aug[season==2017]\n        x_aug_usage = x_aug[season!=2017]\n    \n    x_2017, y_crps_2017, y_mae_2017 = x[season==2017], y_crps[season==2017], y_mae[season==2017]\n    x_usage, y_crps_usage, y_mae_usage = x[season!=2017], y_crps[season!=2017], y_mae[season!=2017]\n    folds = GroupKFold(n_splits=N_SPLITS).split(y_mae_usage, y_mae_usage, groups=game_id[season!=2017])\n\nwith timer('train'):\n    scores = []\n    for n_fold, (train_idx, val_idx) in enumerate(folds):\n        with timer('create model'):\n            x_train, y_train = x_usage[train_idx], y_crps_usage[train_idx]\n            x_val, y_val, y_val_mae = x_usage[val_idx], y_crps_usage[val_idx], y_mae_usage[val_idx]\n            \n            # add 2017 data\n            x_train = np.concatenate([x_train, x_2017], axis=0)\n            y_train = np.concatenate([y_train, y_crps_2017], axis=0)\n            if AUG:\n                x_train = np.concatenate([x_train, x_aug_usage[train_idx], x_aug_2017], axis=0)\n                y_train = np.concatenate([y_train, y_train], axis=0)\n                val_dataset_aug = TensorDataset(torch.tensor(x_aug_usage[val_idx], dtype=torch.float32), torch.tensor(y_val, dtype=torch.float32))\n                val_loader_aug = DataLoader(val_dataset_aug, batch_size=BATCH_SIZE, shuffle=False, num_workers=0, pin_memory=True)\n                \n            y_train = np.argmax(y_train,axis=1)\n            y_val = np.argmax(y_val,axis=1)\n            \n            train_dataset = TensorDataset(torch.tensor(x_train, dtype=torch.float32), torch.tensor(y_train, dtype=torch.float32))\n            train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=0, pin_memory=True,collate_fn=collate_fn)\n\n            val_dataset = TensorDataset(torch.tensor(x_val, dtype=torch.float32), torch.tensor(y_val, dtype=torch.float32))\n            val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=0, pin_memory=True,collate_fn=collate_fn)\n            del train_dataset, val_dataset\n            gc.collect()\n            \n            googlenet_trainer, googlenet_results = train_classifier(model_class=ZooNet,\n                                                        model_name=\"ZooNet\",\n                                                        model_hparams={\"num_classes\": N_CLASSES,\n                                                                       \"act_fn\": nn.relu},\n                                                        optimizer_name=\"adam\",\n                                                        optimizer_hparams={\"lr\": 0.001},\n                                                        exmp_imgs=jax.device_put(\n                                                            next(iter(train_loader))[0]),\n                                                        num_epochs=50)\n#             del batch_stats\n\n#             model = CnnModel(num_classes=N_CLASSES)\n#             #model.to(device)\n\n#             num_steps = len(x_train) // BATCH_SIZE\n#             criterion = torch.nn.CrossEntropyLoss()\n# #             criterion = CRPSLoss()\n#             optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\n# #             scheduler = OneCycleLR(optimizer, num_steps=num_steps, lr_range=(5e-4, 1e-3))\n#             scheduler = None\n\n#         with timer('train fold{}'.format(n_fold)):\n#             best_score = 999\n#             best_epoch = 0\n#             y_pred = np.zeros_like(y_crps)\n#             for epoch in range(1, epochs + 1):\n#                 seed_torch(SEED + epoch)\n\n#                 LOGGER.info(\"Starting {} epoch...\".format(epoch))\n#                 tr_loss = train_one_epoch(model, train_loader, criterion, optimizer, device, scheduler=scheduler)\n#                 LOGGER.info('Mean train loss: {}'.format(round(tr_loss, 5)))\n\n#                 val_pred, y_true, val_loss = validate(model, val_loader, criterion, device)\n#                 if AUG:\n#                     val_pred_aug, _, val_loss_aug = validate(model, val_loader_aug, criterion, device)\n#                     val_loss = [val_loss, val_loss_aug]\n#                     val_pred = val_pred + val_pred_aug\n#                 score = crps(y_val, val_pred)\n#                 LOGGER.info('Mean valid loss: {} score: {}'.format(round(val_loss, 5), round(score, 5)))\n#                 if score < best_score:\n#                     best_score = score\n#                     best_epoch = epoch\n#                     torch.save(model.state_dict(), '{}_fold{}.pth'.format(EXP_ID, n_fold))\n#                     y_pred[val_idx] = val_pred\n            \n#             scores.append(best_score)\n#             LOGGER.info(\"best score={} on epoch={} fold={}\".format(best_score, best_epoch, n_fold))\n#     LOGGER.info(\"score avg={}, score fold0={}, score fold1={}, score fold2={}, score fold3={}, score fold4={}\".format(\n#         np.mean(scores), scores[0], scores[1], scores[2], scores[3], scores[4]))","metadata":{"execution":{"iopub.status.busy":"2022-07-27T06:17:52.383206Z","iopub.execute_input":"2022-07-27T06:17:52.383747Z","iopub.status.idle":"2022-07-27T06:26:14.254663Z","shell.execute_reply.started":"2022-07-27T06:17:52.383715Z","shell.execute_reply":"2022-07-27T06:26:14.253676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}