{"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 gc\nimport os\nimport random\nfrom typing import List, Tuple, Optional, Union\n\nimport numpy as np\nimport pandas as pd\nimport torch\n# import torch.nn as nn\nfrom collections import defaultdict\nfrom torch.utils.data import TensorDataset\nfrom sklearn.preprocessing import StandardScaler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.tensorboard import SummaryWriter\nfrom tqdm import tqdm\n\n\n\nfrom sklearn.model_selection import KFold\nfrom joblib import Parallel, delayed\nfrom sklearn.decomposition import PCA\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts\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","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-27T05:30:06.845841Z","iopub.execute_input":"2022-07-27T05:30:06.846266Z","iopub.status.idle":"2022-07-27T05:30:17.223604Z","shell.execute_reply.started":"2022-07-27T05:30:06.846184Z","shell.execute_reply":"2022-07-27T05:30:17.222471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_train = pd.read_csv('../input/xgb-fraud-with-magic-0-9600/X_train.csv')\nX_test = pd.read_csv('../input/xgb-fraud-with-magic-0-9600/X_test.csv')\ny_train = pd.read_csv('../input/xgb-fraud-with-magic-0-9600/y_train.csv',header=None)[1]","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:30:17.225979Z","iopub.execute_input":"2022-07-27T05:30:17.226847Z","iopub.status.idle":"2022-07-27T05:31:20.087592Z","shell.execute_reply.started":"2022-07-27T05:30:17.226795Z","shell.execute_reply":"2022-07-27T05:31:20.085912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"emb_dim: int = 50\nbatch_size: int = 1024\nmodel_type: str = 'mlp'\nmlp_dropout: float = 0.0\nmlp_hidden: int = 64\nmlp_bn: bool = False\ncnn_hidden: int = 128\ncnn_channel1: int = 32\ncnn_channel2: int = 32\ncnn_channel3: int = 32\ncnn_kernel1: int = 5\ncnn_celu: bool = False\ncnn_weight_norm: bool = False\ndropout_emb: bool = 0.0\nlr: float = 1e-3\nweight_decay: float = 0.0\nmodel_path: str = 'fold_{}.pth'\nscaler_type: str = 'standard'\noutput_dir: str = 'artifacts'\nscheduler_type: str = 'onecycle'\noptimizer_type: str = 'adam'\nmax_lr: float = 0.01\nepochs: int = 30\nseed: int = 42\nn_pca: int = -1\nbatch_double_freq: int = 50\ncnn_dropout: float = 0.1\nna_cols: bool = True\ncnn_leaky_relu: bool = False\npatience: int = 8\nfactor: float = 0.5","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:31:20.089877Z","iopub.execute_input":"2022-07-27T05:31:20.090223Z","iopub.status.idle":"2022-07-27T05:31:20.101322Z","shell.execute_reply.started":"2022-07-27T05:31:20.090190Z","shell.execute_reply":"2022-07-27T05:31:20.100470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NN_VALID_TH = 0.185\nNN_MODEL_TOP_N = 3\nTAB_MODEL_TOP_N = 3\nENSEMBLE_METHOD = 'mean'\nNN_NUM_MODELS = 10\nTABNET_NUM_MODELS = 5\n","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:31:20.103851Z","iopub.execute_input":"2022-07-27T05:31:20.104331Z","iopub.status.idle":"2022-07-27T05:31:20.116256Z","shell.execute_reply.started":"2022-07-27T05:31:20.104301Z","shell.execute_reply":"2022-07-27T05:31:20.115129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed=42):\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(seed)\n    torch.backends.cudnn.deterministic = True","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:31:20.118044Z","iopub.execute_input":"2022-07-27T05:31:20.118396Z","iopub.status.idle":"2022-07-27T05:31:20.130222Z","shell.execute_reply.started":"2022-07-27T05:31:20.118364Z","shell.execute_reply":"2022-07-27T05:31:20.128953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MLP(nn.Module):\n    num_features: int\n    hidden_size: int\n    n_categories: List[int]\n    emb_dim: int = 10\n    dropout_cat: float = 0.2\n    channel_1: int = 256\n    channel_2: int = 512\n    channel_3: int = 512\n    dropout_top: float = 0.1\n    dropout_mid: float = 0.3\n    dropout_bottom: float = 0.2\n    weight_norm: bool = True\n    two_stage: bool = True\n    celu: bool = True\n    kernel1: int = 5\n    leaky_relu: bool = False\n    \n        \n\n    \n    @nn.compact\n    def __call__(self, x, train=True):\n        x = nn.BatchNorm(use_running_average= not train)(x)\n        x = nn.relu(nn.Dense(100)(x))\n        x = nn.relu(nn.Dense(100)(x))\n        x = nn.relu(nn.Dense(100)(x))\n        x = nn.Dropout(self.dropout_top, deterministic= not train)(x)\n        x = nn.Dense(2)(x)\n        return x\n\n    \n# model = MLP(num_features=264,\n#            hidden_size=cnn_hidden,\n#            n_categories=[0],\n#            channel_1=cnn_channel1,\n#                 channel_2=cnn_channel2,\n#                 channel_3=cnn_channel3,\n#                 two_stage=False,\n#                 kernel1=cnn_kernel1,\n#                 celu=cnn_celu,\n#                 dropout_top=cnn_dropout,\n#                 dropout_mid=cnn_dropout,\n#                 dropout_bottom=cnn_dropout,\n#                 weight_norm=cnn_weight_norm,\n#                 leaky_relu=cnn_leaky_relu)\n# batch = jnp.ones((10, 264))  # (N, H, W, C) format\n# variables = model.init({'params': jax.random.PRNGKey(0), 'dropout': jax.random.PRNGKey(1)}, batch)\n# output = model.apply(variables, batch, mutable=['batch_stats'], rngs={'dropout': jax.random.PRNGKey(1)})    \n# one_hot = jax.nn.one_hot(y_tr[:10], 2)\n# loss = jnp.mean(optax.softmax_cross_entropy(logits=output[0], labels=one_hot))","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:31:20.131601Z","iopub.execute_input":"2022-07-27T05:31:20.132452Z","iopub.status.idle":"2022-07-27T05:31:20.151248Z","shell.execute_reply.started":"2022-07-27T05:31:20.132418Z","shell.execute_reply":"2022-07-27T05:31:20.150261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CNN(nn.Module):\n    num_features: int\n    hidden_size: int\n    n_categories: List[int]\n    emb_dim: int = 10\n    dropout_cat: float = 0.2\n    channel_1: int = 256\n    channel_2: int = 512\n    channel_3: int = 512\n    dropout_top: float = 0.1\n    dropout_mid: float = 0.3\n    dropout_bottom: float = 0.2\n    weight_norm: bool = True\n    two_stage: bool = True\n    celu: bool = True\n    kernel1: int = 5\n    leaky_relu: bool = False\n    \n        \n\n    \n    @nn.compact\n    def __call__(self, x, train=True):\n        \n        num_targets = 2\n        cha_1_reshape = int(self.hidden_size / self.channel_1)\n        cha_po_1 = int(self.hidden_size / self.channel_1 / 2)\n        cha_po_2 = int(self.hidden_size / self.channel_1 / 2 / 2) * self.channel_3\n        \n#         dropout_rng = self.make_rng('dropout')\n        \n        x = nn.BatchNorm(use_running_average= not train)(x)\n        x = nn.Dropout(self.dropout_top, deterministic= not train)(x)\n        x = nn.Dense(self.hidden_size)(x)\n        x = nn.relu(x)\n\n        \n        x = x.reshape(x.shape[0], self.channel_1, cha_1_reshape)\n\n        x = nn.BatchNorm(use_running_average= not train)(x)\n        x = nn.Dropout(self.dropout_top, deterministic= not train)(x)\n        x = np.swapaxes(x,1,2)\n        x = nn.Conv(self.channel_2, kernel_size=(self.kernel1,), strides=1, padding=self.kernel1 // 2)(x)\n        x = np.swapaxes(x,2,1)\n        x = nn.relu(x)\n\n        strides = (x.shape[2]//cha_po_1)\n        x = np.swapaxes(x,1,2)\n        x = nn.avg_pool(x, (1,),strides=(strides,))\n        x = np.swapaxes(x,2,1)\n        x = nn.BatchNorm(use_running_average= not train)(x)\n        x = nn.Dropout(self.dropout_top, deterministic= not train)(x)\n        x = nn.relu(x)\n        x = np.swapaxes(x,1,2)\n        x = nn.avg_pool(x,(4,),strides=(2,),padding=((1,1,),))\n        x = np.swapaxes(x,1,2)\n        x = x.reshape((-1,x.shape[1]))\n        x = nn.BatchNorm(use_running_average= not train)(x)\n        x = nn.Dropout(self.dropout_bottom, deterministic= not train)(x)\n        x = nn.Dense(num_targets)(x)\n\n        \n        return x\n\n\n\n\n# model = CNN(num_features=264,\n#            hidden_size=cnn_hidden,\n#            n_categories=[0],\n#            channel_1=cnn_channel1,\n#                 channel_2=cnn_channel2,\n#                 channel_3=cnn_channel3,\n#                 two_stage=False,\n#                 kernel1=cnn_kernel1,\n#                 celu=cnn_celu,\n#                 dropout_top=cnn_dropout,\n#                 dropout_mid=cnn_dropout,\n#                 dropout_bottom=cnn_dropout,\n#                 weight_norm=cnn_weight_norm,\n#                 leaky_relu=cnn_leaky_relu)\n# batch = jnp.ones((10, 264))  # (N, H, W, C) format\n# variables = model.init({'params': jax.random.PRNGKey(0), 'dropout': jax.random.PRNGKey(1)}, batch)\n# output = model.apply(variables, batch, mutable=['batch_stats'], rngs={'dropout': jax.random.PRNGKey(1)})\n# one_hot = jax.nn.one_hot(y_tr[:10], 2)\n# loss = jnp.mean(optax.softmax_cross_entropy(logits=output[0], labels=one_hot))\n\n","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:31:20.152958Z","iopub.execute_input":"2022-07-27T05:31:20.153282Z","iopub.status.idle":"2022-07-27T05:31:20.181016Z","shell.execute_reply.started":"2022-07-27T05:31:20.153251Z","shell.execute_reply":"2022-07-27T05:31:20.179961Z"},"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:31:20.182978Z","iopub.execute_input":"2022-07-27T05:31:20.184048Z","iopub.status.idle":"2022-07-27T05:31:20.199888Z","shell.execute_reply.started":"2022-07-27T05:31:20.184014Z","shell.execute_reply":"2022-07-27T05:31:20.198495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_train['dummy_emb'] = np.random.randint(1)\nX_num = X_train[X_train.columns[~X_train.columns.isin(['dummy_emb'])]]\nscaler = StandardScaler()\nscaled_features = scaler.fit_transform(X_num.values)\nX_num = pd.DataFrame(scaled_features, index=X_num.index, columns=X_num.columns)\n\n","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:31:20.202527Z","iopub.execute_input":"2022-07-27T05:31:20.203651Z","iopub.status.idle":"2022-07-27T05:31:23.178567Z","shell.execute_reply.started":"2022-07-27T05:31:20.203600Z","shell.execute_reply":"2022-07-27T05:31:23.177477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TabularDataset(Dataset):\n    def __init__(self, x_num: np.ndarray, y: Optional[np.ndarray]):\n        super().__init__()\n        self.x_num = x_num\n        self.y = y\n\n    def __len__(self):\n        return len(self.x_num)\n\n    def __getitem__(self, idx):\n        if self.y is None:\n            return self.x_num[idx]\n        else:\n            return self.x_num[idx], self.y[idx]","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:31:23.181875Z","iopub.execute_input":"2022-07-27T05:31:23.182196Z","iopub.status.idle":"2022-07-27T05:31:23.189468Z","shell.execute_reply.started":"2022-07-27T05:31:23.182167Z","shell.execute_reply":"2022-07-27T05:31:23.188298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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:31:23.191045Z","iopub.execute_input":"2022-07-27T05:31:23.191358Z","iopub.status.idle":"2022-07-27T05:31:23.203806Z","shell.execute_reply.started":"2022-07-27T05:31:23.191331Z","shell.execute_reply":"2022-07-27T05:31:23.202821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fast_auc(y_true, y_prob):\n    y_true = np.asarray(y_true)\n    y_true = y_true[np.argsort(y_prob)]\n    nfalse = 0\n    auc = 0\n    n = len(y_true)\n    for i in range(n):\n        y_i = y_true[i]\n        nfalse += (1 - y_i)\n        auc += y_i * nfalse\n    auc /= (nfalse * (n - nfalse))\n    return auc","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:31:23.205494Z","iopub.execute_input":"2022-07-27T05:31:23.206201Z","iopub.status.idle":"2022-07-27T05:31:23.216415Z","shell.execute_reply.started":"2022-07-27T05:31:23.206148Z","shell.execute_reply":"2022-07-27T05:31:23.215102Z"},"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=2)\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            eval_auc = fast_auc(np.concatenate(labels), flax.linen.softmax(jnp.array(np.concatenate(logits)))[:,1])\n            print(f\"epoch {epoch_idx}, valid auc: {eval_auc:.3f}\")\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-27T05:31:23.218223Z","iopub.execute_input":"2022-07-27T05:31:23.218613Z","iopub.status.idle":"2022-07-27T05:31:23.264577Z","shell.execute_reply.started":"2022-07-27T05:31:23.218579Z","shell.execute_reply":"2022-07-27T05:31:23.263264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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, logits_list, labels_list = trainer.eval_model(val_loader)\n    # test_acc = trainer.eval_model(test_loader)\n    return trainer, {'val': val_acc}, logits_list, labels_list","metadata":{"execution":{"iopub.status.busy":"2022-07-27T05:31:23.266423Z","iopub.execute_input":"2022-07-27T05:31:23.266783Z","iopub.status.idle":"2022-07-27T05:31:23.276215Z","shell.execute_reply.started":"2022-07-27T05:31:23.266752Z","shell.execute_reply":"2022-07-27T05:31:23.275210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fold_score = []\n\nkf = KFold(n_splits=10)\n\nfor fold, (train_idx, test_idx) in enumerate(kf.split(X_train)):\n    \n    \n    X_tr = np.array(X_num.loc[train_idx])\n    y_tr = np.array(y_train.iloc[train_idx]).flatten()\n \n    X_va = np.array(X_num.loc[test_idx])\n    y_va = np.array(y_train.iloc[test_idx]).flatten()\n\n    \n    train_dataset = TabularDataset(X_tr, y_tr)\n    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, pin_memory=True,collate_fn=collate_fn)\n\n    val_dataset = TabularDataset(X_va, y_va)\n    val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, pin_memory=True,collate_fn=collate_fn)  \n    \n    googlenet_trainer, googlenet_results, logits, labels = train_classifier(model_class=CNN,\n                                            model_name=\"MLP\",\n                                            model_hparams={\n                                                'num_features':264,\n           'hidden_size':cnn_hidden,\n           'n_categories':[0],\n           'channel_1':cnn_channel1,\n                'channel_2':cnn_channel2,\n                'channel_3':cnn_channel3,\n                'two_stage':False,\n                'kernel1':cnn_kernel1,\n                'celu':cnn_celu,\n                'dropout_top':cnn_dropout,\n                'dropout_mid':cnn_dropout,\n                'dropout_bottom':cnn_dropout,\n                'weight_norm':cnn_weight_norm,\n                'leaky_relu':cnn_leaky_relu},\n                                            optimizer_name=\"adam\",\n                                            optimizer_hparams={\"lr\": lr},\n                                            exmp_imgs=jax.device_put(\n                                                next(iter(train_loader))[0]),\n                                            num_epochs=30)\n    \n# #     rng = jax.random.PRNGKey(0)\n# #     state = create_train_state(rng)\n    \n#     rng = jax.random.PRNGKey(0)\n#     state = create_train_state(rng)\n# # #     for images, labels in val_dataset:\n        \n# # #         grads, loss, accuracy = apply_model(state, images, labels)\n    \n    \n#     epoch_loss = []\n#     epoch_accuracy = []\n#     for epoch in tqdm(range(10)):\n#         for images, labels in train_dataset:\n\n#             grads, loss, accuracy = apply_model(state, images, labels)\n\n# #             state = update_model(state, grads)\n#         epoch_loss.append(loss)\n#         epoch_accuracy.append(accuracy)\n#         print(loss)","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2022-07-27T05:37:36.548191Z","iopub.execute_input":"2022-07-27T05:37:36.549840Z","iopub.status.idle":"2022-07-27T05:38:52.008628Z","shell.execute_reply.started":"2022-07-27T05:37:36.549777Z","shell.execute_reply":"2022-07-27T05:38:52.006969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# submission_df.to_csv('submission.csv', index=False)","metadata":{},"execution_count":null,"outputs":[]}]}