{"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":"tpu1vmV38","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.transforms as transforms\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nimport numpy as np\nimport pandas as pd\nimport joblib\nimport math\nimport cv2\n\nimport matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:42:56.837592Z","iopub.execute_input":"2024-02-25T10:42:56.837940Z","iopub.status.idle":"2024-02-25T10:43:23.683340Z","shell.execute_reply.started":"2024-02-25T10:42:56.837910Z","shell.execute_reply":"2024-02-25T10:43:23.682566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nBASE_PATH = \"/kaggle/input/hms-harmful-brain-activity-classification\"\n\nSPEC_DIR = \"/tmp/dataset/hms-hbac\"\nos.makedirs(os.path.join(SPEC_DIR, 'train_spectrograms'), exist_ok=True)\nos.makedirs(os.path.join(SPEC_DIR, 'test_spectrograms'), exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:43:26.524234Z","iopub.execute_input":"2024-02-25T10:43:26.524676Z","iopub.status.idle":"2024-02-25T10:43:26.531733Z","shell.execute_reply.started":"2024-02-25T10:43:26.524646Z","shell.execute_reply":"2024-02-25T10:43:26.531118Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    verbose = 1  # Verbosity\n    seed = 42  # Random seed\n    preset = \"efficientnetv2_b2_imagenet\"  # Name of pretrained classifier\n    image_size = [400, 300]  # Input image size\n    epochs = 13 # Training epochs\n    batch_size = 64  # Batch size\n    lr_mode = \"cos\" # LR scheduler mode from one of \"cos\", \"step\", \"exp\"\n    drop_remainder = True  # Drop incomplete batches\n    num_classes = 6 # Number of classes in the dataset\n    fold = 0 # Which fold to set as validation data\n    class_names = ['Seizure', 'LPD', 'GPD', 'LRDA','GRDA', 'Other']\n    label2name = dict(enumerate(class_names))\n    name2label = {v:k for k, v in label2name.items()}","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:43:28.716681Z","iopub.execute_input":"2024-02-25T10:43:28.717627Z","iopub.status.idle":"2024-02-25T10:43:28.723016Z","shell.execute_reply.started":"2024-02-25T10:43:28.717567Z","shell.execute_reply":"2024-02-25T10:43:28.722309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PyTorchModel(nn.Module):\n    def __init__(self, num_classes):\n        super(PyTorchModel, self).__init__()\n        self.conv1 = nn.Conv2d(in_channels=400, out_channels=64, kernel_size=(3, 3), padding=(1, 1))\n        self.conv2 = nn.Conv2d(in_channels=64, out_channels=128, kernel_size=(3, 3), padding=(1, 1))\n        self.conv3 = nn.Conv2d(in_channels=128, out_channels=256, kernel_size=(3, 3), padding=(1, 1))\n        self.pool = nn.MaxPool2d(kernel_size=(2, 2), stride=(2, 2))\n        self.flatten = nn.Flatten()\n        self.fc1 = nn.Linear(256 * 25 * 18, 512)\n        self.fc2 = nn.Linear(512, num_classes)\n        self.relu = nn.ReLU()\n        self.dropout = nn.Dropout(0.5)\n        \n    def forward(self, x):\n        x = self.relu(self.conv1(x))\n        x = self.pool(x)\n        x = self.relu(self.conv2(x))\n        x = self.pool(x)\n        x = self.relu(self.conv3(x))\n        x = self.pool(x)\n        x = self.flatten(x)\n        x = self.relu(self.fc1(x))\n        x = self.dropout(x)\n        x = self.fc2(x)\n        return x\n\n# Initialize the model\nnum_classes = CFG.num_classes  # Define the number of output classes\nmodel = PyTorchModel(num_classes=num_classes)\n\n# Move the model to the device\nmodel.to(device)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:43:30.529120Z","iopub.execute_input":"2024-02-25T10:43:30.529825Z","iopub.status.idle":"2024-02-25T10:43:31.075193Z","shell.execute_reply.started":"2024-02-25T10:43:30.529787Z","shell.execute_reply":"2024-02-25T10:43:31.074485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define PyTorch image augmentations\ntrain_transform = transforms.Compose([\n    transforms.ConvertImageDtype(dtype=torch.float),  # Convert image dtype to float\n    transforms.ToPILImage(),  # Convert tensor to PIL Image\n    transforms.Lambda(lambda img: img.convert('RGB')),\n    transforms.RandomHorizontalFlip(),  # Random horizontal flip\n    transforms.RandomVerticalFlip(),  # Random vertical flip\n    transforms.RandomRotation(degrees=15),  # Random rotation\n    transforms.Resize(size=(400, 300)),  # Resize the image\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),  # Color jitter\n    transforms.ToTensor(),  # Convert PIL Image to tensor\n    transforms.Normalize(mean=[0.5], std=[0.5])  # Normalize the tensor\n])\n\nvalid_transform = transforms.Compose([\n    transforms.ConvertImageDtype(dtype=torch.float),  # Convert image dtype to float\n    transforms.ToPILImage(),  # Convert tensor to PIL Image\n    transforms.Lambda(lambda img: img.convert('RGB')),\n    transforms.Resize(size=(400, 300)),  # Resize the image\n    transforms.ToTensor(),  # Convert PIL Image to tensor\n    transforms.Normalize(mean=[0.5], std=[0.5])  # Normalize the tensor\n])\n\n# If you have test-time augmentations, define them here as well\ntest_transform = transforms.Compose([\n    transforms.ConvertImageDtype(dtype=torch.float),  # Convert image dtype to float\n    transforms.ToPILImage(),  # Convert tensor to PIL Image\n    transforms.Lambda(lambda img: img.convert('RGB')),\n    transforms.Resize(size=(400, 300)),  # Resize the image\n    transforms.ToTensor(),  # Convert PIL Image to tensor\n    transforms.Normalize(mean=[0.5], std=[0.5])  # Normalize the tensor\n])","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:43:33.880581Z","iopub.execute_input":"2024-02-25T10:43:33.880943Z","iopub.status.idle":"2024-02-25T10:43:33.889433Z","shell.execute_reply.started":"2024-02-25T10:43:33.880911Z","shell.execute_reply":"2024-02-25T10:43:33.888645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_spec(spec_id, split=\"train\"):\n    spec_path = os.path.join(BASE_PATH, f\"{split}_spectrograms\", f\"{spec_id}.parquet\")\n    spec = pd.read_parquet(spec_path).fillna(0).values[:, 1:].T\n    spec = spec.astype(np.float32)\n    save_dir = os.path.join(SPEC_DIR, f\"{split}_spectrograms\")\n    np.save(os.path.join(save_dir, f\"{spec_id}.npy\"), spec)\n    return os.path.join(save_dir, f\"{spec_id}.npy\")","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:43:37.529812Z","iopub.execute_input":"2024-02-25T10:43:37.530161Z","iopub.status.idle":"2024-02-25T10:43:37.535677Z","shell.execute_reply.started":"2024-02-25T10:43:37.530133Z","shell.execute_reply":"2024-02-25T10:43:37.534754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Read train and test data\ntrain_df = pd.read_csv(os.path.join(BASE_PATH, \"train.csv\"))\ntest_df = pd.read_csv(os.path.join(BASE_PATH, \"test.csv\"))\n\n# Get unique spectrogram IDs for train and test\ntrain_spec_ids = train_df[\"spectrogram_id\"].unique()\ntest_spec_ids = test_df[\"spectrogram_id\"].unique()\n\n# Split train data into train and validation sets\ntrain_ratio = 0.8\nnum_train_samples = int(len(train_spec_ids) * train_ratio)\ntrain_spec_ids_split = train_spec_ids[:num_train_samples]\nvalid_spec_ids_split = train_spec_ids[num_train_samples:]\n\n","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:43:39.305144Z","iopub.execute_input":"2024-02-25T10:43:39.305470Z","iopub.status.idle":"2024-02-25T10:43:39.519714Z","shell.execute_reply.started":"2024-02-25T10:43:39.305443Z","shell.execute_reply":"2024-02-25T10:43:39.518830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install pyarrow","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:44:14.730615Z","iopub.execute_input":"2024-02-25T10:44:14.731043Z","iopub.status.idle":"2024-02-25T10:44:21.730425Z","shell.execute_reply.started":"2024-02-25T10:44:14.731007Z","shell.execute_reply":"2024-02-25T10:44:21.729379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.auto import tqdm\n# Parallel processing for training data\ntrain_paths = joblib.Parallel(n_jobs=-1, backend=\"loky\")(\n    joblib.delayed(process_spec)(spec_id, \"train\") for spec_id in tqdm(train_spec_ids_split)\n)","metadata":{"_kg_hide-output":false,"execution":{"iopub.status.busy":"2024-02-25T10:44:21.732425Z","iopub.execute_input":"2024-02-25T10:44:21.732754Z","iopub.status.idle":"2024-02-25T10:44:50.414045Z","shell.execute_reply.started":"2024-02-25T10:44:21.732724Z","shell.execute_reply":"2024-02-25T10:44:50.412838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Parallel processing for validation data\nvalid_paths = joblib.Parallel(n_jobs=-1, backend=\"loky\")(\n    joblib.delayed(process_spec)(spec_id, \"train\") for spec_id in tqdm(valid_spec_ids_split)\n)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:44:50.416243Z","iopub.execute_input":"2024-02-25T10:44:50.416700Z","iopub.status.idle":"2024-02-25T10:44:54.687805Z","shell.execute_reply.started":"2024-02-25T10:44:50.416648Z","shell.execute_reply":"2024-02-25T10:44:54.687031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Parallel processing for test data\ntest_paths = joblib.Parallel(n_jobs=-1, backend=\"loky\")(\n    joblib.delayed(process_spec)(spec_id, \"test\") for spec_id in tqdm(test_spec_ids)\n)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:44:54.689447Z","iopub.execute_input":"2024-02-25T10:44:54.689954Z","iopub.status.idle":"2024-02-25T10:44:54.779781Z","shell.execute_reply.started":"2024-02-25T10:44:54.689922Z","shell.execute_reply":"2024-02-25T10:44:54.778987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def paths(pths):\n    data_list = []\n    max_length = 0  # Variable to store the maximum length along axis 1\n\n    for pth in pths:\n        data = np.load(pth)\n        data_list.append(data)\n        max_length = max(max_length, data.shape[1])  # Update the maximum length\n\n    # Pad arrays to ensure they have the same shape along axis 1\n    padded_data_list = [np.pad(data, ((0, 0), (0, max_length - data.shape[1])), mode='constant') for data in data_list]\n\n    # Concatenate the data into a single numpy array\n    concatenated_data = np.concatenate(padded_data_list, axis=0)\n\n    # Convert numpy array to pandas DataFrame\n    df = pd.DataFrame(concatenated_data)\n\n    return df","metadata":{"execution":{"iopub.status.busy":"2024-02-24T18:01:31.716319Z","iopub.execute_input":"2024-02-24T18:01:31.716695Z","iopub.status.idle":"2024-02-24T18:01:33.454826Z","shell.execute_reply.started":"2024-02-24T18:01:31.716665Z","shell.execute_reply":"2024-02-24T18:01:33.454021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data= paths(train_paths)\nvalid_data= paths(valid_paths)\ntest_data= paths(test_paths)","metadata":{"execution":{"iopub.status.busy":"2024-02-24T18:01:37.547285Z","iopub.execute_input":"2024-02-24T18:01:37.547658Z","iopub.status.idle":"2024-02-24T18:05:02.098389Z","shell.execute_reply.started":"2024-02-24T18:01:37.547630Z","shell.execute_reply":"2024-02-24T18:05:02.097299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EEGDataset(Dataset):\n    def __init__(self, spec_paths, transform=None):\n        self.spec_paths = spec_paths\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.spec_paths)\n\n    def __getitem__(self, idx):\n        spec_path = self.spec_paths[idx]\n        spec = np.load(spec_path)\n\n        if self.transform:\n            spec = torch.from_numpy(spec)  # Convert to tensor\n            spec = self.transform(spec)\n\n        return spec\n\n","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:00:53.150298Z","iopub.execute_input":"2024-02-25T10:00:53.150966Z","iopub.status.idle":"2024-02-25T10:00:53.155509Z","shell.execute_reply.started":"2024-02-25T10:00:53.150934Z","shell.execute_reply":"2024-02-25T10:00:53.154872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create datasets without resizing\ntrain_dataset = EEGDataset(train_paths, transform=train_transform)\nvalid_dataset = EEGDataset(valid_paths, transform=valid_transform)\ntest_dataset = EEGDataset(test_paths, transform=test_transform)\n\n# print(\"Train set: \", list(train_dataset))\n# print(\"\\nvalid set: \", valid_dataset)\n# print(\"\\nTest set: \", test_dataset)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:00:58.483516Z","iopub.execute_input":"2024-02-25T10:00:58.483899Z","iopub.status.idle":"2024-02-25T10:00:58.487815Z","shell.execute_reply.started":"2024-02-25T10:00:58.483861Z","shell.execute_reply":"2024-02-25T10:00:58.487110Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define data loaders\nbatch_size = CFG.batch_size\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\nvalid_loader = DataLoader(valid_dataset, batch_size=batch_size, shuffle=False)\ntest_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)\n\nprint(\"Train set: \", train_loader)\nprint(\"\\nvalid set: \", valid_loader)\nprint(\"\\nTest set: \", test_loader)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:00:58.937507Z","iopub.execute_input":"2024-02-25T10:00:58.937827Z","iopub.status.idle":"2024-02-25T10:00:58.943161Z","shell.execute_reply.started":"2024-02-25T10:00:58.937800Z","shell.execute_reply":"2024-02-25T10:00:58.942486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Train set:\")\nfor batch_idx, (data, target) in enumerate(train_loader):\n    print(f\"Batch {batch_idx}:\")\n    print(\"Data shape:\", data.shape)\n    print(\"Target shape:\", target.shape)\n    # Optionally, print or visualize the data and target here\n\n# # Iterate over valid loader\n# print(\"\\nValid set:\")\n# for batch_idx, (data, target) in enumerate(valid_loader):\n#     print(f\"Batch {batch_idx}:\")\n#     print(\"Data shape:\", data.shape)\n#     print(\"Target shape:\", target.shape)\n#     # Optionally, print or visualize the data and target here\n\n# # Iterate over test loader\n# print(\"\\nTest set:\")\n# for batch_idx, (data, target) in enumerate(test_loader):\n#     print(f\"Batch {batch_idx}:\")\n#     print(\"Data shape:\", data.shape)\n#     print(\"Target shape:\", target.shape)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define your loss function\nloss_function = nn.KLDivLoss()\n\n# Define your optimizer\noptimizer = optim.Adam(model.parameters(), lr=1e-4)","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:01:04.502598Z","iopub.execute_input":"2024-02-25T10:01:04.503220Z","iopub.status.idle":"2024-02-25T10:01:04.507011Z","shell.execute_reply.started":"2024-02-25T10:01:04.503187Z","shell.execute_reply":"2024-02-25T10:01:04.506394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define your learning rate scheduler\ndef get_lr(epoch, batch_size=8, mode='cos', epochs=10):\n    lr_start, lr_max, lr_min = 5e-5, 6e-6 * batch_size, 1e-5\n    lr_ramp_ep, lr_sus_ep, lr_decay = 3, 0, 0.75\n\n    if epoch < lr_ramp_ep:\n        lr = (lr_max - lr_start) / lr_ramp_ep * epoch + lr_start\n    elif epoch < lr_ramp_ep + lr_sus_ep:\n        lr = lr_max\n    elif mode == 'exp':\n        lr = (lr_max - lr_min) * lr_decay**(epoch - lr_ramp_ep - lr_sus_ep) + lr_min\n    elif mode == 'step':\n        lr = lr_max * lr_decay**((epoch - lr_ramp_ep - lr_sus_ep) // 2)\n    elif mode == 'cos':\n        decay_total_epochs, decay_epoch_index = epochs - lr_ramp_ep - lr_sus_ep + 3, epoch - lr_ramp_ep - lr_sus_ep\n        phase = math.pi * decay_epoch_index / decay_total_epochs\n        lr = (lr_max - lr_min) * 0.5 * (1 + math.cos(phase)) + lr_min\n    return lr\n\nscheduler = CosineAnnealingLR(optimizer, T_max=CFG.epochs, eta_min=1e-5)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:01:04.993353Z","iopub.execute_input":"2024-02-25T10:01:04.993972Z","iopub.status.idle":"2024-02-25T10:01:04.999905Z","shell.execute_reply.started":"2024-02-25T10:01:04.993943Z","shell.execute_reply":"2024-02-25T10:01:04.999182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for batch_idx, batch in enumerate(train_loader):\n    data, target = batch  # Unpack the batch tuple\n    print(f\"Batch {batch_idx}:\")\n    print(\"Data shape:\", data.shape)\n    print(\"Target shape:\", target.shape)\n    ","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:06:45.417064Z","iopub.execute_input":"2024-02-25T10:06:45.417864Z","iopub.status.idle":"2024-02-25T10:06:46.025348Z","shell.execute_reply.started":"2024-02-25T10:06:45.417830Z","shell.execute_reply":"2024-02-25T10:06:46.024463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training loop\nfor epoch in range(CFG.epochs):\n    model.train()\n    train_loss = 0.0\n    for batch_idx, (data, target) in enumerate(train_loader):\n        data, target = data.to(device), target.to(device)\n        optimizer.zero_grad()\n        output = model(data)\n        loss = loss_function(output, target)\n        loss.backward()\n        optimizer.step()\n        train_loss += loss.item() * data.size(0)\n    scheduler.step()\n    \n    model.eval()\n    valid_loss = 0.0\n    for data, target in valid_loader:\n        data, target = data.to(device), target.to(device)\n        output = model(data)\n        loss = loss_function(output, target)\n        valid_loss += loss.item() * data.size(0)\n    \n    train_loss /= len(train_loader.dataset)\n    valid_loss /= len(valid_loader.dataset)\n    \n    print(f'Epoch: {epoch+1}/{CFG.epochs}, Train Loss: {train_loss:.6f}, Valid Loss: {valid_loss:.6f}')\n","metadata":{"execution":{"iopub.status.busy":"2024-02-25T10:12:06.078646Z","iopub.execute_input":"2024-02-25T10:12:06.079287Z","iopub.status.idle":"2024-02-25T10:12:06.700230Z","shell.execute_reply.started":"2024-02-25T10:12:06.079254Z","shell.execute_reply":"2024-02-25T10:12:06.699248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Testing loop\nmodel.eval()\ntest_predictions = []\nfor data in test_loader:\n    data = data.to(device)\n    output = model(data)\n    test_predictions.extend(output.cpu().detach().numpy())\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Convert predictions to the required format\ndef post_process_predictions(predictions):\n    return predictions.tolist()\n\n# Post-process test predictions\ntest_predictions_processed = post_process_predictions(preds)\n\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create DataFrame with EEG IDs and predictions\nsubmission_df = pd.DataFrame({\n    'eeg_id': test_df['eeg_id'],\n    'class_1_vote': test_predictions_processed[:, 0],\n    'class_2_vote': test_predictions_processed[:, 1],\n    'class_3_vote': test_predictions_processed[:, 2],\n    # Add more columns if there are more classes\n})\n\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save DataFrame to submission.csv\nsubmission_df.to_csv(\"submission.csv\", index=False)\n\nprint(\"Submission file saved successfully.\")\n","metadata":{},"execution_count":null,"outputs":[]}]}