{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":2378330,"sourceType":"datasetVersion","datasetId":492658},{"sourceId":8035435,"sourceType":"datasetVersion","datasetId":4312781}],"dockerImageVersionId":30636,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## HMS - Harmful Brain Activity Classification - Inference","metadata":{}},{"cell_type":"markdown","source":"## 1. Setup","metadata":{}},{"cell_type":"code","source":"%%time\nimport sys\n\n!cp ../input/rapids/rapids.0.17.0 /opt/conda/envs/rapids.tar.gz\n!cd /opt/conda/envs/ && tar -xzvf rapids.tar.gz > /dev/null\n!rm /opt/conda/envs/rapids.tar.gz\n\nsys.path += [\"/opt/conda/envs/rapids/lib/python3.7/site-packages\"]\nsys.path += [\"/opt/conda/envs/rapids/lib/python3.7\"]\nsys.path += [\"/opt/conda/envs/rapids/lib\"]\n!cp /opt/conda/envs/rapids/lib/libxgboost.so /opt/conda/lib/","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:02:11.290561Z","iopub.execute_input":"2024-04-05T08:02:11.291390Z","iopub.status.idle":"2024-04-05T08:03:32.186868Z","shell.execute_reply.started":"2024-04-05T08:02:11.291354Z","shell.execute_reply":"2024-04-05T08:03:32.185683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install /kaggle/input/hms-hbac-dataset/packages/timm-0.9.16-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:03:32.188809Z","iopub.execute_input":"2024-04-05T08:03:32.189077Z","iopub.status.idle":"2024-04-05T08:04:08.614842Z","shell.execute_reply.started":"2024-04-05T08:03:32.189052Z","shell.execute_reply":"2024-04-05T08:04:08.613641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport yaml\nimport warnings\nimport timeit\nfrom joblib import Parallel, delayed, cpu_count\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport numpy as np\nimport pandas as pd\nimport scipy\nfrom scipy.signal.windows import dpss\nfrom scipy.signal import detrend\nimport cv2\nimport cusignal\nimport librosa\nimport math\nimport colorcet\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nimport albumentations as A\nfrom albumentations import ImageOnlyTransform\nfrom albumentations.pytorch.transforms import ToTensorV2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-05T08:04:08.616357Z","iopub.execute_input":"2024-04-05T08:04:08.616687Z","iopub.status.idle":"2024-04-05T08:04:17.961016Z","shell.execute_reply.started":"2024-04-05T08:04:08.616656Z","shell.execute_reply":"2024-04-05T08:04:17.959610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"competition_dataset_directory = Path('/kaggle/input/hms-harmful-brain-activity-classification')\nexternal_dataset_directory = Path('/kaggle/input/hms-hbac-dataset')","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:04:17.963656Z","iopub.execute_input":"2024-04-05T08:04:17.964660Z","iopub.status.idle":"2024-04-05T08:04:17.969745Z","shell.execute_reply.started":"2024-04-05T08:04:17.964620Z","shell.execute_reply":"2024-04-05T08:04:17.968801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eeg_directory = competition_dataset_directory / 'test_eegs'\nspectrogram_directory = competition_dataset_directory / 'test_spectrograms'\ndf = pd.read_csv(competition_dataset_directory / 'test.csv')\n\nprint(f'Test Set Shape: {df.shape} - EEG Files: {len(os.listdir(eeg_directory))} - Spectrogram Files: {len(os.listdir(spectrogram_directory))}')\n\ndf","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:04:17.971196Z","iopub.execute_input":"2024-04-05T08:04:17.971463Z","iopub.status.idle":"2024-04-05T08:04:18.060966Z","shell.execute_reply.started":"2024-04-05T08:04:17.971439Z","shell.execute_reply":"2024-04-05T08:04:18.059363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"is_submission = df.shape[0] != 1\nis_submission","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:04:18.064871Z","iopub.execute_input":"2024-04-05T08:04:18.065571Z","iopub.status.idle":"2024-04-05T08:04:18.071547Z","shell.execute_reply.started":"2024-04-05T08:04:18.065537Z","shell.execute_reply":"2024-04-05T08:04:18.070650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_image(image, title, path=None):\n\n    \"\"\"\n    Visualize the given raw EEG or spectrogram\n\n    Parameters\n    ----------\n    image: numpy.ndarray of shape (height, width, channel) or (height, width)\n        Image array\n\n    title: str\n        Title of the plot\n\n    path: str or None\n        Path of the output file or None (if path is None, plot is displayed with selected backend)\n    \"\"\"\n\n    fig, ax = plt.subplots(figsize=(8, 8))\n    ax.imshow(image, cmap=plt.cm.coolwarm)\n    ax.set_xlabel('')\n    ax.set_ylabel('')\n    ax.tick_params(axis='x', labelsize=15, pad=10)\n    ax.tick_params(axis='y', labelsize=15, pad=10)\n    ax.set_title(title, size=15, pad=12.5, loc='center', wrap=True)\n\n    if path is None:\n        plt.show()\n    else:\n        plt.savefig(path)\n        plt.close(fig)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:04:18.072556Z","iopub.execute_input":"2024-04-05T08:04:18.072884Z","iopub.status.idle":"2024-04-05T08:04:18.083391Z","shell.execute_reply.started":"2024-04-05T08:04:18.072859Z","shell.execute_reply":"2024-04-05T08:04:18.082488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Models","metadata":{}},{"cell_type":"markdown","source":"### 2.1. Heads","metadata":{}},{"cell_type":"code","source":"class ClassificationHead(nn.Module):\n\n    def __init__(self, input_dimensions, output_dimensions):\n\n        super(ClassificationHead, self).__init__()\n\n        self.classifier = nn.Linear(input_dimensions, output_dimensions, bias=True)\n\n    def forward(self, x):\n\n        output = self.classifier(x)\n        \n        return output\n","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:04:18.084692Z","iopub.execute_input":"2024-04-05T08:04:18.085251Z","iopub.status.idle":"2024-04-05T08:04:18.097991Z","shell.execute_reply.started":"2024-04-05T08:04:18.085219Z","shell.execute_reply":"2024-04-05T08:04:18.097292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 2.2. EEGNet 1D","metadata":{}},{"cell_type":"code","source":"class ResNet1DBlock(nn.Module):\n\n    def __init__(self, in_channels, out_channels, kernel_size, stride, padding, downsampling):\n\n        super(ResNet1DBlock, self).__init__()\n\n        self.conv1 = nn.Conv1d(in_channels=in_channels, out_channels=out_channels, kernel_size=kernel_size, stride=stride, padding=padding, bias=True)\n        self.bn1 = nn.BatchNorm1d(num_features=in_channels)\n        self.activation = nn.LeakyReLU(inplace=False)\n        self.conv2 = nn.Conv1d(in_channels=out_channels, out_channels=out_channels, kernel_size=kernel_size, stride=stride, padding=padding, bias=True)\n        self.bn2 = nn.BatchNorm1d(num_features=out_channels)\n        self.pooling = nn.MaxPool1d(kernel_size=2, stride=2, padding=0)\n        self.downsampling = downsampling\n\n    def forward(self, x):\n\n        identity = x\n\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.activation(x)\n\n        x = self.conv2(x)\n        x = self.bn2(x)\n        x = self.pooling(x)\n\n        identity = self.downsampling(identity)\n        outputs = x + identity\n        outputs = self.activation(outputs)\n\n        return outputs\n\n\nclass EEGNet(nn.Module):\n\n    def __init__(self, in_channels, kernels, fixed_kernel_size, head_args):\n\n        super(EEGNet, self).__init__()\n\n        self.in_channels = in_channels\n        self.kernels = kernels\n        self.planes = 128\n        self.parallel_conv = nn.ModuleList()\n\n        for i, kernel_size in enumerate(self.kernels):\n            sep_conv = nn.Conv1d(in_channels=in_channels, out_channels=self.planes, kernel_size=kernel_size, stride=1, padding=0, bias=False)\n            self.parallel_conv.append(sep_conv)\n\n        self.bn1 = nn.BatchNorm1d(num_features=self.planes)\n        self.activation = nn.LeakyReLU(inplace=False)\n        self.conv1 = nn.Conv1d(in_channels=self.planes, out_channels=self.planes, kernel_size=fixed_kernel_size, stride=2, padding=2, bias=False)\n        self.res_blocks = self._make_resnet_layer(kernel_size=fixed_kernel_size, stride=1, padding=fixed_kernel_size // 2)\n        self.bn2 = nn.BatchNorm1d(num_features=self.planes)\n        self.pooling = nn.AvgPool1d(kernel_size=6, stride=6, padding=2)\n        self.head = ClassificationHead(input_dimensions=256, **head_args)\n\n    def _make_resnet_layer(self, kernel_size, stride, blocks=11, padding=0):\n\n        layers = []\n\n        for i in range(blocks):\n            downsampling = nn.MaxPool1d(kernel_size=2, stride=2, padding=0)\n            res_block = ResNet1DBlock(\n                in_channels=self.planes,\n                out_channels=self.planes,\n                kernel_size=kernel_size,\n                stride=stride,\n                padding=padding,\n                downsampling=downsampling\n            )\n            layers.append(res_block)\n\n        return nn.Sequential(*layers)\n\n    def forward(self, x):\n\n        out_sep = []\n\n        for i in range(len(self.kernels)):\n            sep = self.parallel_conv[i](x)\n            out_sep.append(sep)\n\n        x = torch.cat(out_sep, dim=2)\n        x = self.bn1(x)\n        x = self.activation(x)\n        x = self.conv1(x)\n\n        x = self.res_blocks(x)\n        x = self.bn2(x)\n        x = self.activation(x)\n        x = self.pooling(x)\n\n        x = x.reshape(x.shape[0], -1)\n        outputs = self.head(x)\n\n        return outputs\n","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:04:18.098933Z","iopub.execute_input":"2024-04-05T08:04:18.099226Z","iopub.status.idle":"2024-04-05T08:04:18.119669Z","shell.execute_reply.started":"2024-04-05T08:04:18.099204Z","shell.execute_reply":"2024-04-05T08:04:18.118827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_eegnet_model(model_directory, model_file_names, device):\n\n    \"\"\"\n    Load model and pretrained weights from the given model directory for EEGNet model\n\n    Parameters\n    ----------\n    model_directory: pathlib.Path\n        Path of the model directory\n\n    model_file_names: list\n        List of names of the model weights files\n\n    device: torch.device\n        Location of the model\n\n    Returns\n    -------\n    model: dict\n        Dictionary of models with pretrained weights loaded\n\n    config: dict\n        Dictionary of configurations\n    \"\"\"\n\n    config = yaml.load(open(model_directory / 'config.yaml', 'r'), Loader=yaml.FullLoader)\n\n    models = {}\n\n    for model_file_name in model_file_names:\n        model = EEGNet(**config['model']['model_args'])\n        model.load_state_dict(torch.load(model_directory / model_file_name))\n        model.to(device)\n        model.eval()\n        models[model_file_name] = model\n        print(f'Loaded EEGNet from {model_directory / model_file_name} to {device}')\n\n    return models, config\n","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:04:18.123318Z","iopub.execute_input":"2024-04-05T08:04:18.123638Z","iopub.status.idle":"2024-04-05T08:04:18.134798Z","shell.execute_reply.started":"2024-04-05T08:04:18.123589Z","shell.execute_reply":"2024-04-05T08:04:18.133955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 2.3. timm Models 2D","metadata":{}},{"cell_type":"code","source":"class EfficientNet(nn.Module):\n\n    def __init__(self, model_name, pretrained, backbone_args, pooling_type, dropout_rate, head_args):\n\n        super(EfficientNet, self).__init__()\n\n        self.backbone = timm.create_model(\n            model_name=model_name,\n            pretrained=pretrained,\n            **backbone_args\n        )\n\n        input_features = self.backbone.get_classifier().in_features\n\n        self.pooling_type = pooling_type\n        self.dropout = nn.Dropout(dropout_rate) if dropout_rate > 0 else nn.Identity()\n        self.head = ClassificationHead(input_dimensions=input_features, **head_args)\n\n    def forward(self, x):\n\n        x = self.backbone.forward_features(x)\n\n        if self.pooling_type == 'avg':\n            x = F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.pooling_type == 'max':\n            x = F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.pooling_type == 'concat':\n            x = torch.cat([\n                F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1),\n                F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n            ], dim=-1)\n\n        x = self.dropout(x)\n        output = self.head(x)\n\n        return output\n\n\nclass ConvNeXt(nn.Module):\n\n    def __init__(self, model_name, pretrained, backbone_args, pooling_type, dropout_rate, head_args):\n\n        super(ConvNeXt, self).__init__()\n\n        self.backbone = timm.create_model(\n            model_name=model_name,\n            pretrained=pretrained,\n            **backbone_args\n        )\n\n        input_features = self.backbone.get_classifier().in_features\n\n        self.pooling_type = pooling_type\n        self.dropout = nn.Dropout(dropout_rate) if dropout_rate > 0 else nn.Identity()\n        self.head = ClassificationHead(input_dimensions=input_features, **head_args)\n\n    def forward(self, x):\n\n        x = self.backbone.forward_features(x)\n\n        if self.pooling_type == 'avg':\n            x = F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.pooling_type == 'max':\n            x = F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.pooling_type == 'concat':\n            x = torch.cat([\n                F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1),\n                F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n            ], dim=-1)\n\n        x = self.dropout(x)\n        output = self.head(x)\n\n        return output\n\n    \nclass CoaT(nn.Module):\n\n    def __init__(self, model_name, pretrained, backbone_args, dropout_rate, head_args):\n\n        super(CoaT, self).__init__()\n\n        self.backbone = timm.create_model(\n            model_name=model_name,\n            pretrained=pretrained,\n            **backbone_args\n        )\n\n        input_features = self.backbone.get_classifier().in_features\n        self.backbone.head_drop = nn.Identity()\n        self.backbone.head = nn.Identity()\n        self.dropout = nn.Dropout(dropout_rate) if dropout_rate > 0 else nn.Identity()\n        self.head = ClassificationHead(input_dimensions=input_features, **head_args)\n\n    def forward(self, x):\n\n        x = self.backbone(x)\n        x = self.dropout(x)\n        output = self.head(x)\n\n        return output\n\n    \nclass CoAtNet(nn.Module):\n\n    def __init__(self, model_name, pretrained, backbone_args, pooling_type, dropout_rate, head_args):\n\n        super(CoAtNet, self).__init__()\n\n        self.backbone = timm.create_model(\n            model_name=model_name,\n            pretrained=pretrained,\n            **backbone_args\n        )\n\n        input_features = self.backbone.get_classifier().in_features\n        self.pooling_type = pooling_type\n        self.dropout = nn.Dropout(dropout_rate) if dropout_rate > 0 else nn.Identity()\n        self.head = ClassificationHead(input_dimensions=input_features, **head_args)\n\n    def forward(self, x):\n\n        x = self.backbone.forward_features(x)\n\n        if self.pooling_type == 'avg':\n            x = F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.pooling_type == 'max':\n            x = F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.pooling_type == 'concat':\n            x = torch.cat([\n                F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1),\n                F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n            ], dim=-1)\n\n        x = self.dropout(x)\n        output = self.head(x)\n\n        return output\n\n\nclass NextViT(nn.Module):\n\n    def __init__(self, model_name, pretrained, backbone_args, pooling_type, dropout_rate, head_args):\n\n        super(NextViT, self).__init__()\n\n        self.backbone = timm.create_model(\n            model_name=model_name,\n            pretrained=pretrained,\n            **backbone_args\n        )\n\n        input_features = self.backbone.get_classifier().in_features\n        self.pooling_type = pooling_type\n        self.dropout = nn.Dropout(dropout_rate) if dropout_rate > 0 else nn.Identity()\n        self.head = ClassificationHead(input_dimensions=input_features, **head_args)\n\n    def forward(self, x):\n\n        x = self.backbone.forward_features(x)\n\n        if self.pooling_type == 'avg':\n            x = F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.pooling_type == 'max':\n            x = F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.pooling_type == 'concat':\n            x = torch.cat([\n                F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1),\n                F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n            ], dim=-1)\n\n        x = self.dropout(x)\n        output = self.head(x)\n\n        return output\n\n    \nclass SwinTransformer(nn.Module):\n\n    def __init__(self, model_name, pretrained, backbone_args, pooling_type, dropout_rate, head_args):\n\n        super(SwinTransformer, self).__init__()\n\n        self.backbone = timm.create_model(\n            model_name=model_name,\n            pretrained=pretrained,\n            **backbone_args\n        )\n\n        input_features = self.backbone.get_classifier().in_features\n        self.pooling_type = pooling_type\n        self.dropout = nn.Dropout(dropout_rate) if dropout_rate > 0 else nn.Identity()\n        self.head = ClassificationHead(input_dimensions=input_features, **head_args)\n\n    def forward(self, x):\n\n        x = self.backbone.forward_features(x)\n        x = x.permute(0, -1, 1, 2)\n\n        if self.pooling_type == 'avg':\n            x = F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.pooling_type == 'max':\n            x = F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n        elif self.pooling_type == 'concat':\n            x = torch.cat([\n                F.adaptive_avg_pool2d(x, output_size=(1, 1)).view(x.size(0), -1),\n                F.adaptive_max_pool2d(x, output_size=(1, 1)).view(x.size(0), -1)\n            ], dim=-1)\n\n        x = self.dropout(x)\n        output = self.head(x)\n\n        return output\n","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:04:18.136068Z","iopub.execute_input":"2024-04-05T08:04:18.136327Z","iopub.status.idle":"2024-04-05T08:04:18.173275Z","shell.execute_reply.started":"2024-04-05T08:04:18.136305Z","shell.execute_reply":"2024-04-05T08:04:18.172543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_timm_model(model_directory, model_file_names, device):\n\n    \"\"\"\n    Load model and pretrained weights from the given model directory\n\n    Parameters\n    ----------\n    model_directory: pathlib.Path\n        Path of the model directory\n\n    model_file_names: list\n        List of names of the model weights files\n\n    device: torch.device\n        Location of the model\n\n    Returns\n    -------\n    model: dict\n        Dictionary of models with pretrained weights loaded\n\n    config: dict\n        Dictionary of configurations\n    \"\"\"\n\n    config = yaml.load(open(model_directory / 'config.yaml', 'r'), Loader=yaml.FullLoader)\n    config['model']['model_args']['pretrained'] = False\n\n    models = {}\n\n    for model_file_name in model_file_names:\n        model = eval(config['model']['model_class'])(**config['model']['model_args'])\n        model.load_state_dict(torch.load(model_directory / model_file_name))\n        model.to(device)\n        model.eval()\n        models[model_file_name] = model\n        print(f'Loaded {config[\"model\"][\"model_class\"]} model from {model_directory / model_file_name} to {device}')\n\n    return models, config\n","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:04:18.174510Z","iopub.execute_input":"2024-04-05T08:04:18.174851Z","iopub.status.idle":"2024-04-05T08:04:18.189191Z","shell.execute_reply.started":"2024-04-05T08:04:18.174822Z","shell.execute_reply":"2024-04-05T08:04:18.188349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. 50 Second Raw EEG Models","metadata":{}},{"cell_type":"code","source":"raw_eeg_50_second_2d_convnext_models, raw_eeg_50_second_2d_convnext_models_config = load_timm_model(\n    model_directory=external_dataset_directory / 'raw_eeg_50_second_2d_convnextbase_384x512_quality_2',\n    model_file_names=[\n        'model_fold_1_epoch_4_best_sq_2_kl_divergence_0.2618.pt',\n        'model_fold_2_epoch_1_best_sq_2_kl_divergence_0.2645.pt',\n        'model_fold_3_epoch_7_best_sq_2_kl_divergence_0.2568.pt',\n        'model_fold_4_epoch_4_best_sq_2_kl_divergence_0.2552.pt',\n        'model_fold_5_epoch_6_best_sq_2_kl_divergence_0.2604.pt'\n    ],\n    device=torch.device('cuda')\n)\n\nraw_eeg_50_second_2d_maxvit_models, raw_eeg_50_second_2d_maxvit_models_config = load_timm_model(\n    model_directory=external_dataset_directory / 'raw_eeg_50_second_2d_maxvittiny_384x512_quality_2',\n    model_file_names=[\n        'model_fold_1_epoch_6_best_sq_2_kl_divergence_0.2669.pt',\n        'model_fold_2_epoch_6_best_sq_2_kl_divergence_0.2663.pt',\n        'model_fold_3_epoch_7_best_sq_2_kl_divergence_0.2592.pt',\n        'model_fold_4_epoch_6_best_sq_2_kl_divergence_0.2654.pt',\n        'model_fold_5_epoch_7_best_sq_2_kl_divergence_0.2592.pt'\n    ],\n    device=torch.device('cuda')\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:04:18.190259Z","iopub.execute_input":"2024-04-05T08:04:18.191110Z","iopub.status.idle":"2024-04-05T08:04:50.439743Z","shell.execute_reply.started":"2024-04-05T08:04:18.191078Z","shell.execute_reply":"2024-04-05T08:04:50.438759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4. 50/10 Second Raw EEG Models","metadata":{}},{"cell_type":"code","source":"raw_eeg_50_10_second_2d_convnext_models, raw_eeg_50_10_second_2d_convnext_models_config = load_timm_model(\n    model_directory=external_dataset_directory / 'raw_eeg_50_10_second_2d_convnextbase_448x512_quality_2',\n    model_file_names=[\n        'model_fold_1_epoch_7_best_sq_2_kl_divergence_0.2532.pt',\n        'model_fold_2_epoch_2_best_sq_2_kl_divergence_0.2710.pt',\n        'model_fold_3_epoch_8_best_sq_2_kl_divergence_0.2515.pt',\n        'model_fold_4_epoch_7_best_sq_2_kl_divergence_0.2544.pt',\n        'model_fold_5_epoch_8_best_sq_2_kl_divergence_0.2602.pt'\n    ],\n    device=torch.device('cuda')\n)\n\nraw_eeg_50_10_second_2d_maxvit_models, raw_eeg_50_10_second_2d_maxvit_models_config = load_timm_model(\n    model_directory=external_dataset_directory / 'raw_eeg_50_10_second_2d_maxvittiny_448x512_quality_2',\n    model_file_names=[\n        'model_fold_1_epoch_6_best_sq_2_kl_divergence_0.2679.pt',\n        'model_fold_2_epoch_5_best_sq_2_kl_divergence_0.2619.pt',\n        'model_fold_3_epoch_6_best_sq_2_kl_divergence_0.2532.pt',\n        'model_fold_4_epoch_6_best_sq_2_kl_divergence_0.2576.pt',\n        'model_fold_5_epoch_5_best_sq_2_kl_divergence_0.2654.pt'\n    ],\n    device=torch.device('cuda')\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:04:50.440852Z","iopub.execute_input":"2024-04-05T08:04:50.441137Z","iopub.status.idle":"2024-04-05T08:05:21.059793Z","shell.execute_reply.started":"2024-04-05T08:04:50.441111Z","shell.execute_reply":"2024-04-05T08:05:21.058753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5. 50 Second Spectrogram Models","metadata":{}},{"cell_type":"code","source":"spectrogram_50_second_2d_convnext_models, spectrogram_50_second_2d_convnext_models_config = load_timm_model(\n    model_directory=external_dataset_directory / 'spectrogram_50_second_2d_convnextbase_512x512_quality_2',\n    model_file_names=[\n        'model_fold_1_epoch_7_best_sq_2_kl_divergence_0.2475.pt',\n        'model_fold_2_epoch_7_best_sq_2_kl_divergence_0.2427.pt',\n        'model_fold_3_epoch_7_best_sq_2_kl_divergence_0.2382.pt',\n        'model_fold_4_epoch_5_best_sq_2_kl_divergence_0.2395.pt',\n        'model_fold_5_epoch_6_best_sq_2_kl_divergence_0.2581.pt'\n    ],\n    device=torch.device('cuda')\n)\n\nspectrogram_50_second_2d_maxvit_models, spectrogram_50_second_2d_maxvit_models_config = load_timm_model(\n    model_directory=external_dataset_directory / 'spectrogram_50_second_2d_maxvittiny_512x512_quality_2',\n    model_file_names=[\n        'model_fold_1_epoch_7_best_sq_2_kl_divergence_0.2554.pt',\n        'model_fold_2_epoch_5_best_sq_2_kl_divergence_0.2333.pt',\n        'model_fold_3_epoch_6_best_sq_2_kl_divergence_0.2431.pt',\n        'model_fold_4_epoch_7_best_sq_2_kl_divergence_0.2546.pt',\n        'model_fold_5_epoch_7_best_sq_2_kl_divergence_0.2440.pt'\n    ],\n    device=torch.device('cuda')\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:05:21.061275Z","iopub.execute_input":"2024-04-05T08:05:21.061696Z","iopub.status.idle":"2024-04-05T08:05:51.400410Z","shell.execute_reply.started":"2024-04-05T08:05:21.061660Z","shell.execute_reply":"2024-04-05T08:05:51.399501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 6. 50/10 Second Spectrogram Models","metadata":{}},{"cell_type":"code","source":"spectrogram_50_10_second_2d_convnext_models, spectrogram_50_10_second_2d_convnext_models_config = load_timm_model(\n    model_directory=external_dataset_directory / 'spectrogram_50_10_second_2d_convnextbase_512x512_quality_2',\n    model_file_names=[\n        'model_fold_1_epoch_7_best_sq_2_kl_divergence_0.2404.pt',\n        'model_fold_2_epoch_4_best_sq_2_kl_divergence_0.2328.pt',\n        'model_fold_3_epoch_2_best_sq_2_kl_divergence_0.2399.pt',\n        'model_fold_4_epoch_8_best_sq_2_kl_divergence_0.2366.pt',\n        'model_fold_5_epoch_5_best_sq_2_kl_divergence_0.2471.pt'\n    ],\n    device=torch.device('cuda')\n)\n\nspectrogram_50_10_second_2d_maxvit_models, spectrogram_50_10_second_2d_maxvit_models_config = load_timm_model(\n    model_directory=external_dataset_directory / 'spectrogram_50_10_second_2d_maxvittiny_512x512_quality_2',\n    model_file_names=[\n        'model_fold_1_epoch_6_best_sq_2_kl_divergence_0.2468.pt',\n        'model_fold_2_epoch_6_best_sq_2_kl_divergence_0.2230.pt',\n        'model_fold_3_epoch_6_best_sq_2_kl_divergence_0.2417.pt',\n        'model_fold_4_epoch_7_best_sq_2_kl_divergence_0.2536.pt',\n        'model_fold_5_epoch_6_best_sq_2_kl_divergence_0.2403.pt'\n    ],\n    device=torch.device('cuda')\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:05:51.401612Z","iopub.execute_input":"2024-04-05T08:05:51.401916Z","iopub.status.idle":"2024-04-05T08:06:22.608946Z","shell.execute_reply.started":"2024-04-05T08:05:51.401889Z","shell.execute_reply":"2024-04-05T08:06:22.607989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 7. 50/30/10 Second Spectrogram Models","metadata":{}},{"cell_type":"code","source":"spectrogram_50_30_10_second_2d_convnext_models, spectrogram_50_30_10_second_2d_convnext_models_config = load_timm_model(\n    model_directory=external_dataset_directory / 'spectrogram_50_30_10_second_2d_convnextbase_512x512_quality_2',\n    model_file_names=[\n        'model_fold_1_epoch_6_best_sq_2_kl_divergence_0.2506.pt',\n        'model_fold_2_epoch_7_best_sq_2_kl_divergence_0.2326.pt',\n        'model_fold_3_epoch_7_best_sq_2_kl_divergence_0.2384.pt',\n        'model_fold_4_epoch_8_best_sq_2_kl_divergence_0.2457.pt',\n        'model_fold_5_epoch_6_best_sq_2_kl_divergence_0.2472.pt'\n    ],\n    device=torch.device('cuda')\n)\n\nspectrogram_50_30_10_second_2d_maxvit_models, spectrogram_50_30_10_second_2d_maxvit_models_config = load_timm_model(\n    model_directory=external_dataset_directory / 'spectrogram_50_30_10_second_2d_maxvittiny_512x512_quality_2',\n    model_file_names=[\n        'model_fold_1_epoch_7_best_sq_2_kl_divergence_0.2487.pt',\n        'model_fold_2_epoch_6_best_sq_2_kl_divergence_0.2352.pt',\n        'model_fold_3_epoch_5_best_sq_2_kl_divergence_0.2394.pt',\n        'model_fold_4_epoch_7_best_sq_2_kl_divergence_0.2422.pt',\n        'model_fold_5_epoch_6_best_sq_2_kl_divergence_0.2475.pt'\n    ],\n    device=torch.device('cuda')\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:06:22.610305Z","iopub.execute_input":"2024-04-05T08:06:22.611022Z","iopub.status.idle":"2024-04-05T08:06:53.202251Z","shell.execute_reply.started":"2024-04-05T08:06:22.610986Z","shell.execute_reply":"2024-04-05T08:06:53.201093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 8. 30/10 Second Spectrogram Models","metadata":{}},{"cell_type":"code","source":"spectrogram_30_10_second_2d_convnext_models, spectrogram_30_10_second_2d_convnext_models_config = load_timm_model(\n    model_directory=external_dataset_directory / 'spectrogram_30_10_second_2d_convnextbase_512x512_quality_2',\n    model_file_names=[\n        'model_fold_1_epoch_7_best_sq_2_kl_divergence_0.2408.pt',\n        'model_fold_2_epoch_4_best_sq_2_kl_divergence_0.2379.pt',\n        'model_fold_3_epoch_5_best_sq_2_kl_divergence_0.2340.pt',\n        'model_fold_4_epoch_8_best_sq_2_kl_divergence_0.2359.pt',\n        'model_fold_5_epoch_6_best_sq_2_kl_divergence_0.2430.pt'\n    ],\n    device=torch.device('cuda')\n)\n\nspectrogram_30_10_second_2d_maxvit_models, spectrogram_30_10_second_2d_maxvit_models_config = load_timm_model(\n    model_directory=external_dataset_directory / 'spectrogram_30_10_second_2d_maxvittiny_512x512_quality_2',\n    model_file_names=[\n        'model_fold_1_epoch_6_best_sq_2_kl_divergence_0.2548.pt',\n        'model_fold_2_epoch_6_best_sq_2_kl_divergence_0.2264.pt',\n        'model_fold_3_epoch_7_best_sq_2_kl_divergence_0.2306.pt',\n        'model_fold_4_epoch_8_best_sq_2_kl_divergence_0.2491.pt',\n        'model_fold_5_epoch_6_best_sq_2_kl_divergence_0.2483.pt'\n    ],\n    device=torch.device('cuda')\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:06:53.203402Z","iopub.execute_input":"2024-04-05T08:06:53.203712Z","iopub.status.idle":"2024-04-05T08:07:02.596425Z","shell.execute_reply.started":"2024-04-05T08:06:53.203686Z","shell.execute_reply":"2024-04-05T08:07:02.595475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:07:29.723338Z","iopub.execute_input":"2024-04-05T08:07:29.723730Z","iopub.status.idle":"2024-04-05T08:07:30.794332Z","shell.execute_reply.started":"2024-04-05T08:07:29.723698Z","shell.execute_reply":"2024-04-05T08:07:30.793253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Raw EEG Transforms","metadata":{}},{"cell_type":"code","source":"class ChannelDifference1D(ImageOnlyTransform):\n\n    def __init__(self, ekg, always_apply=True, p=1.0):\n\n        super(ChannelDifference1D, self).__init__(always_apply=always_apply, p=p)\n\n        self.ekg = ekg\n        fs = 200\n        normal_cutoff = 20.5 / (0.5 * fs)\n        self.sos = scipy.signal.butter(1, normal_cutoff, btype='low', output='sos')\n\n    def apply(self, inputs, **kwargs):\n\n        \"\"\"\n        Create bipolar channel difference features\n\n        Parameters\n        ----------\n        inputs: numpy.ndarray of shape (time, 20)\n             Inputs array\n\n        Returns\n        -------\n        inputs: numpy.ndarray of shape (time, 18 or 19)\n            Inputs array with difference features\n        \"\"\"\n\n        n_channels = 19 if self.ekg else 18\n        inputs_difference = np.zeros((inputs.shape[0], n_channels))\n\n        # Left outside temporal chain\n        inputs_difference[:, 0] = inputs[:, 0] - inputs[:, 4]\n        inputs_difference[:, 1] = inputs[:, 4] - inputs[:, 5]\n        inputs_difference[:, 2] = inputs[:, 5] - inputs[:, 6]\n        inputs_difference[:, 3] = inputs[:, 6] - inputs[:, 7]\n\n        # Right outside temporal chain\n        inputs_difference[:, 4] = inputs[:, 11] - inputs[:, 15]\n        inputs_difference[:, 5] = inputs[:, 15] - inputs[:, 16]\n        inputs_difference[:, 6] = inputs[:, 16] - inputs[:, 17]\n        inputs_difference[:, 7] = inputs[:, 17] - inputs[:, 18]\n\n        # Left inside parasagittal chain\n        inputs_difference[:, 8] = inputs[:, 0] - inputs[:, 1]\n        inputs_difference[:, 9] = inputs[:, 1] - inputs[:, 2]\n        inputs_difference[:, 10] = inputs[:, 2] - inputs[:, 3]\n        inputs_difference[:, 11] = inputs[:, 3] - inputs[:, 7]\n\n        # Right inside parasagittal chain\n        inputs_difference[:, 12] = inputs[:, 11] - inputs[:, 12]\n        inputs_difference[:, 13] = inputs[:, 12] - inputs[:, 13]\n        inputs_difference[:, 14] = inputs[:, 13] - inputs[:, 14]\n        inputs_difference[:, 15] = inputs[:, 14] - inputs[:, 18]\n\n        # Center chain\n        inputs_difference[:, 16] = inputs[:, 8] - inputs[:, 9]\n        inputs_difference[:, 17] = inputs[:, 9] - inputs[:, 10]\n\n        if self.ekg:\n            inputs_difference[:, 18] = inputs[:, 19]\n            \n        inputs_difference = scipy.signal.sosfiltfilt(self.sos, inputs_difference, axis=0)\n\n        return inputs_difference\n\n\nclass ChannelGroupPermute1D(ImageOnlyTransform):\n\n    def __init__(self, ekg, always_apply=False, p=0.5):\n\n        super(ChannelGroupPermute1D, self).__init__(always_apply=always_apply, p=p)\n\n        self.ekg = ekg\n\n    def apply(self, inputs, **kwargs):\n\n        \"\"\"\n        Permute 4 main channel groups\n\n        Parameters\n        ----------\n        inputs: numpy.ndarray of shape (time, 18 or 19)\n            Inputs array with difference features\n\n        Returns\n        -------\n        inputs_permuted: numpy.ndarray of shape (time, 18 or 19)\n            Inputs array with difference features permuted\n        \"\"\"\n\n        groups = np.array([\n            [0, 1, 2, 3],\n            [4, 5, 6, 7],\n            [8, 9, 10, 11],\n            [12, 13, 14, 15]\n        ])\n        inputs_permuted = np.zeros_like(inputs)\n        for group_idx, permuted_group_idx in enumerate(np.random.permutation(np.arange(len(groups)))):\n            # Permute channel groups of LL, RL, LP and RP chains\n            group = groups[group_idx]\n            permuted_group = groups[permuted_group_idx]\n            inputs_permuted[group] = inputs[permuted_group]\n\n        # Center chain and EKG are not permuted\n        inputs_permuted[:, [16, 17]] = inputs[:, [16, 17]]\n        if self.ekg:\n            inputs_permuted[:, 18] = inputs[:, 18]\n\n        return inputs_permuted\n\n\nclass InstanceNormalization2D(ImageOnlyTransform):\n\n    def __init__(self, per_channel, epsilon=1e-15, always_apply=True, p=1.0):\n\n        super(InstanceNormalization2D, self).__init__(always_apply=always_apply, p=p)\n\n        self.per_channel = per_channel\n        self.epsilon = epsilon\n\n    def apply(self, inputs, **kwargs):\n\n        \"\"\"\n        Normalize inputs by mean and standard deviation of itself\n\n        Parameters\n        ----------\n        inputs: numpy.ndarray of shape (height, width, channel)\n             Inputs array\n\n        Returns\n        -------\n        inputs: numpy.ndarray of shape (height, width, channel)\n            Normalized inputs array\n        \"\"\"\n\n        if self.per_channel:\n            axis = (0, 1)\n        else:\n            axis = (0, 1, 2)\n\n        mean = np.mean(inputs, axis=axis)\n        std = np.std(inputs, axis=axis)\n        inputs = (inputs - mean) / (std + self.epsilon)\n\n        return inputs\n\n\nclass CenterTemporalDropout1D(ImageOnlyTransform):\n\n    def __init__(self, n_time_steps, drop_value=0, always_apply=False, p=0.5):\n\n        super(CenterTemporalDropout1D, self).__init__(always_apply=always_apply, p=p)\n\n        self.n_time_steps = n_time_steps\n        self.drop_value = drop_value\n\n    def apply(self, inputs, **kwargs):\n\n        \"\"\"\n        Drop time steps from center 10 seconds randomly\n\n        Parameters\n        ----------\n        inputs: numpy.ndarray of shape (time, channel)\n             Inputs array\n\n        Returns\n        -------\n        inputs: numpy.ndarray of shape (time, channel)\n            Inputs array with dropped time steps from center 10 seconds\n        \"\"\"\n\n        # Uniformly sample N amount of indices to drop on time axis at center 10 seconds\n        center_start = 4000\n        center_end = 6000\n        drop_idx = np.random.choice(np.arange(center_start, center_end + 1), self.n_time_steps, replace=False)\n        inputs[drop_idx, :] = self.drop_value\n\n        return inputs\n\n\nclass NonCenterTemporalDropout1D(ImageOnlyTransform):\n\n    def __init__(self, n_time_steps, drop_value=0, always_apply=False, p=0.5):\n\n        super(NonCenterTemporalDropout1D, self).__init__(always_apply=always_apply, p=p)\n\n        self.n_time_steps = n_time_steps\n        self.drop_value = drop_value\n\n    def apply(self, inputs, **kwargs):\n\n        \"\"\"\n        Drop time steps from other than center 10 seconds randomly\n\n        Parameters\n        ----------\n        inputs: numpy.ndarray of shape (time, channel)\n             Inputs array\n\n        Returns\n        -------\n        inputs: numpy.ndarray of shape (time, channel)\n            Inputs array with dropped time steps except from center 10 seconds\n        \"\"\"\n\n        # Uniformly sample N amount of indices to drop on time axis except from center 10 seconds\n        center_10_seconds_start = 4000\n        center_10_seconds_end = 6000\n        drop_idx = np.random.choice(\n            np.arange(0, center_10_seconds_start).tolist() + np.arange(center_10_seconds_end + 1, inputs.shape[1]).tolist(),\n            self.n_time_steps,\n            replace=False\n        )\n        inputs[drop_idx, :] = self.drop_value\n\n        return inputs\n\n\nclass EEGTo2D(ImageOnlyTransform):\n\n    def __init__(self, ekg, center_stack, always_apply=True, p=1.0):\n\n        super(EEGTo2D, self).__init__(always_apply=always_apply, p=p)\n\n        self.ekg = ekg\n        self.center_stack = center_stack\n\n    def apply(self, inputs, **kwargs):\n\n        \"\"\"\n        Convert 1D EEG to vertically stacked 2D\n\n        Parameters\n        ----------\n        inputs: numpy.ndarray of shape (time, channel)\n             1D inputs array\n\n        Returns\n        -------\n        image: numpy.ndarray of shape (height, width)\n            2D inputs array\n        \"\"\"\n\n        inputs = inputs.T\n        \n        if self.center_stack:\n            center = inputs[:, 4000:6000]\n\n        image = []\n        n_channels = 19 if self.ekg else 18\n\n        for i in range(n_channels):\n            for j in range(20):\n                image.append(inputs[i, j::20])\n\n            if self.center_stack:\n                for j in range(4):\n                    image.append(center[i, j::4])\n\n        image = np.stack(image, axis=0)\n\n        return image","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:07:35.963065Z","iopub.execute_input":"2024-04-05T08:07:35.963429Z","iopub.status.idle":"2024-04-05T08:07:35.999910Z","shell.execute_reply.started":"2024-04-05T08:07:35.963398Z","shell.execute_reply":"2024-04-05T08:07:35.998959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_raw_eeg_2d_transforms(**transform_parameters):\n\n    \"\"\"\n    Get raw EEG 2D transforms for dataset\n\n    Parameters\n    ----------\n    transform_parameters: dict\n        Dictionary of transform parameters\n\n    Returns\n    -------\n    eeg_transforms: dict\n        Transforms for training and inference\n    \"\"\"\n\n    training_transforms = A.Compose([\n        ChannelDifference1D(transform_parameters['ekg'], always_apply=True),\n        ChannelGroupPermute1D(transform_parameters['ekg'], p=transform_parameters['channel_group_permute_probability']),\n        CenterTemporalDropout1D(\n            n_time_steps=transform_parameters['center_temporal_dropout_time_steps'],\n            drop_value=0,\n            p=transform_parameters['center_temporal_dropout_probability']\n        ),\n        NonCenterTemporalDropout1D(\n            n_time_steps=transform_parameters['non_center_temporal_dropout_time_steps'],\n            drop_value=0,\n            p=transform_parameters['non_center_temporal_dropout_probability']\n        ),\n        EEGTo2D(transform_parameters['ekg'], transform_parameters['center_stack'], always_apply=True),\n        A.CoarseDropout(\n            max_holes=transform_parameters['coarse_dropout_max_holes'],\n            min_holes=transform_parameters['coarse_dropout_min_holes'],\n            max_height=transform_parameters['coarse_dropout_max_height'],\n            max_width=transform_parameters['coarse_dropout_max_width'],\n            min_height=transform_parameters['coarse_dropout_min_height'],\n            min_width=transform_parameters['coarse_dropout_min_width'],\n            fill_value=0,\n            p=transform_parameters['coarse_dropout_probability']\n        ),\n        A.VerticalFlip(p=transform_parameters['vertical_flip_probability']),\n        A.HorizontalFlip(p=transform_parameters['horizontal_flip_probability']),\n        A.PadIfNeeded(\n            min_height=transform_parameters['pad_min_height'],\n            min_width=transform_parameters['pad_min_width'],\n            border_mode=cv2.BORDER_CONSTANT,\n            value=0\n        ),\n        InstanceNormalization2D(per_channel=True, always_apply=True),\n        ToTensorV2(always_apply=True)\n    ])\n\n    inference_transforms = A.Compose([\n        ChannelDifference1D(transform_parameters['ekg'], always_apply=True),\n        EEGTo2D(transform_parameters['ekg'], transform_parameters['center_stack'], always_apply=True),\n        A.PadIfNeeded(\n            min_height=transform_parameters['pad_min_height'],\n            min_width=transform_parameters['pad_min_width'],\n            border_mode=cv2.BORDER_CONSTANT,\n            value=0\n        ),\n        InstanceNormalization2D(per_channel=True, always_apply=True),\n        ToTensorV2(always_apply=True)\n    ])\n\n    eeg_transforms = {'training': training_transforms, 'inference': inference_transforms}\n    return eeg_transforms\n","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:07:37.160794Z","iopub.execute_input":"2024-04-05T08:07:37.161442Z","iopub.status.idle":"2024-04-05T08:07:37.172874Z","shell.execute_reply.started":"2024-04-05T08:07:37.161404Z","shell.execute_reply":"2024-04-05T08:07:37.171900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Spectrogram Transforms","metadata":{}},{"cell_type":"code","source":"class InstanceNormalization2D(ImageOnlyTransform):\n\n    def __init__(self, per_channel, epsilon=1e-15, always_apply=True, p=1.0):\n\n        super(InstanceNormalization2D, self).__init__(always_apply=always_apply, p=p)\n\n        self.per_channel = per_channel\n        self.epsilon = epsilon\n\n    def apply(self, inputs, **kwargs):\n\n        \"\"\"\n        Normalize inputs by mean and standard deviation of itself\n\n        Parameters\n        ----------\n        inputs: numpy.ndarray of shape (height, width, channel)\n             Inputs array\n\n        Returns\n        -------\n        inputs: numpy.ndarray of shape (height, width, channel)\n            Normalized inputs array\n        \"\"\"\n\n        if self.per_channel:\n            axis = (0, 1)\n        else:\n            axis = (0, 1, 2)\n\n        mean = np.mean(inputs, axis=axis)\n        std = np.std(inputs, axis=axis)\n        inputs = (inputs - mean) / (std + self.epsilon)\n\n        return inputs\n\n\nclass CenterTemporalDropout2D(ImageOnlyTransform):\n\n    def __init__(self, center_idx, n_time_steps, drop_value=0, always_apply=False, p=0.5):\n\n        super(CenterTemporalDropout2D, self).__init__(always_apply=always_apply, p=p)\n\n        self.center_idx = center_idx\n        self.n_time_steps = n_time_steps\n        self.drop_value = drop_value\n\n    def apply(self, inputs, **kwargs):\n\n        \"\"\"\n        Drop time steps from center 10 seconds randomly\n\n        Parameters\n        ----------\n        inputs: numpy.ndarray of shape (time, channel)\n             Inputs array\n\n        Returns\n        -------\n        inputs: numpy.ndarray of shape (time, channel)\n            Inputs array with dropped time steps from center 10 seconds\n        \"\"\"\n\n        # Uniformly sample N amount of indices to drop on time axis at center 10 seconds\n        center_start = self.center_idx[0]\n        center_end = self.center_idx[1]\n        drop_idx = np.random.choice(np.arange(center_start, center_end), self.n_time_steps, replace=False)\n        inputs[drop_idx, :] = self.drop_value\n\n        return inputs\n\n\nclass NonCenterTemporalDropout2D(ImageOnlyTransform):\n\n    def __init__(self, center_idx, n_time_steps, drop_value=0, always_apply=False, p=0.5):\n\n        super(NonCenterTemporalDropout2D, self).__init__(always_apply=always_apply, p=p)\n\n        self.center_idx = center_idx\n        self.n_time_steps = n_time_steps\n        self.drop_value = drop_value\n\n    def apply(self, inputs, **kwargs):\n\n        \"\"\"\n        Drop time steps except from center 10 seconds randomly\n\n        Parameters\n        ----------\n        inputs: numpy.ndarray of shape (time, channel)\n             Inputs array\n\n        Returns\n        -------\n        inputs: numpy.ndarray of shape (time, channel)\n            Inputs array with dropped time steps except from center 10 seconds\n        \"\"\"\n\n        # Uniformly sample N amount of indices to drop on time axis except from center 10 seconds\n        center_10_seconds_start = self.center_idx[0]\n        center_10_seconds_end = self.center_idx[1]\n        drop_idx = np.random.choice(\n            np.arange(0, center_10_seconds_start).tolist() + np.arange(center_10_seconds_end + 1, image.shape[1]).tolist(),\n            self.n_time_steps,\n            replace=False\n        )\n        inputs[:, drop_idx] = self.drop_value\n\n        return inputs\n\n\nclass FrequencyDropout2D(ImageOnlyTransform):\n\n    def __init__(self, n_frequencies, n_consecutive=0, drop_value=0, always_apply=False, p=0.5):\n\n        super(FrequencyDropout2D, self).__init__(always_apply=always_apply, p=p)\n\n        self.n_frequencies = n_frequencies\n        self.n_consecutive = n_consecutive\n        self.drop_value = drop_value\n\n    def apply(self, inputs, **kwargs):\n\n        \"\"\"\n        Drop frequencies randomly\n\n        Parameters\n        ----------\n        inputs: numpy.ndarray of shape (time, channel)\n             Inputs array\n\n        Returns\n        -------\n        inputs: numpy.ndarray of shape (time, channel)\n             Inputs array with dropped frequencies\n        \"\"\"\n\n        # Uniformly sample N amount of indices to drop on frequency axis\n        drop_idx = np.random.randint(0, inputs.shape[0], self.n_frequencies).tolist()\n\n        if self.n_consecutive > 0:\n\n            consecutive_drop_idx = []\n\n            for idx in drop_idx:\n\n                consecutive_drop_idx.append(idx)\n                # Uniformly sample N amount of consecutive indices to drop for each sampled drop index\n                consecutive_count = np.random.randint(0, self.n_consecutive + 1)\n\n                if consecutive_count > 0:\n                    for consecutive_idx in np.arange(1, consecutive_count + 1):\n                        # Add each consecutive index to current drop index and append it indices that'll be dropped\n                        consecutive_drop_idx.append(min(inputs.shape[0] - 1, (idx + consecutive_idx)))\n\n            inputs[consecutive_drop_idx, :] = self.drop_value\n        else:\n            inputs[drop_idx, :] = self.drop_value\n\n        return inputs\n","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:07:39.029471Z","iopub.execute_input":"2024-04-05T08:07:39.030355Z","iopub.status.idle":"2024-04-05T08:07:39.050520Z","shell.execute_reply.started":"2024-04-05T08:07:39.030315Z","shell.execute_reply":"2024-04-05T08:07:39.049498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_spectrogram_2d_transforms(**transform_parameters):\n\n    \"\"\"\n    Get spectrogram 2D transforms for dataset\n\n    Parameters\n    ----------\n    transform_parameters: dict\n        Dictionary of transform parameters\n\n    Returns\n    -------\n    spectrogram_transforms: dict\n        Transforms for training and inference\n    \"\"\"\n\n    training_transforms = A.Compose([\n        CenterTemporalDropout2D(\n            center_idx=transform_parameters['center_idx'],\n            n_time_steps=transform_parameters['center_temporal_dropout_time_steps'],\n            drop_value=0,\n            p=transform_parameters['center_temporal_dropout_probability']\n        ),\n        NonCenterTemporalDropout2D(\n            center_idx=transform_parameters['center_idx'],\n            n_time_steps=transform_parameters['non_center_temporal_dropout_time_steps'],\n            drop_value=0,\n            p=transform_parameters['non_center_temporal_dropout_probability']\n        ),\n        FrequencyDropout2D(\n           n_frequencies=transform_parameters['n_frequencies'],\n           n_consecutive=transform_parameters['n_consecutive'],\n           drop_value=0,\n           p=transform_parameters['frequency_dropout_probability']\n        ),\n        A.CoarseDropout(\n            max_holes=transform_parameters['coarse_dropout_max_holes'],\n            min_holes=transform_parameters['coarse_dropout_min_holes'],\n            max_height=transform_parameters['coarse_dropout_max_height'],\n            max_width=transform_parameters['coarse_dropout_max_width'],\n            min_height=transform_parameters['coarse_dropout_min_height'],\n            min_width=transform_parameters['coarse_dropout_min_width'],\n            fill_value=0,\n            p=transform_parameters['coarse_dropout_probability']\n        ),\n        A.VerticalFlip(p=transform_parameters['vertical_flip_probability']),\n        A.HorizontalFlip(p=transform_parameters['horizontal_flip_probability']),\n        A.PadIfNeeded(\n            min_height=transform_parameters['pad_min_height'],\n            min_width=transform_parameters['pad_min_width'],\n            border_mode=cv2.BORDER_CONSTANT,\n            value=0\n        ),\n        InstanceNormalization2D(per_channel=True, always_apply=True),\n        ToTensorV2(always_apply=True)\n    ])\n\n    inference_transforms = A.Compose([\n        A.PadIfNeeded(\n            min_height=transform_parameters['pad_min_height'],\n            min_width=transform_parameters['pad_min_width'],\n            border_mode=cv2.BORDER_CONSTANT,\n            value=0\n        ),\n        InstanceNormalization2D(per_channel=True, always_apply=True),\n        ToTensorV2(always_apply=True)\n    ])\n\n    spectrogram_transforms = {'training': training_transforms, 'inference': inference_transforms}\n    return spectrogram_transforms\n","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:07:42.542518Z","iopub.execute_input":"2024-04-05T08:07:42.542904Z","iopub.status.idle":"2024-04-05T08:07:42.553967Z","shell.execute_reply.started":"2024-04-05T08:07:42.542871Z","shell.execute_reply":"2024-04-05T08:07:42.553004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"# MULTITAPER SPECTROGRAM #\ndef multitaper_spectrogram(data, fs, frequency_range=None, time_bandwidth=5, num_tapers=None, window_params=None,\n                           min_nfft=0, detrend_opt='linear', multiprocess=False, n_jobs=None, weighting='unity',\n                           plot_on=True, return_fig=False, clim_scale=True, verbose=True, xyflip=False, ax=None):\n    \"\"\" Compute multitaper spectrogram of timeseries data\n    Usage:\n    mt_spectrogram, stimes, sfreqs = multitaper_spectrogram(data, fs, frequency_range=None, time_bandwidth=5,\n                                                            num_tapers=None, window_params=None, min_nfft=0,\n                                                            detrend_opt='linear', multiprocess=False, cpus=False,\n                                                            weighting='unity', plot_on=True, return_fig=False,\n                                                            clim_scale=True, verbose=True, xyflip=False):\n        Arguments:\n                data (1d np.array): time series data -- required\n                fs (float): sampling frequency in Hz  -- required\n                frequency_range (list): 1x2 list - [<min frequency>, <max frequency>] (default: [0 nyquist])\n                time_bandwidth (float): time-half bandwidth product (window duration*half bandwidth of main lobe)\n                                        (default: 5 Hz*s)\n                num_tapers (int): number of DPSS tapers to use (default: [will be computed\n                                  as floor(2*time_bandwidth - 1)])\n                window_params (list): 1x2 list - [window size (seconds), step size (seconds)] (default: [5 1])\n                detrend_opt (string): detrend data window ('linear' (default), 'constant', 'off')\n                                      (Default: 'linear')\n                min_nfft (int): minimum allowable NFFT size, adds zero padding for interpolation (closest 2^x)\n                                (default: 0)\n                multiprocess (bool): Use multiprocessing to compute multitaper spectrogram (default: False)\n                n_jobs (int): Number of cpus to use if multiprocess = True (default: False). Note: if default is left\n                            as None and multiprocess = True, the number of cpus used for multiprocessing will be\n                            all available - 1.\n                weighting (str): weighting of tapers ('unity' (default), 'eigen', 'adapt');\n                plot_on (bool): plot results (default: True)\n                return_fig (bool): return plotted spectrogram (default: False)\n                clim_scale (bool): automatically scale the colormap on the plotted spectrogram (default: True)\n                verbose (bool): display spectrogram properties (default: True)\n                xyflip (bool): transpose the mt_spectrogram output (default: False)\n                ax (axes): a matplotlib axes to plot the spectrogram on (default: None)\n        Returns:\n                mt_spectrogram (TxF np array): spectral power matrix\n                stimes (1xT np array): timepoints (s) in mt_spectrogram\n                sfreqs (1xF np array)L frequency values (Hz) in mt_spectrogram\n\n        Example:\n        In this example we create some chirp data and run the multitaper spectrogram on it.\n            import numpy as np  # import numpy\n            from scipy.signal import chirp  # import chirp generation function\n            # Set spectrogram params\n            fs = 200  # Sampling Frequency\n            frequency_range = [0, 25]  # Limit frequencies from 0 to 25 Hz\n            time_bandwidth = 3  # Set time-half bandwidth\n            num_tapers = 5  # Set number of tapers (optimal is time_bandwidth*2 - 1)\n            window_params = [4, 1]  # Window size is 4s with step size of 1s\n            min_nfft = 0  # No minimum nfft\n            detrend_opt = 'constant'  # detrend each window by subtracting the average\n            multiprocess = True  # use multiprocessing\n            cpus = 3  # use 3 cores in multiprocessing\n            weighting = 'unity'  # weight each taper at 1\n            plot_on = True  # plot spectrogram\n            return_fig = False  # do not return plotted spectrogram\n            clim_scale = False # don't auto-scale the colormap\n            verbose = True  # print extra info\n            xyflip = False  # do not transpose spect output matrix\n\n            # Generate sample chirp data\n            t = np.arange(1/fs, 600, 1/fs)  # Create 10 min time array from 1/fs to 600 stepping by 1/fs\n            f_start = 1  # Set chirp freq range min (Hz)\n            f_end = 20  # Set chirp freq range max (Hz)\n            data = chirp(t, f_start, t[-1], f_end, 'logarithmic')\n            # Compute the multitaper spectrogram\n            spect, stimes, sfreqs = multitaper_spectrogram(data, fs, frequency_range, time_bandwidth, num_tapers,\n                                                           window_params, min_nfft, detrend_opt, multiprocess,\n                                                           cpus, weighting, plot_on, return_fig, clim_scale,\n                                                           verbose, xyflip):\n\n        This code is companion to the paper:\n        \"Sleep Neurophysiological Dynamics Through the Lens of Multitaper Spectral Analysis\"\n           Michael J. Prerau, Ritchie E. Brown, Matt T. Bianchi, Jeffrey M. Ellenbogen, Patrick L. Purdon\n           December 7, 2016 : 60-92\n           DOI: 10.1152/physiol.00062.2015\n         which should be cited for academic use of this code.\n\n         A full tutorial on the multitaper spectrogram can be found at: # https://www.sleepEEG.org/multitaper\n\n        Copyright 2021 Michael J. Prerau Laboratory. - https://www.sleepEEG.org\n        Authors: Michael J. Prerau, Ph.D., Thomas Possidente, Mingjian He\n\n  __________________________________________________________________________________________________________________\n    \"\"\"\n\n    #  Process user input\n    [data, fs, frequency_range, time_bandwidth, num_tapers,\n     winsize_samples, winstep_samples, window_start,\n     num_windows, nfft, detrend_opt, plot_on, verbose] = process_input(data, fs, frequency_range, time_bandwidth,\n                                                                       num_tapers, window_params, min_nfft,\n                                                                       detrend_opt, plot_on, verbose)\n\n    # Set up spectrogram parameters\n    [window_idxs, stimes, sfreqs, freq_inds] = process_spectrogram_params(fs, nfft, frequency_range, window_start,\n                                                                          winsize_samples)\n    # Display spectrogram parameters\n    if verbose:\n        display_spectrogram_props(fs, time_bandwidth, num_tapers, [winsize_samples, winstep_samples], frequency_range,\n                                  nfft, detrend_opt)\n\n    # Split data into segments and preallocate\n    data_segments = data[window_idxs]\n\n    # COMPUTE THE MULTITAPER SPECTROGRAM\n    #     STEP 1: Compute DPSS tapers based on desired spectral properties\n    #     STEP 2: Multiply the data segment by the DPSS Tapers\n    #     STEP 3: Compute the spectrum for each tapered segment\n    #     STEP 4: Take the mean of the tapered spectra\n\n    # Compute DPSS tapers (STEP 1)\n    dpss_tapers, dpss_eigen = dpss(winsize_samples, time_bandwidth, num_tapers, return_ratios=True)\n    dpss_eigen = np.reshape(dpss_eigen, (num_tapers, 1))\n\n    # pre-compute weights\n    if weighting == 'eigen':\n        wt = dpss_eigen / num_tapers\n    elif weighting == 'unity':\n        wt = np.ones(num_tapers) / num_tapers\n        wt = np.reshape(wt, (num_tapers, 1))  # reshape as column vector\n    else:\n        wt = 0\n\n    tic = timeit.default_timer()  # start timer\n\n    # Set up calc_mts_segment() input arguments\n    mts_params = (dpss_tapers, nfft, freq_inds, detrend_opt, num_tapers, dpss_eigen, weighting, wt)\n\n    if multiprocess:  # use multiprocessing\n        n_jobs = max(cpu_count() - 1, 1) if n_jobs is None else n_jobs\n        mt_spectrogram = np.vstack(Parallel(n_jobs=n_jobs)(delayed(calc_mts_segment)(\n            data_segments[num_window, :], *mts_params) for num_window in range(num_windows)))\n\n    else:  # if no multiprocessing, compute normally\n        mt_spectrogram = np.apply_along_axis(calc_mts_segment, 1, data_segments, *mts_params)\n\n    # Compute one-sided PSD spectrum\n    mt_spectrogram = mt_spectrogram.T\n    dc_select = np.where(sfreqs == 0)[0]\n    nyquist_select = np.where(sfreqs == fs/2)[0]\n    select = np.setdiff1d(np.arange(0, len(sfreqs)), np.concatenate((dc_select, nyquist_select)))\n\n    mt_spectrogram = np.vstack([mt_spectrogram[dc_select, :], 2*mt_spectrogram[select, :],\n                               mt_spectrogram[nyquist_select, :]]) / fs\n\n    # Flip if requested\n    if xyflip:\n        mt_spectrogram = mt_spectrogram.T\n\n    # End timer and get elapsed compute time\n    toc = timeit.default_timer()\n    if verbose:\n        print(\"\\n Multitaper compute time: \" + \"%.2f\" % (toc - tic) + \" seconds\")\n\n    if np.all(mt_spectrogram.flatten() == 0):\n        print(\"\\n Data was all zeros, no output\")\n\n    # Plot multitaper spectrogram\n    if plot_on:\n        # convert from power to dB\n        spect_data = nanpow2db(mt_spectrogram)\n\n        # Set x and y axes\n        dx = stimes[1] - stimes[0]\n        dy = sfreqs[1] - sfreqs[0]\n        extent = [stimes[0]-dx, stimes[-1]+dx, sfreqs[-1]+dy, sfreqs[0]-dy]\n\n        # Plot spectrogram\n        if ax is None:\n            fig, ax = plt.subplots()\n        else:\n            fig = ax.get_figure()\n        im = ax.imshow(spect_data, extent=extent, aspect='auto')\n        fig.colorbar(im, ax=ax, label='PSD (dB)', shrink=0.8)\n        ax.set_xlabel(\"Time (HH:MM:SS)\")\n        ax.set_ylabel(\"Frequency (Hz)\")\n        im.set_cmap(plt.cm.get_cmap('cet_rainbow4'))\n        ax.invert_yaxis()\n\n        # Scale colormap\n        if clim_scale:\n            clim = np.percentile(spect_data, [5, 98])  # from 5th percentile to 98th\n            im.set_clim(clim)  # actually change colorbar scale\n\n        fig.show()\n        if return_fig:\n            return mt_spectrogram, stimes, sfreqs, (fig, ax)\n\n    return mt_spectrogram, stimes, sfreqs\n\n\n# Helper Functions #\n\n# Process User Inputs #\ndef process_input(data, fs, frequency_range=None, time_bandwidth=5, num_tapers=None, window_params=None, min_nfft=0,\n                  detrend_opt='linear', plot_on=True, verbose=True):\n    \"\"\" Helper function to process multitaper_spectrogram() arguments\n            Arguments:\n                    data (1d np.array): time series data-- required\n                    fs (float): sampling frequency in Hz  -- required\n                    frequency_range (list): 1x2 list - [<min frequency>, <max frequency>] (default: [0 nyquist])\n                    time_bandwidth (float): time-half bandwidth product (window duration*half bandwidth of main lobe)\n                                            (default: 5 Hz*s)\n                    num_tapers (int): number of DPSS tapers to use (default: None [will be computed\n                                      as floor(2*time_bandwidth - 1)])\n                    window_params (list): 1x2 list - [window size (seconds), step size (seconds)] (default: [5 1])\n                    min_nfft (int): minimum allowable NFFT size, adds zero padding for interpolation (closest 2^x)\n                                    (default: 0)\n                    detrend_opt (string): detrend data window ('linear' (default), 'constant', 'off')\n                                          (Default: 'linear')\n                    plot_on (True): plot results (default: True)\n                    verbose (True): display spectrogram properties (default: true)\n            Returns:\n                    data (1d np.array): same as input\n                    fs (float): same as input\n                    frequency_range (list): same as input or calculated from fs if not given\n                    time_bandwidth (float): same as input or default if not given\n                    num_tapers (int): same as input or calculated from time_bandwidth if not given\n                    winsize_samples (int): number of samples in single time window\n                    winstep_samples (int): number of samples in a single window step\n                    window_start (1xm np.array): array of timestamps representing the beginning time for each window\n                    num_windows (int): number of windows in the data\n                    nfft (int): length of signal to calculate fft on\n                    detrend_opt ('string'): same as input or default if not given\n                    plot_on (bool): same as input\n                    verbose (bool): same as input\n    \"\"\"\n\n    # Make sure data is 1 dimensional np array\n    if len(data.shape) != 1:\n        if (len(data.shape) == 2) & (data.shape[1] == 1):  # if it's 2d, but can be transferred to 1d, do so\n            data = np.ravel(data[:, 0])\n        elif (len(data.shape) == 2) & (data.shape[0] == 1):  # if it's 2d, but can be transferred to 1d, do so\n            data = np.ravel(data.T[:, 0])\n        else:\n            raise TypeError(\"Input data is the incorrect dimensions. Should be a 1d array with shape (n,) where n is \\\n                            the number of data points. Instead data shape was \" + str(data.shape))\n\n    # Set frequency range if not provided\n    if frequency_range is None:\n        frequency_range = [0, fs / 2]\n\n    # Set detrending method\n    detrend_opt = detrend_opt.lower()\n    if detrend_opt != 'linear':\n        if detrend_opt in ['const', 'constant']:\n            detrend_opt = 'constant'\n        elif detrend_opt in ['none', 'false', 'off']:\n            detrend_opt = 'off'\n        else:\n            raise ValueError(\"'\" + str(detrend_opt) + \"' is not a valid argument for detrend_opt. The choices \" +\n                             \"are: 'constant', 'linear', or 'off'.\")\n    # Check if frequency range is valid\n    if frequency_range[1] > fs / 2:\n        frequency_range[1] = fs / 2\n        warnings.warn('Upper frequency range greater than Nyquist, setting range to [' +\n                      str(frequency_range[0]) + ', ' + str(frequency_range[1]) + ']')\n\n    # Set number of tapers if none provided\n    if num_tapers is None:\n        num_tapers = math.floor(2 * time_bandwidth) - 1\n\n    # Warn if number of tapers is suboptimal\n    if num_tapers != math.floor(2 * time_bandwidth) - 1:\n        warnings.warn('Number of tapers is optimal at floor(2*TW) - 1. consider using ' +\n                      str(math.floor(2 * time_bandwidth) - 1))\n\n    # If no window params provided, set to defaults\n    if window_params is None:\n        window_params = [5, 1]\n\n    # Check if window size is valid, fix if not\n    if window_params[0] * fs % 1 != 0:\n        winsize_samples = round(window_params[0] * fs)\n        warnings.warn('Window size is not divisible by sampling frequency. Adjusting window size to ' +\n                      str(winsize_samples / fs) + ' seconds')\n    else:\n        winsize_samples = window_params[0] * fs\n\n    # Check if window step is valid, fix if not\n    if window_params[1] * fs % 1 != 0:\n        winstep_samples = round(window_params[1] * fs)\n        warnings.warn('Window step size is not divisible by sampling frequency. Adjusting window step size to ' +\n                      str(winstep_samples / fs) + ' seconds')\n    else:\n        winstep_samples = window_params[1] * fs\n\n    # Get total data length\n    len_data = len(data)\n\n    # Check if length of data is smaller than window (bad)\n    if len_data < winsize_samples:\n        raise ValueError(\"\\nData length (\" + str(len_data) + \") is shorter than window size (\" +\n                         str(winsize_samples) + \"). Either increase data length or decrease window size.\")\n\n    # Find window start indices and num of windows\n    window_start = np.arange(0, len_data - winsize_samples + 1, winstep_samples)\n    num_windows = len(window_start)\n\n    # Get num points in FFT\n    if min_nfft == 0:  # avoid divide by zero error in np.log2(0)\n        nfft = max(2 ** math.ceil(np.log2(abs(winsize_samples))), winsize_samples)\n    else:\n        nfft = max(max(2 ** math.ceil(np.log2(abs(winsize_samples))), winsize_samples),\n                   2 ** math.ceil(np.log2(abs(min_nfft))))\n\n    return ([data, fs, frequency_range, time_bandwidth, num_tapers,\n             int(winsize_samples), int(winstep_samples), window_start, num_windows, nfft,\n             detrend_opt, plot_on, verbose])\n\n\n# PROCESS THE SPECTROGRAM PARAMETERS #\ndef process_spectrogram_params(fs, nfft, frequency_range, window_start, datawin_size):\n    \"\"\" Helper function to create frequency vector and window indices\n        Arguments:\n             fs (float): sampling frequency in Hz  -- required\n             nfft (int): length of signal to calculate fft on -- required\n             frequency_range (list): 1x2 list - [<min frequency>, <max frequency>] -- required\n             window_start (1xm np array): array of timestamps representing the beginning time for each\n                                          window -- required\n             datawin_size (float): seconds in one window -- required\n        Returns:\n            window_idxs (nxm np array): indices of timestamps for each window\n                                        (nxm where n=number of windows and m=datawin_size)\n            stimes (1xt np array): array of times for the center of the spectral bins\n            sfreqs (1xf np array): array of frequency bins for the spectrogram\n            freq_inds (1d np array): boolean array of which frequencies are being analyzed in\n                                      an array of frequencies from 0 to fs with steps of fs/nfft\n    \"\"\"\n\n    # create frequency vector\n    df = fs / nfft\n    sfreqs = np.arange(0, fs, df)\n\n    # Get frequencies for given frequency range\n    freq_inds = (sfreqs >= frequency_range[0]) & (sfreqs <= frequency_range[1])\n    sfreqs = sfreqs[freq_inds]\n\n    # Compute times in the middle of each spectrum\n    window_middle_samples = window_start + round(datawin_size / 2)\n    stimes = window_middle_samples / fs\n\n    # Get indexes for each window\n    window_idxs = np.atleast_2d(window_start).T + np.arange(0, datawin_size, 1)\n    window_idxs = window_idxs.astype(int)\n\n    return [window_idxs, stimes, sfreqs, freq_inds]\n\n\n# DISPLAY SPECTROGRAM PROPERTIES\ndef display_spectrogram_props(fs, time_bandwidth, num_tapers, data_window_params, frequency_range, nfft, detrend_opt):\n    \"\"\" Prints spectrogram properties\n        Arguments:\n            fs (float): sampling frequency in Hz  -- required\n            time_bandwidth (float): time-half bandwidth product (window duration*1/2*frequency_resolution) -- required\n            num_tapers (int): number of DPSS tapers to use -- required\n            data_window_params (list): 1x2 list - [window length(s), window step size(s)] -- required\n            frequency_range (list): 1x2 list - [<min frequency>, <max frequency>] -- required\n            nfft(float): number of fast fourier transform samples -- required\n            detrend_opt (str): detrend data window ('linear' (default), 'constant', 'off') -- required\n        Returns:\n            This function does not return anything\n    \"\"\"\n\n    data_window_params = np.asarray(data_window_params) / fs\n\n    # Print spectrogram properties\n    print(\"Multitaper Spectrogram Properties: \")\n    print('     Spectral Resolution: ' + str(2 * time_bandwidth / data_window_params[0]) + 'Hz')\n    print('     Window Length: ' + str(data_window_params[0]) + 's')\n    print('     Window Step: ' + str(data_window_params[1]) + 's')\n    print('     Time Half-Bandwidth Product: ' + str(time_bandwidth))\n    print('     Number of Tapers: ' + str(num_tapers))\n    print('     Frequency Range: ' + str(frequency_range[0]) + \"-\" + str(frequency_range[1]) + 'Hz')\n    print('     NFFT: ' + str(nfft))\n    print('     Detrend: ' + detrend_opt + '\\n')\n\n\n# NANPOW2DB\ndef nanpow2db(y):\n    \"\"\" Power to dB conversion, setting bad values to nans\n        Arguments:\n            y (float or array-like): power\n        Returns:\n            ydB (float or np array): inputs converted to dB with 0s and negatives resulting in nans\n    \"\"\"\n\n    if isinstance(y, int) or isinstance(y, float):\n        if y == 0:\n            return np.nan\n        else:\n            ydB = 10 * np.log10(y)\n    else:\n        if isinstance(y, list):  # if list, turn into array\n            y = np.asarray(y)\n        y = y.astype(float)  # make sure it's a float array so we can put nans in it\n        y[y == 0] = np.nan\n        ydB = 10 * np.log10(y)\n\n    return ydB\n\n\n# Helper #\ndef is_outlier(data):\n    smad = 1.4826 * np.median(abs(data - np.median(data)))  # scaled median absolute deviation\n    outlier_mask = abs(data-np.median(data)) > 3*smad  # outliers are more than 3 smads away from median\n    outlier_mask = (outlier_mask | np.isnan(data) | np.isinf(data))\n    return outlier_mask\n\n\n# CALCULATE MULTITAPER SPECTRUM ON SINGLE SEGMENT\ndef calc_mts_segment(data_segment, dpss_tapers, nfft, freq_inds, detrend_opt, num_tapers, dpss_eigen, weighting, wt):\n    \"\"\" Helper function to calculate the multitaper spectrum of a single segment of data\n        Arguments:\n            data_segment (1d np.array): One window worth of time-series data -- required\n            dpss_tapers (2d np.array): Parameters for the DPSS tapers to be used.\n                                       Dimensions are (num_tapers, winsize_samples) -- required\n            nfft (int): length of signal to calculate fft on -- required\n            freq_inds (1d np array): boolean array of which frequencies are being analyzed in\n                                      an array of frequencies from 0 to fs with steps of fs/nfft\n            detrend_opt (str): detrend data window ('linear' (default), 'constant', 'off')\n            num_tapers (int): number of tapers being used\n            dpss_eigen (np array):\n            weighting (str):\n            wt (int or np array):\n        Returns:\n            mt_spectrum (1d np.array): spectral power for single window\n    \"\"\"\n\n    # If segment has all zeros, return vector of zeros\n    if all(data_segment == 0):\n        ret = np.empty(sum(freq_inds))\n        ret.fill(0)\n        return ret\n\n    if any(np.isnan(data_segment)):\n        ret = np.empty(sum(freq_inds))\n        ret.fill(np.nan)\n        return ret\n\n    # Option to detrend data to remove low frequency DC component\n    if detrend_opt != 'off':\n        data_segment = detrend(data_segment, type=detrend_opt)\n\n    # Multiply data by dpss tapers (STEP 2)\n    tapered_data = np.multiply(np.mat(data_segment).T, np.mat(dpss_tapers.T))\n\n    # Compute the FFT (STEP 3)\n    fft_data = np.fft.fft(tapered_data, nfft, axis=0)\n\n    # Compute the weighted mean spectral power across tapers (STEP 4)\n    spower = np.power(np.imag(fft_data), 2) + np.power(np.real(fft_data), 2)\n    if weighting == 'adapt':\n        # adaptive weights - for colored noise spectrum (Percival & Walden p368-370)\n        tpower = np.dot(np.transpose(data_segment), (data_segment/len(data_segment)))\n        spower_iter = np.mean(spower[:, 0:2], 1)\n        spower_iter = spower_iter[:, np.newaxis]\n        a = (1 - dpss_eigen) * tpower\n        for i in range(3):  # 3 iterations only\n            # Calc the MSE weights\n            b = np.dot(spower_iter, np.ones((1, num_tapers))) / ((np.dot(spower_iter, np.transpose(dpss_eigen))) +\n                                                                 (np.ones((nfft, 1)) * np.transpose(a)))\n            # Calc new spectral estimate\n            wk = (b**2) * np.dot(np.ones((nfft, 1)), np.transpose(dpss_eigen))\n            spower_iter = np.sum((np.transpose(wk) * np.transpose(spower)), 0) / np.sum(wk, 1)\n            spower_iter = spower_iter[:, np.newaxis]\n\n        mt_spectrum = np.squeeze(spower_iter)\n\n    else:\n        # eigenvalue or uniform weights\n        mt_spectrum = np.dot(spower, wt)\n        mt_spectrum = np.reshape(mt_spectrum, nfft)  # reshape to 1D\n\n    return mt_spectrum[freq_inds]\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-04-05T08:07:43.337501Z","iopub.execute_input":"2024-04-05T08:07:43.337882Z","iopub.status.idle":"2024-04-05T08:07:43.401208Z","shell.execute_reply.started":"2024-04-05T08:07:43.337852Z","shell.execute_reply":"2024-04-05T08:07:43.400303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_eeg_differences(eeg):\n    \n    \"\"\"\n    Create bipolar channel difference features\n\n    Parameters\n    ----------\n    eeg: numpy.ndarray of shape (time, 20)\n         EEG array\n\n    Returns\n    -------\n    eeg_difference: numpy.ndarray of shape (time, 18)\n        EEG differences array\n    \"\"\"\n\n    eeg_difference = np.zeros((eeg.shape[0], 18))\n\n    # Left outside temporal chain\n    eeg_difference[:, 0] = eeg[:, 0] - eeg[:, 4]\n    eeg_difference[:, 1] = eeg[:, 4] - eeg[:, 5]\n    eeg_difference[:, 2] = eeg[:, 5] - eeg[:, 6]\n    eeg_difference[:, 3] = eeg[:, 6] - eeg[:, 7]\n\n    # Right outside temporal chain\n    eeg_difference[:, 4] = eeg[:, 11] - eeg[:, 15]\n    eeg_difference[:, 5] = eeg[:, 15] - eeg[:, 16]\n    eeg_difference[:, 6] = eeg[:, 16] - eeg[:, 17]\n    eeg_difference[:, 7] = eeg[:, 17] - eeg[:, 18]\n\n    # Left inside parasagittal chain\n    eeg_difference[:, 8] = eeg[:, 0] - eeg[:, 1]\n    eeg_difference[:, 9] = eeg[:, 1] - eeg[:, 2]\n    eeg_difference[:, 10] = eeg[:, 2] - eeg[:, 3]\n    eeg_difference[:, 11] = eeg[:, 3] - eeg[:, 7]\n\n    # Right inside parasagittal chain\n    eeg_difference[:, 12] = eeg[:, 11] - eeg[:, 12]\n    eeg_difference[:, 13] = eeg[:, 12] - eeg[:, 13]\n    eeg_difference[:, 14] = eeg[:, 13] - eeg[:, 14]\n    eeg_difference[:, 15] = eeg[:, 14] - eeg[:, 18]\n\n    # Center chain\n    eeg_difference[:, 16] = eeg[:, 8] - eeg[:, 9]\n    eeg_difference[:, 17] = eeg[:, 9] - eeg[:, 10]\n\n    return eeg_difference\n\n\ndef get_spectrogram(signal, fs, nperseg, noverlap, frequency_range):\n    \n    \"\"\"\n    Create spectrogram from the given signal\n\n    Parameters\n    ----------\n    signal: numpy.ndarray of shape (time)\n        Signal array\n        \n    fs: int\n        Sampling frequency\n    \n    nperseg: int\n        Length of each segment\n        \n    noverlap: int\n        Number of points to overlap between segments (nperseg - noverlap = stride)\n        \n    frequency_range: tuple of shape (2)\n        Lower and upper bound of frequencies to keep\n        \n    Returns\n    -------\n    spectrogram: numpy.ndarray of shape (frequency, time)\n        Spectrogram array\n    \"\"\"\n\n    frequencies, _, spectrogram = cusignal.spectrogram(\n        signal,\n        fs=fs,\n        nperseg=nperseg,\n        noverlap=noverlap,\n        nfft=None\n    )\n    frequency_mask = (frequencies >= frequency_range[0]) & (frequencies <= frequency_range[1])\n    spectrogram = spectrogram[frequency_mask, :]\n    spectrogram = spectrogram.get().astype(np.float32)\n\n    return spectrogram\n\n\ndef get_spectrogram_50_second(eeg_differences):\n        \n    \"\"\"\n    Create 50 second spectrograms from given EEG differences and stack them on vertical axis\n\n    Parameters\n    ----------\n    eeg_differences: numpy.ndarray of shape (10000, 18)\n        EEG differences array\n        \n    Returns\n    -------\n    spectrogram_50_second: numpy.ndarray of shape (504, 487)\n        Spectrogram 50 second array\n    \"\"\"\n    \n    spectrogram_50_second = []\n    \n    for signal_idx in range(eeg_differences.shape[1]):\n        \n        # Create spectrogram from the entire EEG signal and take frequencies between 0.5 and 20\n        spectrogram = get_spectrogram(\n            signal=eeg_differences[:, signal_idx],\n            fs=200,\n            nperseg=280,\n            noverlap=261,\n            frequency_range=(0.5, 20)\n        )\n        spectrogram_50_second.append(spectrogram)\n        \n    # Concatenate spectrograms on vertical axis and apply log transform\n    spectrogram_50_second = np.concatenate(spectrogram_50_second, axis=0)\n    spectrogram_50_second = np.log1p(spectrogram_50_second)\n    \n    return spectrogram_50_second\n\n\ndef get_spectrogram_50_10_second(eeg_differences):\n    \n    \"\"\"\n    Create 50/10 second spectrograms from given EEG differences and stack them on vertical axis\n\n    Parameters\n    ----------\n    eeg_differences: numpy.ndarray of shape (10000, 18)\n        EEG differences array\n        \n    Returns\n    -------\n    spectrogram_50_10_second: numpy.ndarray of shape (504, 464)\n        Spectrogram 50/10 second array\n    \"\"\"\n    \n    spectrogram_50_10_second = []\n    \n    for signal_idx in range(eeg_differences.shape[1]):\n        \n        # Create spectrogram from the entire EEG signal and take frequencies between 0.5 and 20\n        spectrogram_50_second = get_spectrogram(\n            signal=eeg_differences[:, signal_idx],\n            fs=200,\n            nperseg=145,\n            noverlap=124,\n            frequency_range=(0.5, 20)\n        )[:, 3:-3]\n        spectrogram_50_10_second.append(spectrogram_50_second)\n                \n        # Create spectrogram from the center 10 seconds of EEG signal and take frequencies between 0.5 and 20\n        spectrogram_10_second = get_spectrogram(\n            signal=eeg_differences[4000:6000, signal_idx],\n            fs=200,\n            nperseg=145,\n            noverlap=141,\n            frequency_range=(0.5, 20)\n        )\n        spectrogram_50_10_second.append(spectrogram_10_second)\n        \n    # Concatenate spectrogram on vertical axis and do log transform\n    spectrogram_50_10_second = np.concatenate(spectrogram_50_10_second, axis=0)\n    spectrogram_50_10_second = np.log1p(spectrogram_50_10_second)\n    \n    return spectrogram_50_10_second\n\n\ndef get_spectrogram_50_30_10_second(eeg_differences):\n    \n    \"\"\"\n    Create 50/30/10 second spectrograms from given EEG differences and stack them on vertical axis\n\n    Parameters\n    ----------\n    eeg_differences: numpy.ndarray of shape (10000, 18)\n        EEG differences array\n        \n    Returns\n    -------\n    spectrogram_50_30_10_second: numpy.ndarray of shape (504, 476)\n        Spectrogram 50/30/10 second array\n    \"\"\"\n    \n    spectrogram_50_30_10_second = []\n    \n    for signal_idx in range(eeg_differences.shape[1]):\n        \n        # Create spectrogram from the entire EEG signal and take frequencies between 0.5 and 20\n        spectrogram_50_second = get_spectrogram(\n            signal=eeg_differences[:, signal_idx],\n            fs=200,\n            nperseg=99,\n            noverlap=79,\n            frequency_range=(0.5, 20)\n        )[:, 10:-10]\n        spectrogram_50_30_10_second.append(spectrogram_50_second)\n        \n        # Create spectrogram from the center 30 seconds of EEG signal and take frequencies between 0.5 and 20\n        spectrogram_30_second = get_spectrogram(\n            signal=eeg_differences[2000:8000, signal_idx],\n            fs=200,\n            nperseg=99,\n            noverlap=87,\n            frequency_range=(0.5, 20)\n        )[:, 8:-8]\n        spectrogram_50_30_10_second.append(spectrogram_30_second)\n                \n        # Create spectrogram from the center 10 seconds of EEG signal and take frequencies between 0.5 and 20\n        spectrogram_10_second = get_spectrogram(\n            signal=eeg_differences[4000:6000, signal_idx],\n            fs=200,\n            nperseg=100,\n            noverlap=96,\n            frequency_range=(0.5, 20)\n        )\n        spectrogram_50_30_10_second.append(spectrogram_10_second)\n        \n    # Concatenate spectrogram on vertical axis and do log transform\n    spectrogram_50_30_10_second = np.concatenate(spectrogram_50_30_10_second, axis=0)\n    spectrogram_50_30_10_second = np.log1p(spectrogram_50_30_10_second)\n    \n    return spectrogram_50_30_10_second\n\n\ndef get_spectrogram_30_10_second(eeg_differences):\n    \n    \"\"\"\n    Create 30/10 second spectrograms from given EEG differences and stack them on vertical axis\n\n    Parameters\n    ----------\n    eeg_differences: numpy.ndarray of shape (10000, 18)\n        EEG differences array\n        \n    Returns\n    -------\n    spectrogram_30_10_second: numpy.ndarray of shape (504, 464)\n        Spectrogram 30/10 second array\n    \"\"\"\n    \n    spectrogram_30_10_second = []\n    \n    for signal_idx in range(eeg_differences.shape[1]):\n        \n        # Create spectrogram from the center 30 seconds and take frequencies between 0.5 and 20\n        spectrogram_30_second = get_spectrogram(\n            signal=eeg_differences[2000:8000, signal_idx],\n            fs=200,\n            nperseg=145,\n            noverlap=133,\n            frequency_range=(0.5, 20)\n        )[:, 12:-12]\n        spectrogram_30_10_second.append(spectrogram_30_second)\n\n        # Create spectrogram from the center 10 seconds of EEG signal and take frequencies between 0.5 and 20\n        spectrogram_10_second = get_spectrogram(\n            signal=eeg_differences[4000:6000, signal_idx],\n            fs=200,\n            nperseg=145,\n            noverlap=141,\n            frequency_range=(0.5, 20)\n        )\n        spectrogram_30_10_second.append(spectrogram_10_second)\n\n    # Concatenate spectrogram on vertical axis and do log transform\n    spectrogram_30_10_second = np.concatenate(spectrogram_30_10_second, axis=0)\n    spectrogram_30_10_second = np.log1p(spectrogram_30_10_second)\n\n    return spectrogram_30_10_second\n\n\ndef get_multitaper_spectrograms(eeg_differences):\n    \n    \"\"\"\n    Create 50 and 10 second multitaper spectrograms from given EEG differences and stack them on depth axis\n\n    Parameters\n    ----------\n    eeg_differences: numpy.ndarray of shape (10000, 18)\n        EEG differences array\n        \n    Returns\n    -------\n    spectrograms: numpy.ndarray of shape (500, 372)\n        50 second multitaper spectrogram\n        \n    spectrograms_center: numpy.ndarray of shape (500, 424)\n        10 second multitaper spectrogram\n    \"\"\"\n            \n    spectrograms, spectrograms_center = [], []\n\n    for signal_idx in range(eeg_differences.shape[1]):\n        spect, stimes, sfreqs = multitaper_spectrogram(\n            eeg_differences[:, signal_idx],\n            fs=200,\n            frequency_range=[0.5, 20],\n            window_params=[4, 0.5],\n            multiprocess=True,\n            plot_on=False,\n            verbose=False,\n        )\n        spectrograms.append(spect)\n\n        spect_center, stimes_center, sfreqs_center = multitaper_spectrogram(\n            eeg_differences[4000:6000, signal_idx],\n            fs=200,\n            frequency_range=[0.5, 20],\n            window_params=[4, 0.5],\n            multiprocess=True,\n            plot_on=False,\n            verbose=False,\n        )\n        spectrograms_center.append(spect_center)\n\n    spectrograms = np.stack(spectrograms, axis=0)\n    spectrograms_center = np.stack(spectrograms_center, axis=0)\n    \n    spectrograms = np.log1p(spectrograms)\n    spectrograms_center = np.log1p(spectrograms_center)\n    \n    return spectrograms, spectrograms_center\n","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:13:01.130984Z","iopub.execute_input":"2024-04-05T08:13:01.131383Z","iopub.status.idle":"2024-04-05T08:13:01.165577Z","shell.execute_reply.started":"2024-04-05T08:13:01.131354Z","shell.execute_reply":"2024-04-05T08:13:01.164616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Global configurations\ndevice = torch.device('cuda')\namp = True\ntta = True\n\n# Raw EEG 50 second models, configurations and transforms\nraw_eeg_50_second_2d_model_configs = [\n    raw_eeg_50_second_2d_convnext_models_config,\n    raw_eeg_50_second_2d_maxvit_models_config\n]\nraw_eeg_50_second_2d_models = [\n    raw_eeg_50_second_2d_convnext_models,\n    raw_eeg_50_second_2d_maxvit_models\n]\nraw_eeg_50_second_2d_transforms = [\n    get_raw_eeg_2d_transforms(**config['transforms'])['inference']\n    for config in raw_eeg_50_second_2d_model_configs\n]\n\n# Raw EEG 50/10 second models, configurations and transforms\nraw_eeg_50_10_second_2d_model_configs = [\n    raw_eeg_50_10_second_2d_convnext_models_config,\n    raw_eeg_50_10_second_2d_maxvit_models_config\n]\nraw_eeg_50_10_second_2d_models = [\n    raw_eeg_50_10_second_2d_convnext_models,\n    raw_eeg_50_10_second_2d_maxvit_models\n]\nraw_eeg_50_10_second_2d_transforms = [\n    get_raw_eeg_2d_transforms(**config['transforms'])['inference']\n    for config in raw_eeg_50_10_second_2d_model_configs\n]\n\n# Spectrogram 50 second models, configurations and transforms\nspectrogram_50_second_2d_model_configs = [\n    spectrogram_50_second_2d_convnext_models_config,\n    spectrogram_50_second_2d_maxvit_models_config\n]\nspectrogram_50_second_2d_models = [\n    spectrogram_50_second_2d_convnext_models,\n    spectrogram_50_second_2d_maxvit_models,\n]\nspectrogram_50_second_2d_transforms = [\n    get_spectrogram_2d_transforms(**config['transforms'])['inference']\n    for config in spectrogram_50_second_2d_model_configs\n]\n\n# Spectrogram 50/10 second models, configurations and transforms\nspectrogram_50_10_second_2d_model_configs = [\n    spectrogram_50_10_second_2d_convnext_models_config,\n    spectrogram_50_10_second_2d_maxvit_models_config\n]\nspectrogram_50_10_second_2d_models = [\n    spectrogram_50_10_second_2d_convnext_models,\n    spectrogram_50_10_second_2d_maxvit_models\n]\nspectrogram_50_10_second_2d_transforms = [\n    get_spectrogram_2d_transforms(**config['transforms'])['inference']\n    for config in spectrogram_50_10_second_2d_model_configs\n]\n\n# Spectrogram 50/30/10 second models, configurations and transforms\nspectrogram_50_30_10_second_2d_model_configs = [\n    spectrogram_50_30_10_second_2d_convnext_models_config,\n    spectrogram_50_30_10_second_2d_maxvit_models_config\n]\nspectrogram_50_30_10_second_2d_models = [\n    spectrogram_50_30_10_second_2d_convnext_models,\n    spectrogram_50_30_10_second_2d_maxvit_models\n]\nspectrogram_50_30_10_second_2d_transforms = [\n    get_spectrogram_2d_transforms(**config['transforms'])['inference']\n    for config in spectrogram_50_30_10_second_2d_model_configs\n]\n\n# Spectrogram 30/10 second models, configurations and transforms\nspectrogram_30_10_second_2d_model_configs = [\n    spectrogram_30_10_second_2d_convnext_models_config,\n    spectrogram_30_10_second_2d_maxvit_models_config\n]\nspectrogram_30_10_second_2d_models = [\n    spectrogram_30_10_second_2d_convnext_models,\n    spectrogram_30_10_second_2d_maxvit_models\n]\nspectrogram_30_10_second_2d_transforms = [\n    get_spectrogram_2d_transforms(**config['transforms'])['inference']\n    for config in spectrogram_30_10_second_2d_model_configs\n]\n\n# Other transforms\nget_raw_eeg_50_second_2d_image = lambda x: EEGTo2D(ekg=False, center_stack=False, always_apply=True)(image=get_eeg_differences(x))['image']\nget_raw_eeg_50_10_second_2d_image = lambda x: EEGTo2D(ekg=False, center_stack=True, always_apply=True)(image=get_eeg_differences(x))['image']\nget_multitaper_spectrogram_50_second_2d_image = lambda x: EEG3Dto2D(always_apply=True)(image=x)['image']","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:13:05.357941Z","iopub.execute_input":"2024-04-05T08:13:05.358320Z","iopub.status.idle":"2024-04-05T08:13:05.378743Z","shell.execute_reply.started":"2024-04-05T08:13:05.358290Z","shell.execute_reply":"2024-04-05T08:13:05.377656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_raw_eeg_50_second_2d_models = 2\nn_raw_eeg_50_10_second_2d_models = 2\nn_spectrogram_50_second_2d_models = 2\nn_spectrogram_50_10_second_2d_models = 2\nn_spectrogram_50_30_10_second_2d_models = 2\nn_spectrogram_30_10_second_2d_models = 2\n\ntest_raw_eeg_50_second_2d_predictions = []\ntest_raw_eeg_50_10_second_2d_predictions = []\ntest_spectrogram_50_second_2d_predictions = []\ntest_spectrogram_50_10_second_2d_predictions = []\ntest_spectrogram_50_30_10_second_2d_predictions = []\ntest_spectrogram_30_10_second_2d_predictions = []\n\nfor idx, row in df.iterrows():\n    \n    # Read raw EEG, interpolate center and fill edges with zeros\n    eeg_path = eeg_directory / f'{int(row[\"eeg_id\"])}.parquet'\n    eeg = pd.read_parquet(eeg_path)\n    eeg = eeg.interpolate(method='linear', limit_area='inside').fillna(0).values\n    \n    if is_submission is False:\n        raw_eeg_50_second_2d_image = get_raw_eeg_50_second_2d_image(eeg)\n        print(f'\\nRaw EEG 50 Second Image Shape {raw_eeg_50_second_2d_image.shape}')\n        visualize_image(raw_eeg_50_second_2d_image, title='Raw EEG 50 Second 2D Image')\n        del raw_eeg_50_second_2d_image\n\n    # Create raw EEG 50 second predictions tensors for the current sample\n    raw_eeg_50_second_2d_predictions = torch.zeros(n_raw_eeg_50_second_2d_models, 1, 6)\n\n    for model_idx, (models, transforms) in enumerate(zip(raw_eeg_50_second_2d_models, raw_eeg_50_second_2d_transforms)):\n        \n        #continue\n                                                                                                        \n        raw_eeg_50_second_2d_inputs = transforms(image=eeg)['image'].float()\n        raw_eeg_50_second_2d_inputs = raw_eeg_50_second_2d_inputs.to(device)\n\n        if tta:\n            raw_eeg_50_second_2d_inputs = torch.stack((\n                raw_eeg_50_second_2d_inputs,\n                torch.flip(raw_eeg_50_second_2d_inputs, dims=(1,)),\n                torch.flip(raw_eeg_50_second_2d_inputs, dims=(2,)),\n                torch.flip(raw_eeg_50_second_2d_inputs, dims=(1, 2))\n            ), dim=0)\n        else:\n            raw_eeg_50_second_2d_inputs = torch.unsqueeze(raw_eeg_50_second_2d_inputs, dim=0)\n        \n        for fold, model in enumerate(models.values()):\n            with torch.no_grad():\n                if amp:\n                    with torch.autocast(device_type='cuda', dtype=torch.float16):\n                        raw_eeg_50_second_2d_outputs = model(raw_eeg_50_second_2d_inputs)\n                else:\n                    raw_eeg_50_second_2d_outputs = model(raw_eeg_50_second_2d_inputs)\n\n            raw_eeg_50_second_2d_outputs = raw_eeg_50_second_2d_outputs.cpu()\n            if tta:\n                raw_eeg_50_second_2d_outputs = torch.mean(raw_eeg_50_second_2d_outputs, dim=0)\n            else:\n                raw_eeg_50_second_2d_outputs = torch.squeeze(raw_eeg_50_second_2d_outputs, dim=0)\n            \n            raw_eeg_50_second_2d_predictions[model_idx] += raw_eeg_50_second_2d_outputs / len(models)\n            \n            if is_submission is False:\n                print(f'Predicted with Raw EEG 50 Second Model {model_idx + 1} Fold {fold + 1}')\n    \n    test_raw_eeg_50_second_2d_predictions.append(raw_eeg_50_second_2d_predictions)\n    \n    if is_submission is False:\n        raw_eeg_50_10_second_2d_image = get_raw_eeg_50_10_second_2d_image(eeg)\n        print(f'\\nRaw EEG 50/10 Second Image Shape {raw_eeg_50_10_second_2d_image.shape}')\n        visualize_image(raw_eeg_50_10_second_2d_image, title='Raw EEG 50/10 Second 2D Image')\n        del raw_eeg_50_10_second_2d_image\n\n    # Create raw EEG 50/10 second predictions tensors for the current sample\n    raw_eeg_50_10_second_2d_predictions = torch.zeros(n_raw_eeg_50_10_second_2d_models, 1, 6)\n\n    for model_idx, (models, transforms) in enumerate(zip(raw_eeg_50_10_second_2d_models, raw_eeg_50_10_second_2d_transforms)):\n        \n        #continue\n                                                                                                \n        raw_eeg_50_10_second_2d_inputs = transforms(image=eeg)['image'].float()\n        raw_eeg_50_10_second_2d_inputs = raw_eeg_50_10_second_2d_inputs.to(device)\n\n        if tta:\n            raw_eeg_50_10_second_2d_inputs = torch.stack((\n                raw_eeg_50_10_second_2d_inputs,\n                torch.flip(raw_eeg_50_10_second_2d_inputs, dims=(1,)),\n                torch.flip(raw_eeg_50_10_second_2d_inputs, dims=(2,)),\n                torch.flip(raw_eeg_50_10_second_2d_inputs, dims=(1, 2))\n            ), dim=0)\n        else:\n            raw_eeg_50_10_second_2d_inputs = torch.unsqueeze(raw_eeg_50_10_second_2d_inputs, dim=0)\n        \n        for fold, model in enumerate(models.values()):\n            with torch.no_grad():\n                if amp:\n                    with torch.autocast(device_type='cuda', dtype=torch.float16):\n                        raw_eeg_50_10_second_2d_outputs = model(raw_eeg_50_10_second_2d_inputs)\n                else:\n                    raw_eeg_50_10_second_2d_outputs = model(raw_eeg_50_10_second_2d_inputs)\n\n            raw_eeg_50_10_second_2d_outputs = raw_eeg_50_10_second_2d_outputs.cpu()\n            if tta:\n                raw_eeg_50_10_second_2d_outputs = torch.mean(raw_eeg_50_10_second_2d_outputs, dim=0)\n            else:\n                raw_eeg_50_10_second_2d_outputs = torch.squeeze(raw_eeg_50_10_second_2d_outputs, dim=0)\n            \n            raw_eeg_50_10_second_2d_predictions[model_idx] += raw_eeg_50_10_second_2d_outputs / len(models)\n            \n            if is_submission is False:\n                print(f'Predicted with Raw EEG 50/10 Second Model {model_idx + 1} Fold {fold + 1}')\n    \n    test_raw_eeg_50_10_second_2d_predictions.append(raw_eeg_50_10_second_2d_predictions)\n\n    # Create EEG differences array for spectrograms\n    eeg_differences = get_eeg_differences(eeg)\n    \n    # Create 50 second spectrogram from EEG differences\n    spectrogram_50_second = get_spectrogram_50_second(eeg_differences=eeg_differences)\n    \n    if is_submission is False:\n        print(f'\\nSpectrogram 50 Second Shape {spectrogram_50_second.shape}')\n        visualize_image(spectrogram_50_second, title='Spectrogram 50 Second')\n    \n    # Create 50 second spectrogram predictions tensors for the current sample\n    spectrogram_50_second_2d_predictions = torch.zeros(n_spectrogram_50_second_2d_models, 1, 6)\n    \n    for model_idx, (models, transforms) in enumerate(zip(spectrogram_50_second_2d_models, spectrogram_50_second_2d_transforms)):\n\n        #continue\n                                                                                        \n        spectrogram_50_second_2d_inputs = transforms(image=spectrogram_50_second)['image'].float()\n        spectrogram_50_second_2d_inputs = spectrogram_50_second_2d_inputs.to(device)\n\n        if tta:\n            spectrogram_50_second_2d_inputs = torch.stack((\n                spectrogram_50_second_2d_inputs,\n                torch.flip(spectrogram_50_second_2d_inputs, dims=(1,)),\n                torch.flip(spectrogram_50_second_2d_inputs, dims=(2,)),\n                torch.flip(spectrogram_50_second_2d_inputs, dims=(1, 2))\n            ), dim=0)\n        else:\n            spectrogram_50_second_2d_inputs = torch.unsqueeze(spectrogram_50_second_2d_inputs, dim=0)\n        \n        for fold, model in enumerate(models.values()):\n            with torch.no_grad():\n                if amp:\n                    with torch.autocast(device_type='cuda', dtype=torch.float16):\n                        spectrogram_50_second_2d_outputs = model(spectrogram_50_second_2d_inputs)\n                else:\n                    spectrogram_50_second_2d_outputs = model(spectrogram_50_second_2d_inputs)\n\n            spectrogram_50_second_2d_outputs = spectrogram_50_second_2d_outputs.cpu()\n            if tta:\n                spectrogram_50_second_2d_outputs = torch.mean(spectrogram_50_second_2d_outputs, dim=0)\n            else:\n                spectrogram_50_second_2d_outputs = torch.squeeze(spectrogram_50_second_2d_outputs, dim=0)\n            \n            spectrogram_50_second_2d_predictions[model_idx] += spectrogram_50_second_2d_outputs / len(models)\n            \n            if is_submission is False:\n                print(f'Predicted with Spectrogram 50 Second 2D Model {model_idx + 1} Fold {fold + 1}')\n    \n    test_spectrogram_50_second_2d_predictions.append(spectrogram_50_second_2d_predictions)\n        \n    # Create 50/10 second spectrogram from EEG differences\n    spectrogram_50_10_second = get_spectrogram_50_10_second(eeg_differences=eeg_differences)\n    \n    if is_submission is False:\n        print(f'\\nSpectrogram 50/10 Second Shape {spectrogram_50_10_second.shape}')\n        visualize_image(spectrogram_50_10_second, title='Spectrogram 50/10 Second')\n    \n    # Create 50/10 second spectrogram predictions tensors for the current sample\n    spectrogram_50_10_second_2d_predictions = torch.zeros(n_spectrogram_50_10_second_2d_models, 1, 6)\n    \n    for model_idx, (models, transforms) in enumerate(zip(spectrogram_50_10_second_2d_models, spectrogram_50_10_second_2d_transforms)):\n        \n        #continue\n                                        \n        spectrogram_50_10_second_2d_inputs = transforms(image=spectrogram_50_10_second)['image'].float()\n        spectrogram_50_10_second_2d_inputs = spectrogram_50_10_second_2d_inputs.to(device)\n\n        if tta:\n            spectrogram_50_10_second_2d_inputs = torch.stack((\n                spectrogram_50_10_second_2d_inputs,\n                torch.flip(spectrogram_50_10_second_2d_inputs, dims=(1,)),\n                torch.flip(spectrogram_50_10_second_2d_inputs, dims=(2,)),\n                torch.flip(spectrogram_50_10_second_2d_inputs, dims=(1, 2))\n            ), dim=0)\n        else:\n            spectrogram_50_10_second_2d_inputs = torch.unsqueeze(spectrogram_50_10_second_2d_inputs, dim=0)\n        \n        for fold, model in enumerate(models.values()):\n            with torch.no_grad():\n                if amp:\n                    with torch.autocast(device_type='cuda', dtype=torch.float16):\n                        spectrogram_50_10_second_2d_outputs = model(spectrogram_50_10_second_2d_inputs)\n                else:\n                    spectrogram_50_10_second_2d_outputs = model(spectrogram_50_10_second_2d_inputs)\n\n            spectrogram_50_10_second_2d_outputs = spectrogram_50_10_second_2d_outputs.cpu()\n            if tta:\n                spectrogram_50_10_second_2d_outputs = torch.mean(spectrogram_50_10_second_2d_outputs, dim=0)\n            else:\n                spectrogram_50_10_second_2d_outputs = torch.squeeze(spectrogram_50_10_second_2d_outputs, dim=0)\n            \n            spectrogram_50_10_second_2d_predictions[model_idx] += spectrogram_50_10_second_2d_outputs / len(models)\n            \n            if is_submission is False:\n                print(f'Predicted with Spectrogram 50/10 Second 2D Model {model_idx + 1} Fold {fold + 1}')\n    \n    test_spectrogram_50_10_second_2d_predictions.append(spectrogram_50_10_second_2d_predictions)\n    \n    # Create 50/30/10 second spectrogram from EEG differences\n    spectrogram_50_30_10_second = get_spectrogram_50_30_10_second(eeg_differences=eeg_differences)\n    \n    if is_submission is False:\n        print(f'\\nSpectrogram 50/30/10 Second Shape {spectrogram_50_30_10_second.shape}')\n        visualize_image(spectrogram_50_30_10_second, title='Spectrogram 50/30/10 Second')\n    \n    # Create 50/30/10 second spectrogram predictions tensors for the current sample\n    spectrogram_50_30_10_second_2d_predictions = torch.zeros(n_spectrogram_50_30_10_second_2d_models, 1, 6)\n    \n    for model_idx, (models, transforms) in enumerate(zip(spectrogram_50_30_10_second_2d_models, spectrogram_50_30_10_second_2d_transforms)):\n        \n        #continue\n                                        \n        spectrogram_50_30_10_second_2d_inputs = transforms(image=spectrogram_50_30_10_second)['image'].float()\n        spectrogram_50_30_10_second_2d_inputs = spectrogram_50_30_10_second_2d_inputs.to(device)\n\n        if tta:\n            spectrogram_50_30_10_second_2d_inputs = torch.stack((\n                spectrogram_50_30_10_second_2d_inputs,\n                torch.flip(spectrogram_50_30_10_second_2d_inputs, dims=(1,)),\n                torch.flip(spectrogram_50_30_10_second_2d_inputs, dims=(2,)),\n                torch.flip(spectrogram_50_30_10_second_2d_inputs, dims=(1, 2))\n            ), dim=0)\n        else:\n            spectrogram_50_30_10_second_2d_inputs = torch.unsqueeze(spectrogram_50_30_10_second_2d_inputs, dim=0)\n        \n        for fold, model in enumerate(models.values()):\n            with torch.no_grad():\n                if amp:\n                    with torch.autocast(device_type='cuda', dtype=torch.float16):\n                        spectrogram_50_30_10_second_2d_outputs = model(spectrogram_50_30_10_second_2d_inputs)\n                else:\n                    spectrogram_50_30_10_second_2d_outputs = model(spectrogram_50_30_10_second_2d_inputs)\n\n            spectrogram_50_30_10_second_2d_outputs = spectrogram_50_30_10_second_2d_outputs.cpu()\n            if tta:\n                spectrogram_50_30_10_second_2d_outputs = torch.mean(spectrogram_50_30_10_second_2d_outputs, dim=0)\n            else:\n                spectrogram_50_30_10_second_2d_outputs = torch.squeeze(spectrogram_50_30_10_second_2d_outputs, dim=0)\n            \n            spectrogram_50_30_10_second_2d_predictions[model_idx] += spectrogram_50_30_10_second_2d_outputs / len(models)\n            \n            if is_submission is False:\n                print(f'Predicted with Spectrogram 50/30/10 Second 2D Model {model_idx + 1} Fold {fold + 1}')\n    \n    test_spectrogram_50_30_10_second_2d_predictions.append(spectrogram_50_30_10_second_2d_predictions)\n    \n    # Create 30/10 second spectrogram from EEG differences\n    spectrogram_30_10_second = get_spectrogram_30_10_second(eeg_differences=eeg_differences)\n    \n    if is_submission is False:\n        print(f'\\nSpectrogram 30/10 Second Shape {spectrogram_30_10_second.shape}')\n        visualize_image(spectrogram_30_10_second, title='Spectrogram 30/10 Second')\n    \n    # Create 30/10 second spectrogram predictions tensors for the current sample\n    spectrogram_30_10_second_2d_predictions = torch.zeros(n_spectrogram_30_10_second_2d_models, 1, 6)\n    \n    for model_idx, (models, transforms) in enumerate(zip(spectrogram_30_10_second_2d_models, spectrogram_30_10_second_2d_transforms)):\n        \n        #continue\n\n        spectrogram_30_10_second_2d_inputs = transforms(image=spectrogram_30_10_second)['image'].float()\n        spectrogram_30_10_second_2d_inputs = spectrogram_30_10_second_2d_inputs.to(device)\n\n        if tta:\n            spectrogram_30_10_second_2d_inputs = torch.stack((\n                spectrogram_30_10_second_2d_inputs,\n                torch.flip(spectrogram_30_10_second_2d_inputs, dims=(1,)),\n                torch.flip(spectrogram_30_10_second_2d_inputs, dims=(2,)),\n                torch.flip(spectrogram_30_10_second_2d_inputs, dims=(1, 2))\n            ), dim=0)\n        else:\n            spectrogram_30_10_second_2d_inputs = torch.unsqueeze(spectrogram_30_10_second_2d_inputs, dim=0)\n        \n        for fold, model in enumerate(models.values()):\n            with torch.no_grad():\n                if amp:\n                    with torch.autocast(device_type='cuda', dtype=torch.float16):\n                        spectrogram_30_10_second_2d_outputs = model(spectrogram_30_10_second_2d_inputs)\n                else:\n                    spectrogram_30_10_second_2d_outputs = model(spectrogram_30_10_second_2d_inputs)\n\n            spectrogram_30_10_second_2d_outputs = spectrogram_30_10_second_2d_outputs.cpu()\n            if tta:\n                spectrogram_30_10_second_2d_outputs = torch.mean(spectrogram_30_10_second_2d_outputs, dim=0)\n            else:\n                spectrogram_30_10_second_2d_outputs = torch.squeeze(spectrogram_30_10_second_2d_outputs, dim=0)\n            \n            spectrogram_30_10_second_2d_predictions[model_idx] += spectrogram_30_10_second_2d_outputs / len(models)\n            \n            if is_submission is False:\n                print(f'Predicted with Spectrogram 30/10 Second 2D Model {model_idx + 1} Fold {fold + 1}')\n    \n    test_spectrogram_30_10_second_2d_predictions.append(spectrogram_30_10_second_2d_predictions)\n        \ntest_raw_eeg_50_second_2d_predictions = torch.cat(test_raw_eeg_50_second_2d_predictions, dim=1)\ntest_raw_eeg_50_second_2d_convnext_predictions = test_raw_eeg_50_second_2d_predictions[0, :, :]\ntest_raw_eeg_50_second_2d_maxvit_predictions = test_raw_eeg_50_second_2d_predictions[1, :, :]\n\ntest_raw_eeg_50_10_second_2d_predictions = torch.cat(test_raw_eeg_50_10_second_2d_predictions, dim=1)\ntest_raw_eeg_50_10_second_2d_convnext_predictions = test_raw_eeg_50_10_second_2d_predictions[0, :, :]\ntest_raw_eeg_50_10_second_2d_maxvit_predictions = test_raw_eeg_50_10_second_2d_predictions[1, :, :]\n\ntest_spectrogram_50_second_2d_predictions = torch.cat(test_spectrogram_50_second_2d_predictions, dim=1)\ntest_spectrogram_50_second_2d_convnext_predictions = test_spectrogram_50_second_2d_predictions[0, :, :]\ntest_spectrogram_50_second_2d_maxvit_predictions = test_spectrogram_50_second_2d_predictions[1, :, :]\n\ntest_spectrogram_50_10_second_2d_predictions = torch.cat(test_spectrogram_50_10_second_2d_predictions, dim=1)\ntest_spectrogram_50_10_second_2d_convnext_predictions = test_spectrogram_50_10_second_2d_predictions[0, :, :]\ntest_spectrogram_50_10_second_2d_maxvit_predictions = test_spectrogram_50_10_second_2d_predictions[1, :, :]\n\ntest_spectrogram_50_30_10_second_2d_predictions = torch.cat(test_spectrogram_50_30_10_second_2d_predictions, dim=1)\ntest_spectrogram_50_30_10_second_2d_convnext_predictions = test_spectrogram_50_30_10_second_2d_predictions[0, :, :]\ntest_spectrogram_50_30_10_second_2d_maxvit_predictions = test_spectrogram_50_30_10_second_2d_predictions[1, :, :]\n\ntest_spectrogram_30_10_second_2d_predictions = torch.cat(test_spectrogram_30_10_second_2d_predictions, dim=1)\ntest_spectrogram_30_10_second_2d_convnext_predictions = test_spectrogram_30_10_second_2d_predictions[0, :, :]\ntest_spectrogram_30_10_second_2d_maxvit_predictions = test_spectrogram_30_10_second_2d_predictions[1, :, :]\n    \ntest_raw_eeg_50_second_2d_convnext_predictions = torch.sigmoid(test_raw_eeg_50_second_2d_convnext_predictions).numpy()\ntest_raw_eeg_50_second_2d_convnext_predictions = test_raw_eeg_50_second_2d_convnext_predictions / test_raw_eeg_50_second_2d_convnext_predictions.sum(axis=-1).reshape(-1, 1)\ntest_raw_eeg_50_second_2d_maxvit_predictions = torch.sigmoid(test_raw_eeg_50_second_2d_maxvit_predictions).numpy()\ntest_raw_eeg_50_second_2d_maxvit_predictions = test_raw_eeg_50_second_2d_maxvit_predictions / test_raw_eeg_50_second_2d_maxvit_predictions.sum(axis=-1).reshape(-1, 1)\n\ntest_raw_eeg_50_10_second_2d_convnext_predictions = torch.sigmoid(test_raw_eeg_50_10_second_2d_convnext_predictions).numpy()\ntest_raw_eeg_50_10_second_2d_convnext_predictions = test_raw_eeg_50_10_second_2d_convnext_predictions / test_raw_eeg_50_10_second_2d_convnext_predictions.sum(axis=-1).reshape(-1, 1)\ntest_raw_eeg_50_10_second_2d_maxvit_predictions = torch.sigmoid(test_raw_eeg_50_10_second_2d_maxvit_predictions).numpy()\ntest_raw_eeg_50_10_second_2d_maxvit_predictions = test_raw_eeg_50_10_second_2d_maxvit_predictions / test_raw_eeg_50_10_second_2d_maxvit_predictions.sum(axis=-1).reshape(-1, 1)\n\ntest_spectrogram_50_second_2d_convnext_predictions = torch.sigmoid(test_spectrogram_50_second_2d_convnext_predictions).numpy()\ntest_spectrogram_50_second_2d_convnext_predictions = test_spectrogram_50_second_2d_convnext_predictions / test_spectrogram_50_second_2d_convnext_predictions.sum(axis=-1).reshape(-1, 1)\ntest_spectrogram_50_second_2d_maxvit_predictions = torch.sigmoid(test_spectrogram_50_second_2d_maxvit_predictions).numpy()\ntest_spectrogram_50_second_2d_maxvit_predictions = test_spectrogram_50_second_2d_maxvit_predictions / test_spectrogram_50_second_2d_maxvit_predictions.sum(axis=-1).reshape(-1, 1)\n\ntest_spectrogram_50_10_second_2d_convnext_predictions = torch.sigmoid(test_spectrogram_50_10_second_2d_convnext_predictions).numpy()\ntest_spectrogram_50_10_second_2d_convnext_predictions = test_spectrogram_50_10_second_2d_convnext_predictions / test_spectrogram_50_10_second_2d_convnext_predictions.sum(axis=-1).reshape(-1, 1)\ntest_spectrogram_50_10_second_2d_maxvit_predictions = torch.sigmoid(test_spectrogram_50_10_second_2d_maxvit_predictions).numpy()\ntest_spectrogram_50_10_second_2d_maxvit_predictions = test_spectrogram_50_10_second_2d_maxvit_predictions / test_spectrogram_50_10_second_2d_maxvit_predictions.sum(axis=-1).reshape(-1, 1)\n\ntest_spectrogram_50_30_10_second_2d_convnext_predictions = torch.sigmoid(test_spectrogram_50_30_10_second_2d_convnext_predictions).numpy()\ntest_spectrogram_50_30_10_second_2d_convnext_predictions = test_spectrogram_50_30_10_second_2d_convnext_predictions / test_spectrogram_50_30_10_second_2d_convnext_predictions.sum(axis=-1).reshape(-1, 1)\ntest_spectrogram_50_30_10_second_2d_maxvit_predictions = torch.sigmoid(test_spectrogram_50_30_10_second_2d_maxvit_predictions).numpy()\ntest_spectrogram_50_30_10_second_2d_maxvit_predictions = test_spectrogram_50_30_10_second_2d_maxvit_predictions / test_spectrogram_50_30_10_second_2d_maxvit_predictions.sum(axis=-1).reshape(-1, 1)\n\ntest_spectrogram_30_10_second_2d_convnext_predictions = torch.sigmoid(test_spectrogram_30_10_second_2d_convnext_predictions).numpy()\ntest_spectrogram_30_10_second_2d_convnext_predictions = test_spectrogram_30_10_second_2d_convnext_predictions / test_spectrogram_30_10_second_2d_convnext_predictions.sum(axis=-1).reshape(-1, 1)\ntest_spectrogram_30_10_second_2d_maxvit_predictions = torch.sigmoid(test_spectrogram_30_10_second_2d_maxvit_predictions).numpy()\ntest_spectrogram_30_10_second_2d_maxvit_predictions = test_spectrogram_30_10_second_2d_maxvit_predictions / test_spectrogram_30_10_second_2d_maxvit_predictions.sum(axis=-1).reshape(-1, 1)","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:16:59.626372Z","iopub.execute_input":"2024-04-05T08:16:59.626815Z","iopub.status.idle":"2024-04-05T08:17:10.493505Z","shell.execute_reply.started":"2024-04-05T08:16:59.626785Z","shell.execute_reply":"2024-04-05T08:17:10.492476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_raw_eeg_50_second_2d_convnext_predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:17:22.024324Z","iopub.execute_input":"2024-04-05T08:17:22.024723Z","iopub.status.idle":"2024-04-05T08:17:22.031878Z","shell.execute_reply.started":"2024-04-05T08:17:22.024691Z","shell.execute_reply":"2024-04-05T08:17:22.030922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_raw_eeg_50_second_2d_maxvit_predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:17:22.307539Z","iopub.execute_input":"2024-04-05T08:17:22.308447Z","iopub.status.idle":"2024-04-05T08:17:22.314802Z","shell.execute_reply.started":"2024-04-05T08:17:22.308414Z","shell.execute_reply":"2024-04-05T08:17:22.313809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_raw_eeg_50_10_second_2d_convnext_predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:17:22.779163Z","iopub.execute_input":"2024-04-05T08:17:22.779528Z","iopub.status.idle":"2024-04-05T08:17:22.785822Z","shell.execute_reply.started":"2024-04-05T08:17:22.779499Z","shell.execute_reply":"2024-04-05T08:17:22.784795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_raw_eeg_50_10_second_2d_maxvit_predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:17:23.474503Z","iopub.execute_input":"2024-04-05T08:17:23.474913Z","iopub.status.idle":"2024-04-05T08:17:23.481507Z","shell.execute_reply.started":"2024-04-05T08:17:23.474882Z","shell.execute_reply":"2024-04-05T08:17:23.480631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_spectrogram_50_second_2d_convnext_predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:17:24.133322Z","iopub.execute_input":"2024-04-05T08:17:24.134245Z","iopub.status.idle":"2024-04-05T08:17:24.140620Z","shell.execute_reply.started":"2024-04-05T08:17:24.134211Z","shell.execute_reply":"2024-04-05T08:17:24.139659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_spectrogram_50_second_2d_maxvit_predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:17:24.405474Z","iopub.execute_input":"2024-04-05T08:17:24.405904Z","iopub.status.idle":"2024-04-05T08:17:24.412769Z","shell.execute_reply.started":"2024-04-05T08:17:24.405868Z","shell.execute_reply":"2024-04-05T08:17:24.411622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_spectrogram_50_10_second_2d_convnext_predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:17:25.431393Z","iopub.execute_input":"2024-04-05T08:17:25.431871Z","iopub.status.idle":"2024-04-05T08:17:25.439663Z","shell.execute_reply.started":"2024-04-05T08:17:25.431833Z","shell.execute_reply":"2024-04-05T08:17:25.438432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_spectrogram_50_10_second_2d_maxvit_predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:17:25.815526Z","iopub.execute_input":"2024-04-05T08:17:25.816338Z","iopub.status.idle":"2024-04-05T08:17:25.822899Z","shell.execute_reply.started":"2024-04-05T08:17:25.816284Z","shell.execute_reply":"2024-04-05T08:17:25.821886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_spectrogram_50_30_10_second_2d_convnext_predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:17:26.186974Z","iopub.execute_input":"2024-04-05T08:17:26.187634Z","iopub.status.idle":"2024-04-05T08:17:26.194231Z","shell.execute_reply.started":"2024-04-05T08:17:26.187587Z","shell.execute_reply":"2024-04-05T08:17:26.193242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_spectrogram_50_30_10_second_2d_maxvit_predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:17:28.735889Z","iopub.execute_input":"2024-04-05T08:17:28.736538Z","iopub.status.idle":"2024-04-05T08:17:28.742562Z","shell.execute_reply.started":"2024-04-05T08:17:28.736505Z","shell.execute_reply":"2024-04-05T08:17:28.741518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_spectrogram_30_10_second_2d_maxvit_predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:17:32.944218Z","iopub.execute_input":"2024-04-05T08:17:32.944665Z","iopub.status.idle":"2024-04-05T08:17:32.952222Z","shell.execute_reply.started":"2024-04-05T08:17:32.944621Z","shell.execute_reply":"2024-04-05T08:17:32.951018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Ensemble","metadata":{}},{"cell_type":"code","source":"def normalize_probabilities(df, columns):\n\n    \"\"\"\n    Normalize probabilities to 1 within given columns\n\n    Parameters\n    ----------\n    df: pandas.DataFrame\n        Dataframe with given columns\n\n    columns: list\n        List of column names that have probabilities\n\n    Returns\n    -------\n    df: pandas.DataFrame\n        Dataframe with given columns' sum adjusted to 1\n    \"\"\"\n\n    df_sums = df[columns].sum(axis=1)\n    for column in columns:\n        df[column] /= df_sums\n\n    return df\n","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:17:36.102884Z","iopub.execute_input":"2024-04-05T08:17:36.103322Z","iopub.status.idle":"2024-04-05T08:17:36.110053Z","shell.execute_reply.started":"2024-04-05T08:17:36.103289Z","shell.execute_reply":"2024-04-05T08:17:36.108937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_columns = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\n\n# 50 second raw EEG prediction columns\nraw_eeg_50_second_2d_convnext_prediction_columns = [f'{column}_raw_eeg_50_second_2d_convnext_prediction' for column in target_columns]\nraw_eeg_50_second_2d_maxvit_prediction_columns = [f'{column}_raw_eeg_50_second_2d_maxvit_prediction' for column in target_columns]\n\ndf[raw_eeg_50_second_2d_convnext_prediction_columns] = test_raw_eeg_50_second_2d_convnext_predictions\ndf[raw_eeg_50_second_2d_maxvit_prediction_columns] = test_raw_eeg_50_second_2d_maxvit_predictions\n\ndf = normalize_probabilities(df=df, columns=raw_eeg_50_second_2d_convnext_prediction_columns)\ndf = normalize_probabilities(df=df, columns=raw_eeg_50_second_2d_maxvit_prediction_columns)\n\n# 50/10 second raw EEG prediction columns\nraw_eeg_50_10_second_2d_convnext_prediction_columns = [f'{column}_raw_eeg_50_10_second_2d_convnext_prediction' for column in target_columns]\nraw_eeg_50_10_second_2d_maxvit_prediction_columns = [f'{column}_raw_eeg_50_10_second_2d_maxvit_prediction' for column in target_columns]\n\ndf[raw_eeg_50_10_second_2d_convnext_prediction_columns] = test_raw_eeg_50_10_second_2d_convnext_predictions\ndf[raw_eeg_50_10_second_2d_maxvit_prediction_columns] = test_raw_eeg_50_10_second_2d_maxvit_predictions\n\ndf = normalize_probabilities(df=df, columns=raw_eeg_50_10_second_2d_convnext_prediction_columns)\ndf = normalize_probabilities(df=df, columns=raw_eeg_50_10_second_2d_maxvit_prediction_columns)\n\n# 50 second spectrogram prediction columns\nspectrogram_50_second_2d_convnext_prediction_columns = [f'{column}_spectrogram_50_second_2d_convnext_prediction' for column in target_columns]\nspectrogram_50_second_2d_maxvit_prediction_columns = [f'{column}_spectrogram_50_second_2d_maxvit_prediction' for column in target_columns]\n\ndf[spectrogram_50_second_2d_convnext_prediction_columns] = test_spectrogram_50_second_2d_convnext_predictions\ndf[spectrogram_50_second_2d_maxvit_prediction_columns] = test_spectrogram_50_second_2d_maxvit_predictions\n\ndf = normalize_probabilities(df=df, columns=spectrogram_50_second_2d_convnext_prediction_columns)\ndf = normalize_probabilities(df=df, columns=spectrogram_50_second_2d_maxvit_prediction_columns)\n\n# 50/10 second spectrogram prediction columns\nspectrogram_50_10_second_2d_convnext_prediction_columns = [f'{column}_spectrogram_50_10_second_2d_convnext_prediction' for column in target_columns]\nspectrogram_50_10_second_2d_maxvit_prediction_columns = [f'{column}_spectrogram_50_10_second_2d_maxvit_prediction' for column in target_columns]\n\ndf[spectrogram_50_10_second_2d_convnext_prediction_columns] = test_spectrogram_50_10_second_2d_convnext_predictions\ndf[spectrogram_50_10_second_2d_maxvit_prediction_columns] = test_spectrogram_50_10_second_2d_maxvit_predictions\n\ndf = normalize_probabilities(df=df, columns=spectrogram_50_10_second_2d_convnext_prediction_columns)\ndf = normalize_probabilities(df=df, columns=spectrogram_50_10_second_2d_maxvit_prediction_columns)\n\n# 50/30/10 second spectrogram prediction columns\nspectrogram_50_30_10_second_2d_convnext_prediction_columns = [f'{column}_spectrogram_50_30_10_second_2d_convnext_prediction' for column in target_columns]\nspectrogram_50_30_10_second_2d_maxvit_prediction_columns = [f'{column}_spectrogram_50_30_10_second_2d_maxvit_prediction' for column in target_columns]\n\ndf[spectrogram_50_30_10_second_2d_convnext_prediction_columns] = test_spectrogram_50_30_10_second_2d_convnext_predictions\ndf[spectrogram_50_30_10_second_2d_maxvit_prediction_columns] = test_spectrogram_50_30_10_second_2d_maxvit_predictions\n\ndf = normalize_probabilities(df=df, columns=spectrogram_50_30_10_second_2d_convnext_prediction_columns)\ndf = normalize_probabilities(df=df, columns=spectrogram_50_30_10_second_2d_maxvit_prediction_columns)\n\n# 30/10 second spectrogram prediction columns\nspectrogram_30_10_second_2d_convnext_prediction_columns = [f'{column}_spectrogram_30_10_second_2d_convnext_prediction' for column in target_columns]\nspectrogram_30_10_second_2d_maxvit_prediction_columns = [f'{column}_spectrogram_30_10_second_2d_maxvit_prediction' for column in target_columns]\n\ndf[spectrogram_30_10_second_2d_convnext_prediction_columns] = test_spectrogram_30_10_second_2d_convnext_predictions\ndf[spectrogram_30_10_second_2d_maxvit_prediction_columns] = test_spectrogram_30_10_second_2d_maxvit_predictions\n\ndf = normalize_probabilities(df=df, columns=spectrogram_30_10_second_2d_convnext_prediction_columns)\ndf = normalize_probabilities(df=df, columns=spectrogram_30_10_second_2d_maxvit_prediction_columns)\n\ndf","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:18:04.648138Z","iopub.execute_input":"2024-04-05T08:18:04.648530Z","iopub.status.idle":"2024-04-05T08:18:04.737591Z","shell.execute_reply.started":"2024-04-05T08:18:04.648500Z","shell.execute_reply":"2024-04-05T08:18:04.736566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Raw EEG blend\nraw_eeg_blend_prediction_columns = [f'{column}_raw_eeg_blend_prediction' for column in target_columns]\ndf[raw_eeg_blend_prediction_columns] = (df[raw_eeg_50_second_2d_convnext_prediction_columns] * 0.25).values + \\\n                                       (df[raw_eeg_50_second_2d_maxvit_prediction_columns] * 0.25).values + \\\n                                       (df[raw_eeg_50_10_second_2d_convnext_prediction_columns] * 0.25).values + \\\n                                       (df[raw_eeg_50_10_second_2d_maxvit_prediction_columns] * 0.25).values\n\ndf[raw_eeg_blend_prediction_columns]","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:18:13.388342Z","iopub.execute_input":"2024-04-05T08:18:13.389158Z","iopub.status.idle":"2024-04-05T08:18:13.411536Z","shell.execute_reply.started":"2024-04-05T08:18:13.389119Z","shell.execute_reply":"2024-04-05T08:18:13.410616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Spectrogram blend\nspectrogram_blend_prediction_columns = [f'{column}_spectrogram_blend_prediction' for column in target_columns]\ndf[spectrogram_blend_prediction_columns] = (df[spectrogram_50_second_2d_convnext_prediction_columns] * 0.10).values + \\\n                                           (df[spectrogram_50_second_2d_maxvit_prediction_columns] * 0.10).values + \\\n                                           (df[spectrogram_50_10_second_2d_convnext_prediction_columns] * 0.15).values + \\\n                                           (df[spectrogram_50_10_second_2d_maxvit_prediction_columns] * 0.15).values + \\\n                                           (df[spectrogram_30_10_second_2d_convnext_prediction_columns] * 0.15).values + \\\n                                           (df[spectrogram_30_10_second_2d_maxvit_prediction_columns] * 0.15).values + \\\n                                           (df[spectrogram_50_30_10_second_2d_convnext_prediction_columns] * 0.10).values + \\\n                                           (df[spectrogram_50_30_10_second_2d_maxvit_prediction_columns] * 0.10).values\n\ndf[spectrogram_blend_prediction_columns]","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:18:22.099089Z","iopub.execute_input":"2024-04-05T08:18:22.099841Z","iopub.status.idle":"2024-04-05T08:18:22.123676Z","shell.execute_reply.started":"2024-04-05T08:18:22.099806Z","shell.execute_reply":"2024-04-05T08:18:22.122530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Final blend\nblend_prediction_columns = [f'{column}_blend_prediction' for column in target_columns]\ndf[blend_prediction_columns] = (df[raw_eeg_blend_prediction_columns] * 0.3).values + \\\n                               (df[spectrogram_blend_prediction_columns] * 0.7).values\n\ndf[blend_prediction_columns]","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:18:22.295428Z","iopub.execute_input":"2024-04-05T08:18:22.296299Z","iopub.status.idle":"2024-04-05T08:18:22.313951Z","shell.execute_reply.started":"2024-04-05T08:18:22.296263Z","shell.execute_reply":"2024-04-05T08:18:22.312930Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"df_submission = df.loc[:, ['eeg_id']]\ndf_submission[target_columns] = df[blend_prediction_columns]\ndf_submission","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:18:25.989672Z","iopub.execute_input":"2024-04-05T08:18:25.990701Z","iopub.status.idle":"2024-04-05T08:18:26.009693Z","shell.execute_reply.started":"2024-04-05T08:18:25.990658Z","shell.execute_reply":"2024-04-05T08:18:26.008525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-04-05T08:18:28.373668Z","iopub.execute_input":"2024-04-05T08:18:28.374668Z","iopub.status.idle":"2024-04-05T08:18:28.384084Z","shell.execute_reply.started":"2024-04-05T08:18:28.374625Z","shell.execute_reply":"2024-04-05T08:18:28.383039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}