{"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":"## summary\n\n* Reading large data at once results in insufficient memory.\n* Pytorch has a Dataset class to avoid this.\n* Ten csv files per epoch are read in sequence to avoid running out of memory.\n\n\n* There is an IterableDataset class for sequential access.\n* Looking at the source code, Dataset is random access, so execution speed is slow when csv data is targeted.\n* In this case, IterableDataset is concatenated with ChainDataset.(The effectiveness has not yet been verified.)\n\n\n* In order to reduce the cast time, which takes up most of the data loading time, the csv data was pre-processed as a pickle.\n* By adjusting the batch_size of Pytorch, a lot of training can be done in one epoch, and training can be completed in only a few epochs.\n\n\n* By combining these adjustments, relatively fast learning is possible even with neural nets.","metadata":{}},{"cell_type":"markdown","source":"## Setup","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"%%capture\n!pip install wandb\n!pip install pytorch_lightning","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:32:19.993733Z","iopub.execute_input":"2022-10-09T12:32:19.994281Z","iopub.status.idle":"2022-10-09T12:32:40.527205Z","shell.execute_reply.started":"2022-10-09T12:32:19.994168Z","shell.execute_reply":"2022-10-09T12:32:40.525905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import wandb\ntry:\n    # add-ons -> secrets -> set your wandb api key\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    secret_value_0 = user_secrets.get_secret(\"wandb_api\")\n    wandb.login(key=secret_value_0)\n    anony = None\nexcept:\n    anony = \"must\"\n    print('If you want to use your W&B account, go to Add-ons -> Secrets and provide your W&B access token. Use the Label name as wandb_api. \\nGet your W&B access token from here: https://wandb.ai/authorize')","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:32:40.529623Z","iopub.execute_input":"2022-10-09T12:32:40.530347Z","iopub.status.idle":"2022-10-09T12:32:43.394099Z","shell.execute_reply.started":"2022-10-09T12:32:40.530304Z","shell.execute_reply":"2022-10-09T12:32:43.393052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"project_name = \"tps2210\"\nwandb.init(project=project_name)","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:32:43.396117Z","iopub.execute_input":"2022-10-09T12:32:43.396565Z","iopub.status.idle":"2022-10-09T12:32:50.171873Z","shell.execute_reply.started":"2022-10-09T12:32:43.396524Z","shell.execute_reply":"2022-10-09T12:32:50.170948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Imports and Setting","metadata":{}},{"cell_type":"code","source":"# common\nimport os\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom collections import Counter\nimport time, gc, string, math\nfrom tqdm.notebook import tqdm\nimport warnings\nimport shutil\nfrom collections import defaultdict\nimport heapq\nimport datetime\nimport random\nfrom collections import OrderedDict\nimport glob\nimport copy\nfrom itertools import permutations, chain\n\n# sklearn\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.model_selection import KFold\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.preprocessing import StandardScaler\nfrom sklearn.metrics import roc_auc_score\n\n# pytorch\nimport torch\nimport torch.nn as nn\nfrom torch.nn import functional as F\nfrom torch.utils.data import DataLoader, Dataset, IterableDataset, ChainDataset\nfrom torch import optim\nfrom torch.optim import lr_scheduler\n\n# pytorch lightning\nimport pytorch_lightning as pl\nfrom pytorch_lightning.loggers import WandbLogger","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:32:50.177312Z","iopub.execute_input":"2022-10-09T12:32:50.178277Z","iopub.status.idle":"2022-10-09T12:32:55.680198Z","shell.execute_reply.started":"2022-10-09T12:32:50.178236Z","shell.execute_reply":"2022-10-09T12:32:55.679145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.set_option('display.max_columns', 200)\npd.set_option('display.max_rows', 200)\nwarnings.simplefilter('ignore')\npl.seed_everything(42)","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:32:55.681521Z","iopub.execute_input":"2022-10-09T12:32:55.682434Z","iopub.status.idle":"2022-10-09T12:32:55.703478Z","shell.execute_reply.started":"2022-10-09T12:32:55.682394Z","shell.execute_reply":"2022-10-09T12:32:55.702479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## read csv and FE","metadata":{}},{"cell_type":"code","source":"dtypes_df = pd.read_csv('/kaggle/input/tabular-playground-series-oct-2022/train_dtypes.csv')\ndtypes = {k: v for (k, v) in zip(dtypes_df.column, dtypes_df.dtype)}\ntrain = pd.read_csv('/kaggle/input/tabular-playground-series-oct-2022/train_0.csv', dtype=dtypes)\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:32:55.707810Z","iopub.execute_input":"2022-10-09T12:32:55.709372Z","iopub.status.idle":"2022-10-09T12:33:29.977510Z","shell.execute_reply.started":"2022-10-09T12:32:55.709334Z","shell.execute_reply":"2022-10-09T12:33:29.976259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del train\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:33:29.979162Z","iopub.execute_input":"2022-10-09T12:33:29.979545Z","iopub.status.idle":"2022-10-09T12:33:30.261342Z","shell.execute_reply.started":"2022-10-09T12:33:29.979506Z","shell.execute_reply":"2022-10-09T12:33:30.259897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dtypes_df = pd.read_csv('/kaggle/input/tabular-playground-series-oct-2022/test_dtypes.csv')\ndtypes = {k: v for (k, v) in zip(dtypes_df.column, dtypes_df.dtype)}\ntest = pd.read_csv('/kaggle/input/tabular-playground-series-oct-2022/test.csv', dtype=dtypes)\ntest.head()","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:33:30.266819Z","iopub.execute_input":"2022-10-09T12:33:30.268417Z","iopub.status.idle":"2022-10-09T12:33:39.448264Z","shell.execute_reply.started":"2022-10-09T12:33:30.268377Z","shell.execute_reply":"2022-10-09T12:33:39.447242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_col_list = ['team_A_scoring_within_10sec', 'team_B_scoring_within_10sec']\ntrain_col_list = list(test.columns)\ntrain_col_list.remove('id')\nprint(train_col_list)","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:33:39.450584Z","iopub.execute_input":"2022-10-09T12:33:39.451565Z","iopub.status.idle":"2022-10-09T12:33:39.459265Z","shell.execute_reply.started":"2022-10-09T12:33:39.451525Z","shell.execute_reply":"2022-10-09T12:33:39.458050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv_path_list = []\nfor i in range(10):\n    train_csv_path = '/kaggle/input/tabular-playground-series-oct-2022/train_{}.csv'.format(i)\n    train_csv_path_list.append(train_csv_path)","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:33:39.463941Z","iopub.execute_input":"2022-10-09T12:33:39.464246Z","iopub.status.idle":"2022-10-09T12:33:39.471502Z","shell.execute_reply.started":"2022-10-09T12:33:39.464220Z","shell.execute_reply":"2022-10-09T12:33:39.470276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## preprocessing\n\nPre-processing to save memory and speed up data loading.\n\nAfter na-filling, pickle the data (eliminating the time required for casting).","metadata":{}},{"cell_type":"code","source":"# preprocess test\n# fill na -> scaling\ntest = test.fillna(test.median())\nskaler_of = defaultdict(lambda: None)\nfor col in train_col_list:\n    skaler = skaler_of[col]\n    if skaler is None:\n        skaler = StandardScaler()\n        skaler.fit(test[[col]])\n        skaler_of[col] = skaler\n    test[col] = skaler.transform(test[[col]])\n    \n# save pickle\ntest.to_pickle('test.pkl')","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:33:39.472888Z","iopub.execute_input":"2022-10-09T12:33:39.473553Z","iopub.status.idle":"2022-10-09T12:33:42.018045Z","shell.execute_reply.started":"2022-10-09T12:33:39.473513Z","shell.execute_reply":"2022-10-09T12:33:42.016965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# preprocess train\ndtypes_df = pd.read_csv('/kaggle/input/tabular-playground-series-oct-2022/train_dtypes.csv')\ndtypes = {k: v for (k, v) in zip(dtypes_df.column, dtypes_df.dtype)}\nfor i, train_csv_path in tqdm(enumerate(train_csv_path_list)):\n    train = pd.read_csv(train_csv_path, dtype=dtypes)\n    train = train.fillna(train.median())\n    for col in train_col_list:\n        skaler = skaler_of[col]\n        train[col] = skaler.transform(train[[col]])\n\n    # shuffle\n    train = train.sample(frac=1, random_state=42)\n    \n    # not work int8. Loss need float32\n    train[target_col_list] = train[target_col_list].astype(np.float32)\n    \n    # save pickle\n    train.to_pickle('train_{}.pkl'.format(i))\n\n    del train\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:33:42.021063Z","iopub.execute_input":"2022-10-09T12:33:42.021853Z","iopub.status.idle":"2022-10-09T12:38:48.673028Z","shell.execute_reply.started":"2022-10-09T12:33:42.021808Z","shell.execute_reply":"2022-10-09T12:38:48.672013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pkl_path_list = []\nfor i in range(len(train_csv_path_list)):\n    pkl_path = 'train_{}.pkl'.format(i)\n    pkl_path_list.append(pkl_path)","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:38:48.678860Z","iopub.execute_input":"2022-10-09T12:38:48.681186Z","iopub.status.idle":"2022-10-09T12:38:48.689640Z","shell.execute_reply.started":"2022-10-09T12:38:48.681149Z","shell.execute_reply":"2022-10-09T12:38:48.688649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Pytorch Lightning","metadata":{}},{"cell_type":"code","source":"class CFG:\n    num_workers = 1  # colabは4, kaggleは2?, IterableDatasetは1?\n    weight_decay=1e-4\n    print_epoch_freq=1\n    max_epochs=10\n    batch_size=1024*16\n    lr = 1e-2\n    min_lr = 1e-6\n    is_debug = True\n    dataaug = False\n\nif CFG.is_debug:\n    CFG.max_epochs=5  # 後続でfoldを0にしている","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:41:18.343542Z","iopub.execute_input":"2022-10-09T12:41:18.343922Z","iopub.status.idle":"2022-10-09T12:41:18.352763Z","shell.execute_reply.started":"2022-10-09T12:41:18.343889Z","shell.execute_reply":"2022-10-09T12:41:18.351642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TrainDataset(IterableDataset):\n    def __init__(self, train_col_list, target_col, pkl_path, dataaug=False, dataaug_sample=0.1):\n        self.train_col_list = train_col_list\n        self.target_col = target_col\n        self.pkl_path = pkl_path\n        self.perm_list = list(permutations([0, 1, 2], 3))[1:]\n        self.pi_col_pos_of = defaultdict(list)  # train_col_list index of A team\n        self.xcol_pos = []\n        self.ycol_pos = []\n        self.boost_even_pos = []\n        self.boost_odd_pos = []\n        self.dataaug = dataaug\n        self.dataaug_sample = dataaug_sample\n\n        for idx, col in enumerate(train_col_list):\n            for i in range(6):\n                if 'p{}_'.format(i) in col:\n                    self.pi_col_pos_of[i].append(idx)\n            if '_x' in col:\n                self.xcol_pos.append(idx)\n            if '_y' in col:\n                self.ycol_pos.append(idx)\n            for i in range(6):\n                if i % 2 == 0:\n                    if 'boost{}_timer'.format(i) in col:\n                        self.boost_even_pos.append(idx)\n                else:\n                    if 'boost{}_timer'.format(i) in col:\n                        self.boost_odd_pos.append(idx)\n\n    def data_augment(self, X, y):\n        # p0~2 swap, p3~5 swap, x reverse, y reverse and AB swap\n        # p0~2 swap\n        pi_col_pos_of = self.pi_col_pos_of\n        orig_X = copy.deepcopy(X)\n        for perm in self.perm_list:\n            p0, p1, p2 = perm\n            X[:, pi_col_pos_of[p0]], X[:, pi_col_pos_of[p1]], X[:, pi_col_pos_of[p2]] = orig_X[:, pi_col_pos_of[0]], orig_X[:, pi_col_pos_of[1]], orig_X[:, pi_col_pos_of[2]]\n            yield X, y\n        # p3~5 swap\n        X = copy.deepcopy(orig_X)\n        for perm in self.perm_list:\n            p3, p4, p5 = perm\n            X[:, pi_col_pos_of[p3]], X[:, pi_col_pos_of[p4]], X[:, pi_col_pos_of[p5]] = orig_X[:, pi_col_pos_of[3]], orig_X[:, pi_col_pos_of[4]], orig_X[:, pi_col_pos_of[5]]\n            yield X, y\n        # x reverse\n        # boost0 <-> 1, 2 <-> 3, 4 <-> 5\n        X = copy.deepcopy(orig_X)\n        X[:, self.xcol_pos] *= -1\n        X[:, self.boost_even_pos], X[:, self.boost_odd_pos] = orig_X[:, self.boost_odd_pos], orig_X[:, self.boost_even_pos]\n        yield X, y\n        # y reverse and AB swajp\n        # TODO\n        # orig\n        yield orig_X, y\n\n    def orig_chain(self, dat):\n        for X, y in dat:\n            for row_x, row_y in zip(X, y):\n                if random.random() >= self.dataaug_sample:\n                    continue\n                yield row_x, row_y\n\n    def __iter__(self):\n        data = pd.read_pickle(self.pkl_path)\n        print('load {}'.format(self.pkl_path), 'dataaug: ', self.dataaug)\n        X = data[self.train_col_list].values\n        if self.target_col in data.columns:\n            y = data[self.target_col].values\n        else:\n            y = np.zeros(X.shape[0])\n        del data\n        gc.collect()\n        if self.dataaug:\n            return self.orig_chain(self.data_augment(X, y))\n        else:\n            return zip(X, y)\n\n\"\"\" HOW TO USE\nds = TrainDataset(train_col_list, target_col_list[0], pkl_path_list[0])\nfor ip, op in tqdm(ds):\n    print(ip, op)\n    break\n\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:45:22.957829Z","iopub.execute_input":"2022-10-09T12:45:22.958519Z","iopub.status.idle":"2022-10-09T12:45:23.217393Z","shell.execute_reply.started":"2022-10-09T12:45:22.958480Z","shell.execute_reply":"2022-10-09T12:45:23.216336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# chain dataset\nds_list = []\nfor pkl_path in pkl_path_list:\n    ds = TrainDataset(train_col_list, target_col_list[0], pkl_path)\n    ds_list.append(ds)\nds = ChainDataset(ds_list)","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:38:48.751763Z","iopub.execute_input":"2022-10-09T12:38:48.754624Z","iopub.status.idle":"2022-10-09T12:38:48.766254Z","shell.execute_reply.started":"2022-10-09T12:38:48.754588Z","shell.execute_reply":"2022-10-09T12:38:48.765285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# SPEEDY?\n\"\"\"\n# runtime 1min\nfor ip, op in tqdm(ds):\n    pass\n\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:38:48.770632Z","iopub.execute_input":"2022-10-09T12:38:48.773766Z","iopub.status.idle":"2022-10-09T12:38:48.783716Z","shell.execute_reply.started":"2022-10-09T12:38:48.773727Z","shell.execute_reply":"2022-10-09T12:38:48.782864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DataModule(pl.LightningDataModule):\n    # train, val, testの3つのDataLoaderを定義する\n    # trainerにこれを渡すと、train, val, testのそれぞれのステップでこれを渡してくれる\n    def __init__(self, train_col_list, target_col, train_pkl_path_list, valid_pkl_path_list, test_pkl_path, batch_size):\n        self.train_col_list = train_col_list\n        self.target_col = target_col\n        self.train_pkl_path_list = train_pkl_path_list\n        self.valid_pkl_path_list = valid_pkl_path_list\n        self.test_pkl_path = test_pkl_path\n        self.batch_size = batch_size\n        self._log_hyperparams = None  # ナニコレ・・・\n\n    def train_dataloader(self):\n        # chain dataset\n        ds_list = []\n        for pkl_path in self.train_pkl_path_list:\n            ds = TrainDataset(self.train_col_list, self.target_col, pkl_path, dataaug=CFG.dataaug)\n            ds_list.append(ds)\n        ds = ChainDataset(ds_list)\n        dl = DataLoader(ds, batch_size=self.batch_size, shuffle=False, pin_memory=True, drop_last=True, num_workers=CFG.num_workers, persistent_workers=False)\n        return dl\n\n    def val_dataloader(self):\n        # chain dataset\n        ds_list = []\n        for pkl_path in self.valid_pkl_path_list:\n            ds = TrainDataset(self.train_col_list, self.target_col, pkl_path)\n            ds_list.append(ds)\n        ds = ChainDataset(ds_list)\n        dl = DataLoader(ds, batch_size=self.batch_size, shuffle=False, pin_memory=True, drop_last=False, num_workers=CFG.num_workers, persistent_workers=False)\n        return dl\n\n    def predict_dataloader(self):\n        ds = TrainDataset(self.train_col_list, self.target_col, self.test_pkl_path)\n        # Why does setting num_worker to 2 double the amount of data?\n        dl = DataLoader(ds, batch_size=self.batch_size, shuffle=False, pin_memory=True, drop_last=False, num_workers=1, persistent_workers=False)\n        return dl\n\n    def prepare_data_per_node(self):\n        # TODO 本来要らないはずなんだけど・・・\n        pass\n\n    def teardown(self, stage=None):\n        torch.cuda.empty_cache()  # TODO: これであってるのか不明　何も出てこないんだよね\n        gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:38:48.787723Z","iopub.execute_input":"2022-10-09T12:38:48.789819Z","iopub.status.idle":"2022-10-09T12:38:48.813146Z","shell.execute_reply.started":"2022-10-09T12:38:48.789784Z","shell.execute_reply":"2022-10-09T12:38:48.812112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"# モデルの実装部分\nclass DNN(nn.Module):\n    def __init__(self, input_size):\n        super().__init__()\n        hidden_size = 120\n        output_size = 1\n        self.fc1 = nn.Linear(input_size, hidden_size*4)\n        self.bn1 = nn.BatchNorm1d(hidden_size*4)\n        self.fc2 = nn.Linear(hidden_size*4, hidden_size*4)\n        self.fc3 = nn.Linear(hidden_size*4, hidden_size*2)\n        self.fc4 = nn.Linear(hidden_size*2, hidden_size)\n        self.fc5 = nn.Linear(hidden_size, output_size)\n    \n    def forward(self, x):\n        # dropoutとbnの併用禁止\n        # bnは活性化関数の前に\n        x = F.silu(self.bn1((self.fc1(x))))\n        x = F.silu(self.fc2(x))\n        x = F.silu(self.fc3(x))\n        x = F.silu(self.fc4(x))\n        x = self.fc5(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:38:48.817614Z","iopub.execute_input":"2022-10-09T12:38:48.820453Z","iopub.status.idle":"2022-10-09T12:38:48.834734Z","shell.execute_reply.started":"2022-10-09T12:38:48.820414Z","shell.execute_reply":"2022-10-09T12:38:48.833850Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class NNModel(pl.LightningModule):\n    def __init__(self, model: nn.Module):\n        super().__init__()\n        self.model = model\n        self.criterion = nn.MSELoss()\n        self.lr = CFG.lr\n\n    def forward(self, x) -> torch.Tensor:\n        return self.model(x)\n\n    # Setup Optimizer and Scheduler\n    def configure_optimizers(self):\n        model_params = [p for n, p in self.model.named_parameters()]\n        optimizer_params = [\n            {\"params\":  model_params,\n             \"weight_decay\": CFG.weight_decay,\n             \"lr\": CFG.lr\n            },\n        ]\n\n        optimizer = optim.Adam(optimizer_params)\n\n        scheduler = lr_scheduler.ReduceLROnPlateau(optimizer,\n                                                'min',\n                                                patience=3,\n                                                factor=0.5\n                                                )\n        interval = \"epoch\"\n        monitor = \"valid_avg_loss\"\n\n        return [optimizer], [{\"scheduler\": scheduler, \"interval\": interval, \"monitor\": monitor}]\n\n    # training valid test steps\n    def training_step(self, batch_data, batch_idx):\n        # batch_data: DataModuleで定義したtrain_dataloaderの結果\n        # 戻値: lossであることが必須(裏でoptimizerに渡すため)\n        X, y = batch_data\n        op = self(X).squeeze()\n        loss = self.criterion(op, y)\n        return loss\n\n    def training_epoch_end(self, outputs):\n        # 1epoch分の処理(全バッチの処理)のreturn値をlistで受け取る\n        loss_list = [x['loss'] for x in outputs]\n        avg_loss = torch.stack(loss_list).mean()\n        self.log('train_avg_loss', avg_loss, prog_bar=True)\n        if (self.current_epoch+1) % CFG.print_epoch_freq == 0:\n            print(\"epoch:\", self.current_epoch, \"train_avg_loss:\", avg_loss.item())\n\n    def validation_step(self, batch_data, batch_idx):\n        # 戻値: 任意の辞書\n        X, y = batch_data\n        op = self(X).squeeze()\n        loss = self.criterion(op, y)\n        return {'valid_loss': loss}\n\n    def validation_epoch_end(self, outputs):\n        loss_list = [x['valid_loss'] for x in outputs]\n        avg_loss = torch.stack(loss_list).mean()\n        self.log('valid_avg_loss', avg_loss, prog_bar=True)\n        if (self.current_epoch+1) % CFG.print_epoch_freq == 0:\n            print(\"epoch:\", self.current_epoch, \"valid_avg_loss:\", avg_loss.item())\n        return avg_loss\n\n    def predict_step(self, batch_data, batch_idx):\n        # 実際に予測させるときに使う\n        X, _ = batch_data\n        outputs = self(X).squeeze()\n        # criterionがwithLogit系の場合は、sigmoidを追加する。\n        # outputs = torch.sigmoid(outputs)\n        return outputs","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:38:48.840557Z","iopub.execute_input":"2022-10-09T12:38:48.846550Z","iopub.status.idle":"2022-10-09T12:38:48.878004Z","shell.execute_reply.started":"2022-10-09T12:38:48.846514Z","shell.execute_reply":"2022-10-09T12:38:48.876815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"checkpoint_path_of = defaultdict(str)\nmodel_name_prefix = 'TPS2210'\nall_preds_of = {}\n\nif CFG.is_debug:\n    fold_list = [0]\nelse:\n    fold_list = range(5)\n\nfor target_col in target_col_list:\n    all_preds = []\n    for fold in fold_list:\n        print(\"=\"*10, \"{}/{}\".format(fold+1, len(fold_list)), \"=\"*10)\n        wandb.init(project=project_name)\n\n        # fold=0 -> valid = 0, 5\n        train_pkl_path_list = pkl_path_list[:]\n        train_pkl_path_list.pop(fold+5)\n        train_pkl_path_list.pop(fold)\n        valid_pkl_path = [pkl_path_list[fold], pkl_path_list[fold+5]]\n        test_pkl_path = 'test.pkl'\n\n        # data module\n        batch_size = CFG.batch_size  # len(X) == batch then raise errer.\n        dm = DataModule(train_col_list, target_col, train_pkl_path_list, valid_pkl_path, test_pkl_path, batch_size)\n\n        # create model\n        cur_model_name = \"model\" + model_name_prefix + \"_\" + str(fold)\n        dirpath = \"./model/\"\n        dnn = DNN(len(train_col_list))\n        model = NNModel(dnn)\n\n        # train\n        logger = WandbLogger()\n        logger.log_hyperparams(CFG.__dict__)\n        callbacks = [\n                    pl.callbacks.EarlyStopping('valid_avg_loss', patience=3),  # validation_epoch_endの戻値が10ターン改善がなかったら打ち止め\n                    pl.callbacks.ModelCheckpoint(dirpath=\"./model/\", filename=cur_model_name, save_top_k=1, monitor=\"valid_avg_loss\", save_weights_only=False),  # model保存の設定\n                    pl.callbacks.LearningRateMonitor(),  # ログに学習率を吐き出す設定\n        ]\n        trainer = pl.Trainer(accelerator=\"auto\", devices=\"auto\", max_epochs=CFG.max_epochs, logger=logger, callbacks=callbacks, enable_progress_bar=False)\n        trainer.fit(model, datamodule=dm)\n        wandb.finish()\n\n        # load_best_model\n        checkpoint_path = glob.glob(dirpath+cur_model_name+\"*.ckpt\")[0]\n        print('load model:', checkpoint_path)\n        model.load_from_checkpoint(checkpoint_path, model=dnn)\n        checkpoint_path_of[cur_model_name] = checkpoint_path\n\n        \"\"\"\n        # stack_valid\n        dm = DataModule(X_train, y_train, X_valid, y_valid, X_valid, batch_size)\n        results = trainer.predict(model=model, datamodule=dm)\n        preds = []\n        for batch in results:\n            preds.append(batch)\n        outputs = torch.cat(preds, dim=0)\n        train[\"stack\"].loc[train['fold']==fold] = outputs.tolist()\n        auc = roc_auc_score(y_valid, outputs.tolist())\n        print(\"auc :\", auc)\n        all_auc.append(auc)\n        \"\"\"\n\n        # predict\n        dm = DataModule(train_col_list, target_col, train_pkl_path_list, valid_pkl_path, test_pkl_path, batch_size)\n        results = trainer.predict(model=model, datamodule=dm)\n        preds = []\n        for batch in results:\n            preds.append(batch)\n        outputs = torch.cat(preds, dim=0)\n\n        # write result\n        all_preds.append(outputs.tolist())\n\n        torch.cuda.empty_cache()\n        gc.collect()\n    all_preds_of[target_col] = all_preds","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:45:30.296122Z","iopub.execute_input":"2022-10-09T12:45:30.296497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submit","metadata":{}},{"cell_type":"code","source":"sub = pd.read_csv('../input/tabular-playground-series-oct-2022/sample_submission.csv')\nfor target_col in target_col_list:\n    all_preds = all_preds_of[target_col]\n    all_preds = np.mean(all_preds, axis=0)\n\n    sub[target_col] = all_preds\nsub.to_csv(\"submission.csv\", index=False)\nsub","metadata":{"execution":{"iopub.status.busy":"2022-10-09T12:41:09.004031Z","iopub.status.idle":"2022-10-09T12:41:09.004773Z","shell.execute_reply.started":"2022-10-09T12:41:09.004516Z","shell.execute_reply":"2022-10-09T12:41:09.004544Z"},"trusted":true},"execution_count":null,"outputs":[]}]}