{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":59093,"databundleVersionId":7469972},{"sourceType":"datasetVersion","sourceId":7550085,"datasetId":4397253,"databundleVersionId":7644269},{"sourceType":"datasetVersion","sourceId":7664979,"datasetId":4459104,"databundleVersionId":7762101},{"sourceType":"kernelVersion","sourceId":162349255}],"dockerImageVersionId":30648,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# HMS - PyTorch Baseline Inference\n\n**This script is modified from others' work, and the original note is below**\n\nOne of my goals in this competition is to learn more PyTorch.\n\nThis is an **inference** notebook; the respetive training notebook is [HMS - PyTorch Baseline Training](https://www.kaggle.com/code/morodertobias/hms-pytorch-baseline-training/notebook), and its trained models have been registered as a versioned dataset [HMS - PyTorch Baseline Training Dataset](https://www.kaggle.com/datasets/morodertobias/hms-pytorch-baseline-training-dataset).\n\nThe model uses squashed spectrograms, as done in the reference notebooks. I try to use my way of coding, but naturally it is similar. \n\nThis version uses the current version of the notebook, version 1 of dataset, and the last successful notebook run, version 8, hence 10 models in total. Each one is an EfficientNetB0 which have been fined-tuned from noisy student weights.","metadata":{}},{"cell_type":"markdown","source":"## Table of Contents\n- [Imports](#Imports)\n- [Config](#Config)\n- [Prepare data](#Prepare-data)\n- [Prepare model](#Prepare-model)\n- [Predict](#Predict)\n- [Finalize submission](#Finalize-submission)","metadata":{}},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport pathlib\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm\nimport torch\nimport torch.nn.functional as F\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nimport timm\nfrom torch.utils.data import Dataset, DataLoader","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-02-20T16:14:38.771928Z","iopub.execute_input":"2024-02-20T16:14:38.772634Z","iopub.status.idle":"2024-02-20T16:14:47.222382Z","shell.execute_reply.started":"2024-02-20T16:14:38.772596Z","shell.execute_reply":"2024-02-20T16:14:47.221584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"class CFG:\n    base_dir = pathlib.Path(\"/kaggle/input/hms-harmful-brain-activity-classification\")\n    path_test = base_dir / \"test.csv\"\n    path_submission = base_dir / \"sample_submission.csv\"\n    spec_dir = base_dir / \"test_spectrograms\"\n    model_name = \"tf_efficientnet_b4_ns\"\n    model_weights = sorted(\n        list(pathlib.Path(\"/kaggle/input/studentresnext\").glob(\"*.pt\"))\n\n    )\n    transform = transforms.Resize((512, 512), antialias=False)\n    batch_size = 16\n    label_columns = [\n        \"seizure_vote\",\n        \"lpd_vote\",\n        \"gpd_vote\",\n        \"lrda_vote\",\n        \"grda_vote\",\n        \"other_vote\",\n    ]\n\n\nCFG.model_weights","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:47.224134Z","iopub.execute_input":"2024-02-20T16:14:47.224412Z","iopub.status.idle":"2024-02-20T16:14:47.235671Z","shell.execute_reply.started":"2024-02-20T16:14:47.224388Z","shell.execute_reply":"2024-02-20T16:14:47.234894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare data\n- Load test dataframe.\n- Prepare Dataset and DataLoader.\n- Check one example to see that everything is correct.","metadata":{}},{"cell_type":"code","source":"test = pd.read_csv(CFG.path_test)\nsubmission = pd.read_csv(CFG.path_submission)\nsubmission = pd.merge(submission, test, how=\"inner\", on=\"eeg_id\")\nsubmission[\"path\"] = submission[\"spectrogram_id\"].map(lambda x: CFG.spec_dir / f\"{x}.parquet\")\nsubmission","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:47.236698Z","iopub.execute_input":"2024-02-20T16:14:47.236981Z","iopub.status.idle":"2024-02-20T16:14:47.276756Z","shell.execute_reply.started":"2024-02-20T16:14:47.236953Z","shell.execute_reply":"2024-02-20T16:14:47.275895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess(x):\n    x = np.clip(x, np.exp(-6), np.exp(10))\n    x = np.log(x)\n    m, s = x.mean(), x.std()\n    x = (x - m) / (s + 1e-6)\n    return x\n\n\nclass SpecDataset(Dataset):\n    \n    def __init__(self, df, transform=CFG.transform):\n        self.df = df\n        self.transform = transform\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        # input\n        x = pd.read_parquet(row.path)\n        x = x.fillna(-1).values[:, 1:].T\n        x = preprocess(x)\n        x = torch.Tensor(x[None, :])\n        if self.transform:\n            x = self.transform(x)\n        # output\n        y = np.array(row.loc[CFG.label_columns].values, 'float32')\n        y = torch.Tensor(y)\n        return x, y","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:47.277879Z","iopub.execute_input":"2024-02-20T16:14:47.278154Z","iopub.status.idle":"2024-02-20T16:14:47.287745Z","shell.execute_reply.started":"2024-02-20T16:14:47.278131Z","shell.execute_reply":"2024-02-20T16:14:47.286907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_ds = SpecDataset(df=submission)\ndata_loader = DataLoader(dataset=data_ds, num_workers=os.cpu_count())\ndata_loader","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:47.290688Z","iopub.execute_input":"2024-02-20T16:14:47.291299Z","iopub.status.idle":"2024-02-20T16:14:47.299737Z","shell.execute_reply.started":"2024-02-20T16:14:47.291274Z","shell.execute_reply":"2024-02-20T16:14:47.298903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x, y = next(iter(data_loader))\nx.shape, x","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:47.300716Z","iopub.execute_input":"2024-02-20T16:14:47.301006Z","iopub.status.idle":"2024-02-20T16:14:47.758964Z","shell.execute_reply.started":"2024-02-20T16:14:47.300982Z","shell.execute_reply":"2024-02-20T16:14:47.757873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(x[0, 0])\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:47.760276Z","iopub.execute_input":"2024-02-20T16:14:47.760589Z","iopub.status.idle":"2024-02-20T16:14:48.051074Z","shell.execute_reply.started":"2024-02-20T16:14:47.760561Z","shell.execute_reply":"2024-02-20T16:14:48.050113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare model","metadata":{}},{"cell_type":"code","source":"DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"DEVICE: {DEVICE}\")","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:48.052288Z","iopub.execute_input":"2024-02-20T16:14:48.052572Z","iopub.status.idle":"2024-02-20T16:14:48.079400Z","shell.execute_reply.started":"2024-02-20T16:14:48.052548Z","shell.execute_reply":"2024-02-20T16:14:48.078349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def model_B(num_classes):\n    class Block(nn.Module):\n        def __init__(self,in_channels, out_channels, stride=1, is_shortcut=False):\n            super(Block,self).__init__()\n            self.relu = nn.ReLU(inplace=True)\n            self.is_shortcut = is_shortcut\n            self.conv1 = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels // 2, kernel_size=1,stride=stride,bias=False),\n                nn.BatchNorm2d(out_channels // 2),\n                nn.ReLU()\n            )\n            self.conv2 = nn.Sequential(\n                nn.Conv2d(out_channels // 2, out_channels // 2, kernel_size=3, stride=1, padding=1, groups=32,\n                                       bias=False),\n                nn.BatchNorm2d(out_channels // 2),\n                nn.ReLU()\n            )\n            self.conv3 = nn.Sequential(\n                nn.Conv2d(out_channels // 2, out_channels, kernel_size=1,stride=1,bias=False),\n                nn.BatchNorm2d(out_channels),\n            )\n            if is_shortcut:\n                self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels,out_channels,kernel_size=1,stride=stride,bias=1),\n                nn.BatchNorm2d(out_channels)\n            )\n        def forward(self, x):\n            x_shortcut = x\n            x = self.conv1(x)\n            x = self.conv2(x)\n            x = self.conv3(x)\n            if self.is_shortcut:\n                x_shortcut = self.shortcut(x_shortcut)\n            x = x + x_shortcut\n            x = self.relu(x)\n            return x\n    \n    class Resnext(nn.Module):\n        def __init__(self,num_classes,layer=[3,4,6,3,3]):\n            super(Resnext,self).__init__()\n            self.conv1 = nn.Sequential(\n                nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False),\n                nn.BatchNorm2d(64),\n                nn.ReLU(),\n                nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n            )\n            self.conv2 = self._make_layer(64,256,1,num=layer[0])\n            self.conv3 = self._make_layer(256,512,2,num=layer[1])\n            self.conv4 = self._make_layer(512,1024,2,num=layer[2])\n            self.conv5 = self._make_layer(1024,2048,3,num=layer[3])\n            self.conv6 = self._make_layer(2048,4096,2,num=layer[4])\n            self.global_average_pool = nn.AvgPool2d(kernel_size=6, stride=1)\n            self.fc = nn.Linear(4096,num_classes)\n        def forward(self, x):\n            x = self.conv1(x)\n            x = self.conv2(x)\n            x = self.conv3(x)\n            x = self.conv4(x)\n            x = self.conv5(x)\n            x = self.conv6(x)\n            x = self.global_average_pool(x)\n            x = torch.flatten(x,1)\n            x = self.fc(x)\n            return x\n        def _make_layer(self,in_channels,out_channels,stride,num):\n            layers = []\n            block_1=Block(in_channels, out_channels,stride=stride,is_shortcut=True)\n            layers.append(block_1)\n            for i in range(1, num):\n                layers.append(Block(out_channels,out_channels,stride=1,is_shortcut=False))\n            return nn.Sequential(*layers)\n\n    model = Resnext(num_classes)\n    return model","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:48.080875Z","iopub.execute_input":"2024-02-20T16:14:48.081243Z","iopub.status.idle":"2024-02-20T16:14:48.102040Z","shell.execute_reply.started":"2024-02-20T16:14:48.081210Z","shell.execute_reply":"2024-02-20T16:14:48.101083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = timm.create_model(model_name=CFG.model_name, pretrained=False, num_classes=6, in_chans=1)\n#model = model_B(num_classes=6)\nmodel.to(DEVICE)\nnum_parameter = sum(x.numel() for x in model.parameters())\nprint(f\"Model has {num_parameter} parameters.\")","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:48.103175Z","iopub.execute_input":"2024-02-20T16:14:48.103441Z","iopub.status.idle":"2024-02-20T16:14:48.675382Z","shell.execute_reply.started":"2024-02-20T16:14:48.103419Z","shell.execute_reply":"2024-02-20T16:14:48.674458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict\n- Load weights and compute individual predictions.\n- Note, the output of the model are logits.\n- Final predicition is the ensemble of all invidiual predictions.","metadata":{}},{"cell_type":"code","source":"prediction = pd.DataFrame(0.0, columns=CFG.label_columns, index=submission.index)\nfor i, path_weight in enumerate(CFG.model_weights):\n    print(f\"Model {i}: {path_weight}\")\n    model.load_state_dict(torch.load(path_weight))\n    model.eval()\n    with torch.no_grad():\n        res = []\n        for x, y in data_loader:\n            x = x.to(DEVICE)\n            pred = model(x)\n            pred = F.softmax(pred, dim=1)\n            pred = pred.detach().cpu().numpy()\n            res.append(pred)\n        res = np.concatenate(res)\n        res = pd.DataFrame(res, columns=CFG.label_columns, index=submission.index)\n        display(res)\n        prediction = prediction + res\n        print(\"\\n\")\nprediction = prediction / len(CFG.model_weights)","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:48.676766Z","iopub.execute_input":"2024-02-20T16:14:48.677522Z","iopub.status.idle":"2024-02-20T16:14:50.895434Z","shell.execute_reply.started":"2024-02-20T16:14:48.677488Z","shell.execute_reply":"2024-02-20T16:14:50.894327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"prediction","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:50.896958Z","iopub.execute_input":"2024-02-20T16:14:50.897260Z","iopub.status.idle":"2024-02-20T16:14:50.908286Z","shell.execute_reply.started":"2024-02-20T16:14:50.897233Z","shell.execute_reply":"2024-02-20T16:14:50.907291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Finalize submission","metadata":{"execution":{"iopub.status.busy":"2024-02-04T07:41:51.765545Z","iopub.execute_input":"2024-02-04T07:41:51.765946Z","iopub.status.idle":"2024-02-04T07:41:51.774172Z","shell.execute_reply.started":"2024-02-04T07:41:51.765901Z","shell.execute_reply":"2024-02-04T07:41:51.773122Z"}}},{"cell_type":"code","source":"submission[CFG.label_columns] = prediction\nsubmission = submission[[\"eeg_id\"] + CFG.label_columns]\nsubmission","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:50.909384Z","iopub.execute_input":"2024-02-20T16:14:50.909695Z","iopub.status.idle":"2024-02-20T16:14:50.926946Z","shell.execute_reply.started":"2024-02-20T16:14:50.909653Z","shell.execute_reply":"2024-02-20T16:14:50.926030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv(\"submission.csv\", index=None)","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:50.930044Z","iopub.execute_input":"2024-02-20T16:14:50.930329Z","iopub.status.idle":"2024-02-20T16:14:50.938614Z","shell.execute_reply.started":"2024-02-20T16:14:50.930306Z","shell.execute_reply":"2024-02-20T16:14:50.937832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head submission.csv","metadata":{"execution":{"iopub.status.busy":"2024-02-20T16:14:50.939758Z","iopub.execute_input":"2024-02-20T16:14:50.940146Z","iopub.status.idle":"2024-02-20T16:14:51.895514Z","shell.execute_reply.started":"2024-02-20T16:14:50.940121Z","shell.execute_reply":"2024-02-20T16:14:51.894301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}