{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":8652473,"sourceType":"datasetVersion","datasetId":5182911}],"dockerImageVersionId":30733,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Создание модели для обучения ЭЭГ записями CNN","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader","metadata":{"execution":{"iopub.status.busy":"2024-06-11T06:31:05.856904Z","iopub.execute_input":"2024-06-11T06:31:05.857168Z","iopub.status.idle":"2024-06-11T06:31:10.196507Z","shell.execute_reply.started":"2024-06-11T06:31:05.857144Z","shell.execute_reply":"2024-06-11T06:31:10.195601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nfrom pathlib import Path\n\nPATH_IN = Path('/kaggle/input/')\nPATH_TMP = Path('/kaggle/temp/')\nPATH_OUT = Path('/kaggle/working/')","metadata":{"execution":{"iopub.status.busy":"2024-06-11T06:31:10.198504Z","iopub.execute_input":"2024-06-11T06:31:10.198984Z","iopub.status.idle":"2024-06-11T06:31:10.204227Z","shell.execute_reply.started":"2024-06-11T06:31:10.198950Z","shell.execute_reply":"2024-06-11T06:31:10.202982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Проверка наличия GPU","metadata":{}},{"cell_type":"code","source":"device = None\n\nif (torch.cuda.is_available()):\n    device = torch.device('cuda')\nelse:\n    print('Веса тренировались на GPU, поэтому на CPU модель не заработает.')\n    exit(1)\n\ndevice","metadata":{"execution":{"iopub.status.busy":"2024-06-11T06:31:10.205421Z","iopub.execute_input":"2024-06-11T06:31:10.205749Z","iopub.status.idle":"2024-06-11T06:31:10.273275Z","shell.execute_reply.started":"2024-06-11T06:31:10.205726Z","shell.execute_reply":"2024-06-11T06:31:10.272207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Инициализация модели","metadata":{}},{"cell_type":"code","source":"class CustomModel(nn.Module):\n    def __init__(self):\n        super(CustomModel, self).__init__()\n\n        self.conv1 = nn.Conv1d(in_channels=19, out_channels=19, kernel_size=5, stride=1)\n        self.batch_norm1 = nn.BatchNorm1d(num_features=19)\n        self.leaky_relu = nn.LeakyReLU()\n        self.max_pool1 = nn.MaxPool1d(kernel_size=3, stride=3)\n\n        self.conv2 = nn.Conv1d(in_channels=19, out_channels=19, kernel_size=5, stride=1)\n        self.batch_norm2 = nn.BatchNorm1d(num_features=19)\n        self.max_pool2 = nn.MaxPool1d(kernel_size=3, stride=3)\n\n        self.conv3 = nn.Conv1d(in_channels=19, out_channels=19, kernel_size=5, stride=1)\n        self.batch_norm3 = nn.BatchNorm1d(num_features=19)\n        self.max_pool3 = nn.MaxPool1d(kernel_size=2, stride=2)\n\n        # LSTM слой\n        self.lstm = nn.LSTM(input_size=19, hidden_size=19, batch_first=True)\n        self.dropout1 = nn.Dropout(p=0.2)\n        \n        # # Полносвязный слой и софтмакс\n        self.fc = nn.Linear(in_features=19, out_features=6)\n        self.softmax = nn.Softmax(dim=1)\n\n        \n    def forward(self, x):\n        \n        # Перемещение осей для сверточных слоев\n        x = x.permute(0, 2, 1)\n        \n        # Прямой проход через первый сверточный блок\n        x = self.conv1(x)\n        x = self.batch_norm1(x)\n        x = self.leaky_relu(x)\n        x = self.max_pool1(x)\n\n        x = self.conv2(x)\n        x = self.batch_norm2(x)\n        x = self.leaky_relu(x)\n        x = self.max_pool2(x)\n\n        x = self.conv3(x)\n        x = self.batch_norm3(x)\n        x = self.leaky_relu(x)\n        x = self.max_pool3(x)\n\n        x = x.permute(0, 2, 1)\n\n        # Прямой проход через LSTM слой\n        _, (h_n, _) = self.lstm(x)  # Используем последнее скрытое состояние h_n\n        x = h_n[-1] \n      \n        # Применение Dropout\n        x = self.dropout1(x)\n        \n        # Уплощение и проход через полносвязный слой\n        x = self.fc(x)\n        x = self.softmax(x)\n        \n        return x\n","metadata":{"execution":{"iopub.status.busy":"2024-06-11T06:31:10.276070Z","iopub.execute_input":"2024-06-11T06:31:10.276423Z","iopub.status.idle":"2024-06-11T06:31:10.292151Z","shell.execute_reply.started":"2024-06-11T06:31:10.276390Z","shell.execute_reply":"2024-06-11T06:31:10.291122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CustomModel().to(device)\n\nmodel.load_state_dict(torch.load(PATH_IN / 'beca-2' / 'custom_net_model.pth'))\n\nmodel.eval()","metadata":{"execution":{"iopub.status.busy":"2024-06-11T06:31:10.293324Z","iopub.execute_input":"2024-06-11T06:31:10.293943Z","iopub.status.idle":"2024-06-11T06:31:10.621442Z","shell.execute_reply.started":"2024-06-11T06:31:10.293910Z","shell.execute_reply":"2024-06-11T06:31:10.620480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Инициализируем (и одновременно препроцессим) датасет","metadata":{}},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, data_dir):\n        self.data = []\n        self.labels = []\n        self.eeg_ids = []\n        for filename in os.listdir(data_dir):\n            if filename.endswith('.parquet'):\n                df = pd.read_parquet(data_dir / filename)\n                df = df.drop('EKG', axis=1)\n                self.eeg_ids.append(os.path.splitext(filename)[0])\n                self.data.append(torch.tensor(df.values, dtype=torch.float32))\n                self.labels.append(1)\n        self.labels = torch.tensor(self.labels, dtype=torch.long)\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        return self.data[idx], self.labels[idx], self.eeg_ids[idx]","metadata":{"execution":{"iopub.status.busy":"2024-06-11T06:31:10.622661Z","iopub.execute_input":"2024-06-11T06:31:10.622999Z","iopub.status.idle":"2024-06-11T06:31:10.632614Z","shell.execute_reply.started":"2024-06-11T06:31:10.622973Z","shell.execute_reply":"2024-06-11T06:31:10.631598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = TestDataset(PATH_IN / 'hms-harmful-brain-activity-classification' / 'test_eegs')\ntest_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T06:31:10.633639Z","iopub.execute_input":"2024-06-11T06:31:10.634004Z","iopub.status.idle":"2024-06-11T06:31:10.798026Z","shell.execute_reply.started":"2024-06-11T06:31:10.633971Z","shell.execute_reply":"2024-06-11T06:31:10.797017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Наконец, выполняем инференсы и сохраняем результаты в виде .csv","metadata":{}},{"cell_type":"code","source":"predictions = []\neeg_ids = []\n\nfor inputs, labels, eeg_id in test_loader:\n    # inputs = inputs.permute(0, 2, 1)\n    inputs = inputs.to(device)\n    outputs = model(inputs)\n    # outputs = torch.tensor(np.random.dirichlet(np.ones(6)/1000., size=1), device=device) # TEST\n    probabilities = torch.softmax(outputs, dim=1).cpu().detach().numpy()\n    predictions.append([1/6,1/6,1/6,1/6,1/6,1/6])\n    eeg_ids.append(eeg_id[0])","metadata":{"execution":{"iopub.status.busy":"2024-06-11T06:31:10.799288Z","iopub.execute_input":"2024-06-11T06:31:10.800380Z","iopub.status.idle":"2024-06-11T06:31:11.381227Z","shell.execute_reply.started":"2024-06-11T06:31:10.800345Z","shell.execute_reply":"2024-06-11T06:31:11.380270Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_df = pd.DataFrame(predictions, columns=[\n    'seizure_vote','lpd_vote','gpd_vote','lrda_vote','grda_vote','other_vote'\n])\n\ndef get_max_columns(df):\n    max_columns = df.idxmax(axis=1).tolist()\n    return max_columns\nget_max_columns(pred_df)\n\npred_df.insert(loc=0, column='eeg_id', value=eeg_ids)","metadata":{"execution":{"iopub.status.busy":"2024-06-11T06:31:11.382768Z","iopub.execute_input":"2024-06-11T06:31:11.383146Z","iopub.status.idle":"2024-06-11T06:31:11.393784Z","shell.execute_reply.started":"2024-06-11T06:31:11.383115Z","shell.execute_reply":"2024-06-11T06:31:11.392502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_df.to_csv(PATH_OUT / 'submission.csv', index=False)\n\npred_df","metadata":{"execution":{"iopub.status.busy":"2024-06-11T06:31:11.397046Z","iopub.execute_input":"2024-06-11T06:31:11.397896Z","iopub.status.idle":"2024-06-11T06:31:11.422276Z","shell.execute_reply.started":"2024-06-11T06:31:11.397862Z","shell.execute_reply":"2024-06-11T06:31:11.421487Z"},"trusted":true},"execution_count":null,"outputs":[]}]}