{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":8062848,"sourceType":"datasetVersion","datasetId":4730481}],"dockerImageVersionId":30673,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# here is the code for ResNet inference ","metadata":{}},{"cell_type":"code","source":"device='cuda'","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:13.002526Z","iopub.execute_input":"2024-04-08T12:45:13.003538Z","iopub.status.idle":"2024-04-08T12:45:13.007684Z","shell.execute_reply.started":"2024-04-08T12:45:13.003495Z","shell.execute_reply":"2024-04-08T12:45:13.006633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport math\n\nclass SincConv_fast(nn.Module):\n    \"\"\"Sinc-based convolution\n    Parameters for EEG-based emotion recognition\n    ----------\n    in_channels : `int`\n        Number of input channels. Must be 20.\n    out_channels : `int`\n        Number of filters.\n    kernel_size : `int`\n        Filter length.\n    sample_rate : `int`, optional\n        Sample rate. Defaults to 500.\n    Usage\n    -----\n    See `torch.nn.Conv2d`\n    Reference\n    ---------\n    Mirco Ravanelli, Yoshua Bengio,\n    \"Speaker Recognition from raw waveform with SincNet\".\n    https://arxiv.org/abs/1808.00158\n    \"\"\"\n\n    @staticmethod\n    def to_mel(hz):\n        return 2595 * np.log10(1 + hz / 700)\n\n    @staticmethod\n    def to_hz(mel):\n        return 700 * (10 ** (mel / 2595) - 1)\n\n    def __init__(self, out_channels, kernel_size, sample_rate=500, in_channels=20,\n                 stride=1, padding=0, dilation=1, bias=False, groups=1, min_low_hz=0.1, min_band_hz=10):\n\n        super(SincConv_fast, self).__init__()\n\n        #if in_channels != 20:\n        #    raise ValueError(\"SincConv only supports 20 input channels, got {}.\".format(in_channels))\n\n        self.out_channels = out_channels\n        self.kernel_size = kernel_size\n\n        # Forcing the filters to be odd (i.e., perfectly symmetrical)\n        if kernel_size % 2 == 0:\n            self.kernel_size += 1\n\n        self.stride = stride\n        self.padding = padding\n        self.dilation = dilation\n\n        if bias:\n            raise ValueError('SincConv does not support bias.')\n        if groups > 1:\n            raise ValueError('SincConv does not support groups.')\n\n        self.sample_rate = sample_rate\n        self.min_low_hz = min_low_hz\n        self.min_band_hz = min_band_hz\n\n        # Initialize filterbanks such that they are equally spaced in Mel scale around the complete range of the pre-filtered EEG trial\n        low_hz = 0.1\n        high_hz = 30\n\n        mel = np.linspace(self.to_mel(low_hz),\n                          self.to_mel(high_hz),\n                          self.out_channels + 1)\n        hz = self.to_hz(mel)\n\n        self.hz = hz\n\n        # Filter lower frequency (out_channels, 1)\n        self.low_hz_ = nn.Parameter(torch.Tensor(hz[:-1]).view(-1, 1))\n\n        # Filter frequency band (out_channels, 1)\n        self.band_hz_ = nn.Parameter(torch.Tensor(np.diff(hz)).view(-1, 1))\n\n        # Setting the Hamming window\n        # Computing only half of the window\n        n_lin = torch.linspace(0, (self.kernel_size / 2) - 1,\n                               steps=int((self.kernel_size / 2)))\n        self.window_ = 0.54 - 0.46 * torch.cos(2 * math.pi * n_lin / self.kernel_size)\n\n        n = (self.kernel_size - 1) / 2.0\n        # Due to symmetry, only need half of the time axes\n        self.n_ = 2 * math.pi * torch.arange(-n, 0).view(1, -1) / self.sample_rate\n\n    def forward(self, waveforms):\n        \"\"\"\n        Parameters\n        ----------\n        waveforms : `torch.Tensor` (batch_size, num_channels, num_samples)\n            Batch of waveforms.\n        Returns\n        -------\n        features : `torch.Tensor` (batch_size, out_channels, num_channels, n_samples_out)\n            Batch of sinc filter activations.\n        \"\"\"\n\n        self.n_ = self.n_.to(waveforms.device)\n        self.window_ = self.window_.to(waveforms.device)\n\n        low = self.min_low_hz + torch.abs(self.low_hz_)\n        high = torch.clamp(low + self.min_band_hz +\n                           torch.abs(self.band_hz_), self.min_low_hz, self.sample_rate / 2)\n        band = (high - low)[:, 0]\n\n        f_times_t_low = torch.matmul(low, self.n_)\n        f_times_t_high = torch.matmul(high, self.n_)\n\n        # Equivalent of Eq.4 of the reference paper (SPEAKER RECOGNITION FROM RAW WAVEFORM WITH SINCNET).\n        # I just have expanded the sinc and simplified the terms. This way I avoid several useless computations (Mirco).\n        band_pass_left = ((torch.sin(f_times_t_high) -\n                           torch.sin(f_times_t_low)) / (self.n_ / 2)) * self.window_\n        band_pass_center = 2 * band.view(-1, 1)\n        band_pass_right = torch.flip(band_pass_left, dims=[1])\n\n        band_pass = torch.cat(\n            [band_pass_left, band_pass_center, band_pass_right], dim=1)\n\n        band_pass = band_pass / (2 * band[:, None])\n\n        # Expand the filters for each input channel\n        #filters = band_pass.unsqueeze(2).permute(2, 0, 1)\n        f = (band_pass).view(self.out_channels, 1, self.kernel_size)\n\n        # Apply the filters to each input channel separately\n        outputs = []\n        for i in range(waveforms.size(1)):\n            output = F.conv1d(waveforms[:, i:i+1, :], f, stride=self.stride,\n                              padding=self.padding, dilation=self.dilation,\n                              bias=None, groups=1)\n            outputs.append(output)\n\n        # Concatenate the outputs along the channel dimension\n        outputs = torch.stack(outputs, dim=2)\n\n        return outputs\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:13.914389Z","iopub.execute_input":"2024-04-08T12:45:13.914765Z","iopub.status.idle":"2024-04-08T12:45:18.265101Z","shell.execute_reply.started":"2024-04-08T12:45:13.914713Z","shell.execute_reply":"2024-04-08T12:45:18.264211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ResNet_1D_Block(nn.Module):\n\n    def __init__(self, in_channels, out_channels, kernel_size, stride, padding, downsampling):\n        super(ResNet_1D_Block, self).__init__()\n        self.bn1 = nn.BatchNorm2d(num_features=in_channels)\n        self.relu = nn.ReLU(inplace=False)\n        self.dropout = nn.Dropout(p=0.0, inplace=False)\n        self.conv1 = nn.Conv2d(in_channels=in_channels, out_channels=out_channels, kernel_size=kernel_size,\n                               stride=stride, padding=padding, bias=False)\n        self.bn2 = nn.BatchNorm2d(num_features=out_channels)\n        self.conv2 = nn.Conv2d(in_channels=out_channels, out_channels=out_channels, kernel_size=kernel_size,\n                               stride=stride, padding=padding, bias=False)\n        self.maxpool = nn.MaxPool2d(kernel_size=(1, 1), stride=2, padding=0)\n        self.downsampling = downsampling\n\n    def forward(self, x):\n        identity = x\n\n        out = self.bn1(x)\n        out = self.relu(out)\n        out = self.dropout(out)\n        out = self.conv1(out)\n        out = self.bn2(out)\n        out = self.relu(out)\n        out = self.dropout(out)\n        out = self.conv2(out)\n\n        out = self.maxpool(out)\n        identity = self.downsampling(x)\n        \n        out += identity\n        return out","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:18.266875Z","iopub.execute_input":"2024-04-08T12:45:18.267289Z","iopub.status.idle":"2024-04-08T12:45:18.278377Z","shell.execute_reply.started":"2024-04-08T12:45:18.267261Z","shell.execute_reply":"2024-04-08T12:45:18.277259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport math\n\nclass EEGNet(nn.Module):\n\n    def __init__(self, kernels, sample_rate=500, in_channels=8, num_classes=6):\n        super(EEGNet, self).__init__()\n        self.kernels = kernels\n        self.planes = 24\n        self.parallel_conv = nn.ModuleList()\n        self.in_channels = in_channels\n        \n        # Replace initial convolution with SincConv_fast\n        self.conv0 = SincConv_fast(out_channels=self.planes, kernel_size=251, sample_rate=sample_rate, in_channels=in_channels)\n        \n        for i, kernel_size in enumerate(list(self.kernels)):\n            sep_conv = SincConv_fast(out_channels=self.planes, kernel_size=251, sample_rate=sample_rate, in_channels=in_channels)\n            #sep_conv = nn.Conv1d(in_channels=in_channels, out_channels=self.planes, kernel_size=(kernel_size),\n            #                   stride=1, padding=0, bias=False,)\n            self.parallel_conv.append(sep_conv)\n\n        self.bn1 = nn.BatchNorm2d(num_features=self.planes)\n        self.relu = nn.ReLU(inplace=False)\n        self.block = self._make_resnet_layer(kernel_size=(17, 17), stride=(1, 1), padding=(8,8))\n        self.bn2 = nn.BatchNorm2d(num_features=self.planes)\n        self.avgpool = nn.AvgPool2d(kernel_size=(1,1), stride=1, padding=0)\n        self.rnn = nn.GRU(input_size=self.in_channels, hidden_size=128, num_layers=1, bidirectional=True)\n        self.fc = nn.Linear(in_features=736, out_features=num_classes)\n        self.fc_two = nn.Linear(in_features=736, out_features=500)\n\n\n    def _make_resnet_layer(self, kernel_size, stride, blocks=9, padding=0):\n        layers = []\n        base_width = self.planes\n\n        for i in range(blocks):\n            downsampling = nn.Sequential(\n                    nn.MaxPool2d(kernel_size=(1, 1), stride=2, padding=0)\n                )\n            layers.append(ResNet_1D_Block(in_channels=self.planes, out_channels=self.planes, kernel_size=kernel_size,\n                                       stride=stride, padding=padding, downsampling=downsampling))\n\n        return nn.Sequential(*layers)\n    \n    def extract_features(self, x):\n        x = x.permute(0, 2, 1)\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        out = torch.cat(out_sep, dim=2)\n        \n        out = self.bn1(out)\n        out = self.relu(out)\n        \n        \n        out = self.block(out)\n        out = self.bn2(out)\n        out = self.relu(out)\n        out = self.avgpool(out)  \n        \n        out = out.reshape(out.shape[0], -1)  \n        rnn_out, _ = self.rnn(x.permute(0, 2, 1))\n        new_rnn_h = rnn_out[:, -1, :]  \n        \n        new_out = torch.cat([out, new_rnn_h], dim=1) \n        return new_out\n    \n    def forward(self, x):\n        new_out = self.extract_features(x)\n        result = self.fc(new_out)\n        two_result = self.fc_two(new_out)\n\n        return result, two_result\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:18.279619Z","iopub.execute_input":"2024-04-08T12:45:18.279969Z","iopub.status.idle":"2024-04-08T12:45:18.302489Z","shell.execute_reply.started":"2024-04-08T12:45:18.279938Z","shell.execute_reply":"2024-04-08T12:45:18.301488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"weights = torch.load('/kaggle/input/sinc-resnet/resnet1d_gru_fold1_best_version9_stage1.pth')\nweights.keys()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:18.305115Z","iopub.execute_input":"2024-04-08T12:45:18.305469Z","iopub.status.idle":"2024-04-08T12:45:18.763274Z","shell.execute_reply.started":"2024-04-08T12:45:18.305434Z","shell.execute_reply":"2024-04-08T12:45:18.762263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = EEGNet(kernels=[3,5,7,9], in_channels=8, num_classes=6)\n#model = nn.DataParallel(model)\nmodel.load_state_dict(weights['model'])","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:18.764950Z","iopub.execute_input":"2024-04-08T12:45:18.765344Z","iopub.status.idle":"2024-04-08T12:45:18.921205Z","shell.execute_reply.started":"2024-04-08T12:45:18.765309Z","shell.execute_reply":"2024-04-08T12:45:18.920114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd \n\ntest = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/test.csv')","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:18.922320Z","iopub.execute_input":"2024-04-08T12:45:18.922621Z","iopub.status.idle":"2024-04-08T12:45:20.196843Z","shell.execute_reply.started":"2024-04-08T12:45:18.922595Z","shell.execute_reply":"2024-04-08T12:45:20.195939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eeg_from_parquet(parquet_path: str, display: bool = False) -> np.ndarray:\n    \"\"\"\n    This function reads a parquet file and extracts the middle 50 seconds of readings. Then it fills NaN values\n    with the mean value (ignoring NaNs).\n    :param parquet_path: path to parquet file.\n    :param display: whether to display EEG plots or not.\n    :return data: np.array of shape  (time_steps, eeg_features) -> (10_000, 8)\n    \"\"\"\n    # === Extract middle 50 seconds ===\n    eeg = pd.read_parquet(parquet_path, columns=eeg_features)\n    rows = len(eeg)\n    offset = (rows - 10_000) // 2 # 50 * 200 = 10_000\n    eeg = eeg.iloc[offset:offset+10_000] # middle 50 seconds, has the same amount of readings to left and right\n    if display: \n        plt.figure(figsize=(10,5))\n        offset = 0\n    # === Convert to numpy ===\n    data = np.zeros((10_000, len(eeg_features))) # create placeholder of same shape with zeros\n    for index, feature in enumerate(eeg_features):\n        x = eeg[feature].values.astype('float32') # convert to float32\n        mean = np.nanmean(x) # arithmetic mean along the specified axis, ignoring NaNs\n        nan_percentage = np.isnan(x).mean() # percentage of NaN values in feature\n        # === Fill nan values ===\n        if nan_percentage < 1: # if some values are nan, but not all\n            x = np.nan_to_num(x, nan=mean)\n        else: # if all values are nan\n            x[:] = 0\n        data[:, index] = x\n        if display: \n            if index != 0:\n                offset += x.max()\n            plt.plot(range(10_000), x-offset, label=feature)\n            offset -= x.min()\n    if display:\n        plt.legend()\n        name = parquet_path.split('/')[-1].split('.')[0]\n        plt.yticks([])\n        plt.title(f'EEG {name}',size=16)\n        plt.show()    \n    return data","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:20.198148Z","iopub.execute_input":"2024-04-08T12:45:20.198498Z","iopub.status.idle":"2024-04-08T12:45:20.211154Z","shell.execute_reply.started":"2024-04-08T12:45:20.198468Z","shell.execute_reply":"2024-04-08T12:45:20.210110Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from glob import glob\nfrom tqdm.auto import tqdm\nimport matplotlib.pyplot as plt \n\neeg_features = ['Fp1','T3','C3','O1','Fp2','C4','T4','O2']\n\n\nCREATE_EEGS = True\nall_eegs = {}\nvisualize = 0\ndata_root = '/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/'\neeg_paths = glob('/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/' + \"*.parquet\")\neeg_ids = test.eeg_id.unique()\n\nfor i, eeg_id in tqdm(enumerate(eeg_ids)):  \n    # Save EEG to Python dictionary of numpy arrays\n    eeg_path = data_root + str(eeg_id) + \".parquet\"\n    data = eeg_from_parquet(eeg_path, display=i<visualize)              \n    all_eegs[eeg_id] = data\n    \n    if i == visualize:\n        if CREATE_EEGS:\n            print(f'Processing {test.eeg_id.nunique()} eeg parquets... ',end='')\n        else:\n            print(f'Reading {len(eeg_ids)} eeg NumPys from disk.')\n            break\n            \nif CREATE_EEGS: \n    np.save('eegs', all_eegs)\nelse:\n    all_eegs = np.load('/kaggle/input/brain-eegs/eegs.npy',allow_pickle=True).item()\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:20.212588Z","iopub.execute_input":"2024-04-08T12:45:20.212961Z","iopub.status.idle":"2024-04-08T12:45:21.340250Z","shell.execute_reply.started":"2024-04-08T12:45:20.212928Z","shell.execute_reply":"2024-04-08T12:45:21.339174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader, Dataset\nfrom typing import Dict, List\n\n\nclass EEGDataset(Dataset):\n    def __init__(\n        self, df: pd.DataFrame, config, mode: str = 'train',\n        eegs: Dict[int, np.ndarray] = all_eegs, downsample: int = None\n    ): \n        self.df = df\n        self.config = config\n        self.mode = mode\n        self.eegs = eegs\n        self.downsample = downsample\n        \n    def __len__(self):\n        \"\"\"\n        Length of dataset.\n        \"\"\"\n        return len(self.df)\n        \n    def __getitem__(self, index):\n        \"\"\"\n        Get one item.\n        \"\"\"\n        if self.mode != 'test': \n            X, y_prob = self.__data_generation(index)\n        else: \n            X = self.__data_generation(index)\n        if self.downsample is not None:\n            X = X[::self.downsample,:]\n        if self.mode != 'test': \n            output = {\n                \"eeg\": torch.tensor(X, dtype=torch.float32),\n                \"labels\": torch.tensor(y_prob, dtype=torch.float32)\n            }\n        else: \n            output = {\n                \"eeg\": torch.tensor(X, dtype=torch.float32),\n            }    \n        return output\n                        \n    def __data_generation(self, index):\n        row = self.df.iloc[index]\n        X = np.zeros((10_000, 8), dtype='float32')\n        y = np.zeros(6, dtype='float32')\n        data = self.eegs[row.eeg_id]\n\n        # === Feature engineering ===\n        X[:,0] = data[:,feature_to_index['Fp1']] - data[:,feature_to_index['T3']]\n        X[:,1] = data[:,feature_to_index['T3']] - data[:,feature_to_index['O1']]\n\n        X[:,2] = data[:,feature_to_index['Fp1']] - data[:,feature_to_index['C3']]\n        X[:,3] = data[:,feature_to_index['C3']] - data[:,feature_to_index['O1']]\n\n        X[:,4] = data[:,feature_to_index['Fp2']] - data[:,feature_to_index['C4']]\n        X[:,5] = data[:,feature_to_index['C4']] - data[:,feature_to_index['O2']]\n\n        X[:,6] = data[:,feature_to_index['Fp2']] - data[:,feature_to_index['T4']]\n        X[:,7] = data[:,feature_to_index['T4']] - data[:,feature_to_index['O2']]\n\n        # === Standarize ===\n        X = np.clip(X,-1024, 1024)\n        X = np.nan_to_num(X, nan=0) / 32.0\n\n        # === Butter Low-pass Filter ===\n        X = butter_lowpass_filter(X)\n        if self.mode != 'test':\n            y_prob = row[self.config.target_cols].values.astype(np.float32)\n            return X, y_prob\n        else: \n            return X","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:21.341786Z","iopub.execute_input":"2024-04-08T12:45:21.342163Z","iopub.status.idle":"2024-04-08T12:45:21.362838Z","shell.execute_reply.started":"2024-04-08T12:45:21.342128Z","shell.execute_reply":"2024-04-08T12:45:21.361697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#config \n\nfeature_to_index = {x:y for x,y in zip(eeg_features, range(len(eeg_features)))}\n\n\nclass CFG:\n    wandb = False\n    debug = False\n    train=True\n    apex=True\n    visualize=True\n    stage1_pop1=True\n    stage2_pop2=False\n    scheduler='CosineAnnealingWarmRestarts' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts','OneCycleLR']\n    # CosineAnnealingLR params\n    cosanneal_params={\n        'T_max':6,\n        'eta_min':1e-5,\n        'last_epoch':-1\n    }\n    #ReduceLROnPlateau params\n    reduce_params={\n        'mode':'min',\n        'factor':0.2,\n        'patience':4,\n        'eps':1e-6,\n        'verbose':True\n    }\n    # CosineAnnealingWarmRestarts params\n    cosanneal_res_params={\n        'T_0':20,\n        'eta_min':1e-6,\n        'T_mult':1,\n        'last_epoch':-1\n    }\n    print_freq=50\n    num_workers = 1\n    model_name = 'resnet1d_gru'\n    optimizer='Adan'\n    epochs = 10\n    factor = 0.9\n    patience = 2\n    eps = 1e-6\n    lr = 1e-3\n    min_lr = 1e-6\n    in_channels = 8\n    batch_size = 64\n    weight_decay = 1e-2\n    batch_scheduler = True\n    gradient_accumulation_steps = 1\n    max_grad_norm = 1e7\n    seed = 2024\n    target_cols = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\n    target_size = 6\n    pred_cols = ['pred_seizure_vote', 'pred_lpd_vote', 'pred_gpd_vote', 'pred_lrda_vote', 'pred_grda_vote', 'pred_other_vote']\n    n_fold = 2\n    trn_fold = [0, 1]\n    PATH = '/kaggle/input/hms-harmful-brain-activity-classification/'\n    data_root = \"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/\"\n    raw_eeg_path = \"/kaggle/input/brain-eegs/eegs.npy\"","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:21.365934Z","iopub.execute_input":"2024-04-08T12:45:21.366280Z","iopub.status.idle":"2024-04-08T12:45:21.379706Z","shell.execute_reply.started":"2024-04-08T12:45:21.366254Z","shell.execute_reply":"2024-04-08T12:45:21.378610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from scipy.signal import butter, lfilter\n\ndef quantize_data(data, classes):\n    mu_x = mu_law_encoding(data, classes)\n    return mu_x#quantized\n\ndef mu_law_encoding(data, mu):\n    mu_x = np.sign(data) * np.log(1 + mu * np.abs(data)) / np.log(mu + 1)\n    return mu_x\n\ndef mu_law_expansion(data, mu):\n    s = np.sign(data) * (np.exp(np.abs(data) * np.log(mu + 1)) - 1) / mu\n    return s\n\ndef butter_lowpass_filter(data, cutoff_freq=20, sampling_rate=200, order=4):\n    nyquist = 0.5 * sampling_rate\n    normal_cutoff = cutoff_freq / nyquist\n    b, a = butter(order, normal_cutoff, btype='low', analog=False)\n    filtered_data = lfilter(b, a, data, axis=0)\n    return filtered_data","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:21.381106Z","iopub.execute_input":"2024-04-08T12:45:21.381502Z","iopub.status.idle":"2024-04-08T12:45:22.579692Z","shell.execute_reply.started":"2024-04-08T12:45:21.381464Z","shell.execute_reply":"2024-04-08T12:45:22.578683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#test = pd.concat([test, test, test, test, test, test, test, test, test, test, test])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = EEGDataset(test, CFG, mode='test')","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:22.580957Z","iopub.execute_input":"2024-04-08T12:45:22.581425Z","iopub.status.idle":"2024-04-08T12:45:22.586144Z","shell.execute_reply.started":"2024-04-08T12:45:22.581396Z","shell.execute_reply":"2024-04-08T12:45:22.585155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataloader = DataLoader(\n    test_df,\n    batch_size=1,\n    shuffle=False,\n    num_workers=CFG.num_workers, pin_memory=True, drop_last=True\n)\noutput = test_df[0]\nX = output[\"eeg\"]\nprint(f\"X shape: {X.shape}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:44.687800Z","iopub.execute_input":"2024-04-08T12:45:44.688642Z","iopub.status.idle":"2024-04-08T12:45:44.700052Z","shell.execute_reply.started":"2024-04-08T12:45:44.688607Z","shell.execute_reply":"2024-04-08T12:45:44.698798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\n\ndevice='cuda'\ndef clear_gpu_memory():\n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:45.438221Z","iopub.execute_input":"2024-04-08T12:45:45.438923Z","iopub.status.idle":"2024-04-08T12:45:45.443351Z","shell.execute_reply.started":"2024-04-08T12:45:45.438881Z","shell.execute_reply":"2024-04-08T12:45:45.442220Z"},"_kg_hide-output":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn.functional as F\nimport gc\n\nmodel.eval()\nmodel.to(device)\nsub=pd.DataFrame()\nwith torch.no_grad():\n    for batch in test_dataloader: \n        batch = batch['eeg'].to(device)\n        output = model(batch)\n        prob = F.softmax(output[0], dim=1) #dim for torch \n        print(prob)\n        s = pd.DataFrame(prob.cpu().numpy())\n        sub=pd.concat([sub, s])\n        clear_gpu_memory()\n        del output\n\nf_sub = pd.concat([test.eeg_id.reset_index(drop=True), sub.reset_index(drop=True)],axis=1)\nf_sub.columns = ['eeg_id', 'seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\nf_sub.to_csv('submission.csv', index=False)\n#a = pd.read_csv('/kaggle/working/submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-04-08T12:45:45.467584Z","iopub.execute_input":"2024-04-08T12:45:45.467933Z","iopub.status.idle":"2024-04-08T12:45:46.946907Z","shell.execute_reply.started":"2024-04-08T12:45:45.467906Z","shell.execute_reply":"2024-04-08T12:45:46.945620Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# here is the code for two-head model ","metadata":{}},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2024-04-08T11:39:28.219559Z","iopub.execute_input":"2024-04-08T11:39:28.220577Z","iopub.status.idle":"2024-04-08T11:39:28.227686Z","shell.execute_reply.started":"2024-04-08T11:39:28.220544Z","shell.execute_reply":"2024-04-08T11:39:28.226788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}