{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392733,"sourceType":"datasetVersion","datasetId":4297749},{"sourceId":7801925,"sourceType":"datasetVersion","datasetId":4568467},{"sourceId":164443259,"sourceType":"kernelVersion"},{"sourceId":166906548,"sourceType":"kernelVersion"},{"sourceId":168749976,"sourceType":"kernelVersion"}],"dockerImageVersionId":30665,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# About this notebook\n\nStrat based folds on number of evaluators and patient id in 5 bins\n\n## Models notebook\n\n* GRU only 0.6976942163420315 CV\n* GRU stage 2 0.7840440699870758\n* GRU new feats stage 2 0.6932664658512089 LB 0.33\n* GRU new feats stage 2 0.6074094803600332 LB 0.33","metadata":{}},{"cell_type":"code","source":"! pip install /kaggle/input/wheel-albumentation/albumentations-1.4.0-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:53:34.127848Z","iopub.execute_input":"2024-03-26T15:53:34.128488Z","iopub.status.idle":"2024-03-26T15:54:07.771345Z","shell.execute_reply.started":"2024-03-26T15:53:34.128463Z","shell.execute_reply":"2024-03-26T15:54:07.770358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Libraries","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nfrom glob import glob\nimport sys\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom typing import Dict, List\nfrom scipy.stats import entropy\nfrom scipy.signal import butter, lfilter, freqz,iirnotch, filtfilt\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nsys.path.append('/kaggle/input/kaggle-kl-div')\nfrom kaggle_kl_div import score\nimport numpy as np\nimport pandas as pd\nfrom sklearn import preprocessing\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.metrics import accuracy_score, log_loss\nfrom tqdm.auto import tqdm\nfrom functools import partial\nimport cv2\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport pytorch_lightning as pl\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD, AdamW\nimport torchvision.models as models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau, OneCycleLR, CosineAnnealingLR, CosineAnnealingWarmRestarts\nfrom sklearn.preprocessing import LabelEncoder\nfrom torchvision.transforms import v2\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.model_selection import train_test_split\nimport albumentations as A\nfrom albumentations import (Compose, Normalize, Resize, RandomResizedCrop, HorizontalFlip, VerticalFlip, ShiftScaleRotate, Transpose)\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\nimport timm\nimport warnings \nwarnings.filterwarnings('ignore')\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nfrom matplotlib import pyplot as plt\nimport joblib\nos.environ['CUDA_VISIBLE_DEVICES'] = \"0,1\"\nVERSION=2\nENSEMBLE = False","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:07.773119Z","iopub.execute_input":"2024-03-26T15:54:07.773419Z","iopub.status.idle":"2024-03-26T15:54:19.338561Z","shell.execute_reply.started":"2024-03-26T15:54:07.773392Z","shell.execute_reply":"2024-03-26T15:54:19.337632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"class CFG:\n    PATH = '/kaggle/resnet50_gru_pathhms-harmful-brain-activity-classification/'\n    test_eeg = \"/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/\"\n    test_csv = \"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\"\n    test_spectrograms = \"/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/\"    \n    SparK = False\n    FREEZE = False\n    seed = 2024\n    wavenet_gru_path = \"/kaggle/input/hms-wavenet-gru-train-v3-diff-feats/pop_1_weight_oof/\"#\"/kaggle/input/hms-wavenet-gru-train-v3/pop_1_weight_oof/\"\n    tf_effnetb0_ns_path = \"/kaggle/input/new-tf-effnetb0-ns-weights-v1/pop_2_weight_oof/\"\n    resnet_gru_in_channels = 8\n    wavenet_gru_in_channels = 1\n    target_size = 6\n    filter_order = 5\n    lowcut = 0.7  # 0.85  \n    highcut = 20  # 25.0\n    nyquist_freq = 0.5 * 200\n    low_cut_freq_normalized = lowcut / nyquist_freq\n    high_cut_freq_normalized = highcut / nyquist_freq\n    batch_size = 32\n    TTA=True\n    num_workers = 1\n    target_cols = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\n    pred_cols = ['pred_seizure_vote', 'pred_lpd_vote', 'pred_gpd_vote', 'pred_lrda_vote', 'pred_grda_vote', 'pred_other_vote'] ","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:19.339796Z","iopub.execute_input":"2024-03-26T15:54:19.340252Z","iopub.status.idle":"2024-03-26T15:54:19.347530Z","shell.execute_reply.started":"2024-03-26T15:54:19.340213Z","shell.execute_reply":"2024-03-26T15:54:19.346595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"def compute_cv(oof_df, txt):\n    \n    label_cols = CFG.target_cols\n    gt = oof_df[[\"eeg_id\"] + CFG.target_cols]\n    gt.sort_values(by=\"eeg_id\", inplace=True)\n    gt.reset_index(inplace=True, drop=True)\n\n    preds = oof_df[[\"eeg_id\"] + CFG.pred_cols]\n    preds.columns = [\"eeg_id\"] + CFG.target_cols\n    preds.sort_values(by=\"eeg_id\", inplace=True)\n    preds.reset_index(inplace=True, drop=True)\n\n    y_trues = gt[CFG.target_cols]\n    y_preds = preds[CFG.target_cols]\n\n    oof = pd.DataFrame(y_preds.copy())\n    oof['id'] = np.arange(len(oof))\n\n    true = pd.DataFrame(y_trues.copy())\n    true['id'] = np.arange(len(true))\n\n    cv = score(solution=true, submission=oof, row_id_column_name='id')\n    print(f'CV Stage1 Score with {txt} =',cv)\n    \ndef goto_conversion(listOfOdds, total = 1, eps = 1e-6, isAmericanOdds = False):\n\n    #Convert American Odds to Decimal Odds\n    if isAmericanOdds:\n        for i in range(len(listOfOdds)):\n            currOdds = listOfOdds[i]\n            isNegativeAmericanOdds = currOdds < 0\n            if isNegativeAmericanOdds:\n                currDecimalOdds = 1 + (100/(currOdds*-1))\n            else: #Is non-negative American Odds\n                currDecimalOdds = 1 + (currOdds/100)\n            listOfOdds[i] = currDecimalOdds\n\n    #Error Catchers\n    #if len(listOfOdds) < 2:\n        #raise ValueError('len(listOfOdds) must be >= 2')\n    if any(x < 1 for x in listOfOdds):\n        raise ValueError('All odds must be >= 1, set isAmericanOdds parameter to True if using American Odds')\n\n    #Computation\n    listOfProbabilities = [1/x for x in listOfOdds] #initialize probabilities using inverse odds\n    listOfSe = [pow((x-x**2)/x,0.5) for x in listOfProbabilities] #compute the standard error (SE) for each probability\n    step = (sum(listOfProbabilities) - total)/sum(listOfSe) #compute how many steps of SE the probabilities should step back by\n    outputListOfProbabilities = [min(max(x - (y*step),eps),1) for x,y in zip(listOfProbabilities, listOfSe)]\n    return outputListOfProbabilities","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:19.350324Z","iopub.execute_input":"2024-03-26T15:54:19.350893Z","iopub.status.idle":"2024-03-26T15:54:19.366196Z","shell.execute_reply.started":"2024-03-26T15:54:19.350858Z","shell.execute_reply":"2024-03-26T15:54:19.365447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"def Fix_EEG(eeg):\n    msg = 'False'   # no errors exist in the row\n    length_eeg = len(eeg)\n\n    # Check if the length of eeg is less than 10000\n    if length_eeg < 10000:\n\n        rows_to_add = 10000 - length_eeg\n        # Create a DataFrame with noise\n        noise_df = pd.DataFrame(np.random.normal(loc=0.00001, scale=0.0001, size=(rows_to_add, eeg.shape[1])), \n                                columns=eeg.columns)\n        # Concatenate the original DataFrame with the noise DataFrame\n        eeg = pd.concat([eeg, noise_df], ignore_index=True)\n        msg = 'Not 10K'\n\n    # Convert inf values to NaN\n    eeg.replace([np.inf, -np.inf], np.nan, inplace=True)\n\n    # Check for NaN values\n    has_nan = eeg.isna().any().any()\n    if has_nan:\n        msg = 'NAN Present'\n        # Replace NaN values with a specified method, e.g., ffill, bfill, or a constant value\n        # e.g., eeg.fillna(method='ffill', inplace=True)\n        eeg.fillna(eeg.mean(), inplace=True)  # Replacing NaNs with the mean of each column\n\n\n    return eeg, msg\n\ndef eeg_from_parquet(parquet_path: str) -> 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    #eeg, msg = Fix_EEG(eeg)\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   \n    return data\n\nimport pywt, librosa\n\nUSE_WAVELET = None \n\nNAMES = ['LL','LP','RP','RR']\n\nFEATS = [['Fp1','F7','T3','T5','O1'],\n         ['Fp1','F3','C3','P3','O1'],\n         ['Fp2','F8','T4','T6','O2'],\n         ['Fp2','F4','C4','P4','O2']]\n\n# DENOISE FUNCTION\ndef maddest(d, axis=None):\n    return np.mean(np.absolute(d - np.mean(d, axis)), axis)\n\ndef denoise(x, wavelet='haar', level=1):    \n    coeff = pywt.wavedec(x, wavelet, mode=\"per\")\n    sigma = (1/0.6745) * maddest(coeff[-level])\n\n    uthresh = sigma * np.sqrt(2*np.log(len(x)))\n    coeff[1:] = (pywt.threshold(i, value=uthresh, mode='hard') for i in coeff[1:])\n\n    ret=pywt.waverec(coeff, wavelet, mode='per')\n    \n    return ret\n\ndef seed_everything(seed: int):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed) \n    \n    \ndef sep():\n    print(\"-\"*100)\n\n    \ntarget_preds = [x + \"_pred\" for x in ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']]\nlabel_to_num = {'Seizure': 0, 'LPD': 1, 'GPD': 2, 'LRDA': 3, 'GRDA': 4, 'Other':5}\nnum_to_label = {v: k for k, v in label_to_num.items()}\nseed_everything(CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:19.367185Z","iopub.execute_input":"2024-03-26T15:54:19.367465Z","iopub.status.idle":"2024-03-26T15:54:19.452707Z","shell.execute_reply.started":"2024-03-26T15:54:19.367443Z","shell.execute_reply":"2024-03-26T15:54:19.451974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load data","metadata":{}},{"cell_type":"code","source":"test_df = pd.read_csv(CFG.test_csv)\nprint(f\"Test dataframe shape is: {test_df.shape}\")\ntest_df.head()\n\neeg_parquet_paths = glob(CFG.test_eeg+ \"*.parquet\")\neeg_df = pd.read_parquet(eeg_parquet_paths[0])\neeg_features = eeg_df.columns\nprint(f'There are {len(eeg_features)} raw eeg features')\nprint(list(eeg_features))\neeg_features =  ['Fp1', 'Fp2', 'F3', 'F4', 'F7', 'F8', 'C3', 'C4',  'T3', 'T4', 'P3', 'P4', 'O1', 'O2', 'T5', 'T6'] # 17 feats CZ\nfeature_to_index = {x:y for x,y in zip(eeg_features, range(len(eeg_features)))}","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:19.453656Z","iopub.execute_input":"2024-03-26T15:54:19.454103Z","iopub.status.idle":"2024-03-26T15:54:19.648634Z","shell.execute_reply.started":"2024-03-26T15:54:19.454080Z","shell.execute_reply":"2024-03-26T15:54:19.647728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\nCREATE_EEGS = False\nall_eegs = {}\nvisualize = 1\neeg_paths = glob(CFG.test_eeg + \"*.parquet\")\neeg_ids = test_df.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 = CFG.test_eeg + str(eeg_id) + \".parquet\"\n    data = eeg_from_parquet(eeg_path)              \n    all_eegs[eeg_id] = data","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:19.649750Z","iopub.execute_input":"2024-03-26T15:54:19.650047Z","iopub.status.idle":"2024-03-26T15:54:19.686606Z","shell.execute_reply.started":"2024-03-26T15:54:19.650023Z","shell.execute_reply":"2024-03-26T15:54:19.685598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# READ ALL SPECTROGRAMS\nfiles2 = os.listdir(CFG.test_spectrograms)\nprint(f'There are {len(files2)} test spectrogram parquets')\n    \nall_spectrograms = {}\nfor i,f in enumerate(files2):\n    if i%100==0: print(i,', ',end='')\n    tmp = pd.read_parquet(f'{CFG.test_spectrograms}{f}')\n    name = int(f.split('.')[0])\n    all_spectrograms[name] = tmp.iloc[:,1:].values\n    \n# RENAME FOR DATALOADER\ntest_df = test_df.rename({'spectrogram_id':'spec_id'},axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:19.687954Z","iopub.execute_input":"2024-03-26T15:54:19.688290Z","iopub.status.idle":"2024-03-26T15:54:19.738798Z","shell.execute_reply.started":"2024-03-26T15:54:19.688259Z","shell.execute_reply":"2024-03-26T15:54:19.737984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# READ ALL EEG SPECTROGRAMS\nDISPLAY = 1\nEEG_IDS2 = test_df.eeg_id.unique()\ncustom_eegs = {}\n\nprint('Converting Test EEG to Spectrograms...'); print() \nfrom scipy.signal import butter, lfilter\n\ndef butter_lowpass_filter(data, cutoff_freq: int = 20, sampling_rate: int = 200, order: int = 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-03-26T15:54:19.739992Z","iopub.execute_input":"2024-03-26T15:54:19.740341Z","iopub.status.idle":"2024-03-26T15:54:19.747829Z","shell.execute_reply.started":"2024-03-26T15:54:19.740316Z","shell.execute_reply":"2024-03-26T15:54:19.746944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class EEGDataset(Dataset):\n    def __init__(\n        self, df: pd.DataFrame, config, mode: str = 'train',\n        eegs: Dict[int, np.ndarray] = all_eegs, downsample: int = 5\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        X, y = self.__data_generation(index)\n        if self.downsample is not None:\n            X = X[::self.downsample,:]\n        output = {\n            \"eeg\": torch.tensor(X, dtype=torch.float32),\n            \"labels\": torch.tensor(y, dtype=torch.float32)\n        }\n        return output\n                        \n    def __data_generation(self, index):\n        row = self.df.iloc[index]\n        y = np.zeros(6, dtype='float32')\n        data = self.eegs[row.eeg_id]\n\n        # === Feature engineering ===\n        signal_pairs = [\n            ('Fp1', 'F7'),\n            ('Fp2', 'F8'),\n            ('F7', 'T3'),\n            ('F8', 'T4'),\n            ('T3', 'T5'),\n            ('T4', 'T6'),\n            ('T5', 'O1'),\n            ('T6', 'O2'),\n            #('T3', 'C3'),\n            ('C4', 'T4'),\n            #('C3', 'Cz'),\n            #('Cz', 'C4'),\n            ('Fp1', 'F3'),\n            ('Fp2', 'F4'),\n            ('F3', 'C3'),\n            ('F4', 'C4'), \n            #('C3', 'P3'),\n            #('C4', 'P4'),\n            ('P3', 'O1'),\n            ('P4', 'O2'),\n        ]\n        X = np.zeros((10_000, len(signal_pairs)), dtype='float32')\n        #X = np.zeros((eeg.shape[0], len(signal_pairs)))\n        for i, (signal1, signal2) in enumerate(signal_pairs):\n            X[:, i] = data[:, feature_to_index[signal1]] - data[:, feature_to_index[signal2]]\n\n        # === Standarize ===\n        X = np.clip(X,-1024, 1024)\n        X = np.nan_to_num(X, nan=0) / 32.0\n        # === Notch and Butter Low-pass Filter ===\n        # Bandpass and notch filter\n        bandpass_coefficients = butter(CFG.filter_order, [CFG.low_cut_freq_normalized, CFG.high_cut_freq_normalized], btype='band')\n        notch_coefficients = iirnotch(w0=60, Q=30, fs=200)\n        # Filter bandpass and notch\n        X = filtfilt(*notch_coefficients, X, axis=0)\n        X = filtfilt(*bandpass_coefficients, X, axis=0)\n          \n        if self.mode != 'test':\n            y = row[self.config.target_cols].values.astype(np.float32)\n        return X.copy(), y ","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:19.750638Z","iopub.execute_input":"2024-03-26T15:54:19.750894Z","iopub.status.idle":"2024-03-26T15:54:19.766445Z","shell.execute_reply.started":"2024-03-26T15:54:19.750871Z","shell.execute_reply":"2024-03-26T15:54:19.765621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Wavenet GRU Inference","metadata":{}},{"cell_type":"markdown","source":"#### Model","metadata":{}},{"cell_type":"code","source":"class SeqPool(nn.Module):\n    def __init__(self, emb_dim=192):\n        super().__init__()\n        self.dense = nn.Linear(emb_dim, 1)\n        self.softmax = nn.Softmax(dim=-1)\n\n    def forward(self, x):\n        bs, seq_len, emb_dim = x.shape\n        identity = x\n        x = self.dense(x)\n        x = x.permute(0, 2, 1)\n        x = self.softmax(x)\n        x = x @ identity\n        x = x.reshape(x.shape[0], -1)\n        return x\n    \n\n\nclass Wave_Block(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int, dilation_rates: int, kernel_size: int = 3):\n        \"\"\"\n        WaveNet building block.\n        :param in_channels: number of input channels.\n        :param out_channels: number of output channels.\n        :param dilation_rates: how many levels of dilations are used.\n        :param kernel_size: size of the convolving kernel.\n        \"\"\"\n        super(Wave_Block, self).__init__()\n        self.num_rates = dilation_rates\n        self.convs = nn.ModuleList()\n        self.filter_convs = nn.ModuleList()\n        self.gate_convs = nn.ModuleList()\n        self.convs.append(nn.Conv1d(in_channels, out_channels, kernel_size=1, bias=True))\n        \n        \n        dilation_rates = [2 ** i for i in range(dilation_rates)]\n        for dilation_rate in dilation_rates:\n            self.filter_convs.append(\n                nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size,\n                          padding=int((dilation_rate*(kernel_size-1))/2), dilation=dilation_rate))\n            self.gate_convs.append(\n                nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size,\n                          padding=int((dilation_rate*(kernel_size-1))/2), dilation=dilation_rate))\n            self.convs.append(nn.Conv1d(out_channels, out_channels, kernel_size=1, bias=True))\n        \n        for i in range(len(self.convs)):\n            nn.init.xavier_uniform_(self.convs[i].weight, gain=nn.init.calculate_gain('relu'))\n            nn.init.zeros_(self.convs[i].bias)\n\n        for i in range(len(self.filter_convs)):\n            nn.init.xavier_uniform_(self.filter_convs[i].weight, gain=nn.init.calculate_gain('relu'))\n            nn.init.zeros_(self.filter_convs[i].bias)\n\n        for i in range(len(self.gate_convs)):\n            nn.init.xavier_uniform_(self.gate_convs[i].weight, gain=nn.init.calculate_gain('relu'))\n            nn.init.zeros_(self.gate_convs[i].bias)\n\n    def forward(self, x):\n        x = self.convs[0](x)\n        res = x\n        for i in range(self.num_rates):\n            tanh_out = torch.tanh(self.filter_convs[i](x))\n            sigmoid_out = torch.sigmoid(self.gate_convs[i](x))\n            x = tanh_out * sigmoid_out\n            x = self.convs[i + 1](x) \n            res = res + x\n            \n        return res\n    \nclass WaveNet(nn.Module):\n    def __init__(self, input_channels: int = 1, kernel_size: int = 3):\n        super(WaveNet, self).__init__()\n        self.wave_blocks = nn.Sequential(\n                Wave_Block(input_channels, 8, 12, kernel_size),\n                Wave_Block(8, 16, 8, kernel_size),\n                Wave_Block(16, 32, 4, kernel_size),\n                Wave_Block(32, 64, 1, kernel_size)\n                )\n        self.gru = nn.GRU(input_size=64, hidden_size=128, num_layers=1, bidirectional=True)\n        #self.gru = nn.LSTM(input_size=64, hidden_size=128, num_layers=1, bidirectional=True)\n        self.seqpool = SeqPool(128*2)\n        \n    def forward(self, x: torch.Tensor) -> torch.Tensor: \n        x = x.permute(0, 2, 1)\n        output = self.wave_blocks(x)\n        out, _ = self.gru(output.permute(0, 2, 1))\n        out = self.seqpool(out)\n        return out\n    \nclass CustomModel(nn.Module):\n    def __init__(self):\n        super(CustomModel, self).__init__()\n        self.model = WaveNet()\n        self.dropout = 0.0\n        self.head = nn.Sequential(\n            nn.Linear(1024, 128),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Dropout(self.dropout),\n            nn.Linear(128, 6)\n        )\n        #CUDA_VISIBLE_DEVICES=7 python3 ./2_train.py \n        #--project-name alexnet_v4_raw_sincnet --model alexnet_v4 --task-type binary --optim adam --window-size 4 --window-shift 1 --eeg-type bipolar --enc-model sincnet --sincnet-bandnum 7 --binary-sampler-type 6types --binary-target-groups 2 --epoch 8 --batch-size 32 --seizure-wise-eval-for-binary True\n        \"\"\"class args:\n            sincnet_layer_num = 3  # Number of SincNet convolutional layers\n            window_size_sig = 2  # Length of input signal\n            sincnet_kernel_size = 7  # Kernel size for the first SincConv layer\n            cnn_channel_sizes = [8, 10, 16]  # Number of output channels for each CNN layer\n            sincnet_stride = 1  # Stride for the first SincConv layer \n            sincnet_input_normalize = \"none\"  # Input normalization method  choices=[\"none\",\"layernorm\",\"batchnorm\"])\"\"\"\n        #self.sincnet  = PSD_FEATURE2()#SINCNET_FEATURE(args=args,  num_eeg_channel=8) # padding to 0 or (kernel_size-1)//2  \n        # This is a dictionary that maps feature pairs to their index in x\n        self.index_dict = {('Fp1', 'F7'): 0, ('Fp2', 'F8'): 1, ('F7', 'T3'): 2, ('F8', 'T4'): 3,('T3', 'T5'): 4,\n                      ('T4', 'T6'): 5, ('T5', 'O1'): 6,('T6', 'O2'): 7, ('C4', 'T4'): 8, ('Fp1', 'F3'): 9,('Fp2', 'F4'): 10,('F3', 'C3'): 11,\n                      ('F4', 'C4'): 12 ,('P3', 'O1'): 13,('P4', 'O2'): 14}\n    def extract_features(self, x):\n        # Left part\n        x1 = self.model(x[:, :, self.index_dict[('Fp1', 'F7')]:self.index_dict[('Fp1', 'F7')]+1])\n        x2 = self.model(x[:, :, self.index_dict[('F7', 'T3')]:self.index_dict[('F7', 'T3')]+1])\n        x3 = self.model(x[:, :, self.index_dict[('T3', 'T5')]:self.index_dict[('T3', 'T5')]+1])\n        x4 = self.model(x[:, :, self.index_dict[('T5', 'O1')]:self.index_dict[('T5', 'O1')]+1])\n        #x5 = self.model(x[:, :, self.index_dict[('T3', 'C3')]:self.index_dict[('T3', 'C3')]+1])\n        x6 = self.model(x[:, :, self.index_dict[('Fp1', 'F3')]:self.index_dict[('Fp1', 'F3')]+1])\n        x7 = self.model(x[:, :, self.index_dict[('F3', 'C3')]:self.index_dict[('F3', 'C3')]+1])\n        #x8 = self.model(x[:, :, self.index_dict[('C3', 'P3')]:self.index_dict[('C3', 'P3')]+1])\n        x9 = self.model(x[:, :, self.index_dict[('P3', 'O1')]:self.index_dict[('P3', 'O1')]+1])\n        z1 = torch.mean(torch.stack([x1, x2, x3, x4, x6, x7,   x9]), dim=0)\n\n        # Right part\n        x1 = self.model(x[:, :, self.index_dict[('Fp2', 'F8')]:self.index_dict[('Fp2', 'F8')]+1])\n        x2 = self.model(x[:, :, self.index_dict[('F8', 'T4')]:self.index_dict[('F8', 'T4')]+1])\n        x3 = self.model(x[:, :, self.index_dict[('T4', 'T6')]:self.index_dict[('T4', 'T6')]+1])\n        x4 = self.model(x[:, :, self.index_dict[('T6', 'O2')]:self.index_dict[('T6', 'O2')]+1])\n        x5 = self.model(x[:, :, self.index_dict[('C4', 'T4')]:self.index_dict[('C4', 'T4')]+1])\n        x6 = self.model(x[:, :, self.index_dict[('Fp2', 'F4')]:self.index_dict[('Fp2', 'F4')]+1])\n        x7 = self.model(x[:, :, self.index_dict[('F4', 'C4')]:self.index_dict[('F4', 'C4')]+1])\n        #x8 = self.model(x[:, :, self.index_dict[('C4', 'P4')]:self.index_dict[('C4', 'P4')]+1])\n        x9 = self.model(x[:, :, self.index_dict[('P4', 'O2')]:self.index_dict[('P4', 'O2')]+1])\n        z2 = torch.mean(torch.stack([x1, x2, x3, x4, x5, x6, x7,   x9]), dim=0)\n\n        # Front part\n        x1 = self.model(x[:, :, self.index_dict[('Fp1', 'F7')]:self.index_dict[('Fp1', 'F7')]+1])\n        x2 = self.model(x[:, :, self.index_dict[('Fp2', 'F8')]:self.index_dict[('Fp2', 'F8')]+1])\n        x3 = self.model(x[:, :, self.index_dict[('Fp1', 'F3')]:self.index_dict[('Fp1', 'F3')]+1])\n        x4 = self.model(x[:, :, self.index_dict[('Fp2', 'F4')]:self.index_dict[('Fp2', 'F4')]+1])\n        z3 = torch.mean(torch.stack([x1, x2, x3, x4]), dim=0)\n\n        # Back part\n        x1 = self.model(x[:, :, self.index_dict[('T5', 'O1')]:self.index_dict[('T5', 'O1')]+1])\n        x2 = self.model(x[:, :, self.index_dict[('T6', 'O2')]:self.index_dict[('T6', 'O2')]+1])\n        x3 = self.model(x[:, :, self.index_dict[('P3', 'O1')]:self.index_dict[('P3', 'O1')]+1])\n        x4 = self.model(x[:, :, self.index_dict[('P4', 'O2')]:self.index_dict[('P4', 'O2')]+1])\n        z4 = torch.mean(torch.stack([x1, x2, x3, x4]), dim=0)\n\n        y = torch.cat([z1, z2, z3, z4], dim=1)\n        \n        return y\n    def forward(self, x: torch.Tensor):\n        \"\"\"\n        Forwward pass.\n        \"\"\" \n        y = self.extract_features(x)\n        y = self.head(y)\n        \n        return y","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:19.767621Z","iopub.execute_input":"2024-03-26T15:54:19.767912Z","iopub.status.idle":"2024-03-26T15:54:19.811423Z","shell.execute_reply.started":"2024-03-26T15:54:19.767890Z","shell.execute_reply":"2024-03-26T15:54:19.810465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Inference Function","metadata":{}},{"cell_type":"code","source":"def inference_function(test_loader, model, device):\n    model.eval() # set model in evaluation mode\n    softmax = nn.Softmax(dim=1)\n    prediction_dict = {}\n    preds = []\n    with tqdm(test_loader, unit=\"test_batch\", desc='Inference') as tqdm_test_loader:\n        for step, batch in enumerate(tqdm_test_loader):\n            X = batch.pop(\"eeg\").to(device) # send inputs to `device`\n            batch_size = X.size(0)\n            with torch.no_grad():\n                y_preds = model(X) # forward propagation pass\n            y_preds = softmax(y_preds)\n            preds.append(y_preds.to('cpu').numpy()) # save predictions\n                \n    prediction_dict[\"predictions\"] = np.concatenate(preds) # np.array() of shape (fold_size, target_cols)\n    return prediction_dict","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:19.812482Z","iopub.execute_input":"2024-03-26T15:54:19.812945Z","iopub.status.idle":"2024-03-26T15:54:19.826752Z","shell.execute_reply.started":"2024-03-26T15:54:19.812900Z","shell.execute_reply":"2024-03-26T15:54:19.825961Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Inference","metadata":{}},{"cell_type":"code","source":"wavenet_gru_weights = [x for x in glob(CFG.wavenet_gru_path + '*.pth')]\n\nwavenet_gru_preds = []\n\n\nfor model_weight in wavenet_gru_weights:\n    test_dataset = EEGDataset(test_df, CFG, mode='test', downsample=5)\n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=CFG.batch_size,\n        shuffle=False,\n        num_workers=CFG.num_workers,\n        pin_memory=True,\n        drop_last=False\n    )\n    model = CustomModel()\n    checkpoint = torch.load(model_weight, map_location=device)\n    model.load_state_dict(checkpoint[\"model\"])\n    model.to(device)\n    prediction_dict = inference_function(test_loader, model, device)\n    wavenet_gru_preds.append(prediction_dict[\"predictions\"])\n    torch.cuda.empty_cache()\n    gc.collect()\n    \nwavenet_gru_preds = np.array(wavenet_gru_preds)\nwavenet_gru_preds = np.mean(wavenet_gru_preds, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:19.827695Z","iopub.execute_input":"2024-03-26T15:54:19.827985Z","iopub.status.idle":"2024-03-26T15:54:50.072570Z","shell.execute_reply.started":"2024-03-26T15:54:19.827957Z","shell.execute_reply":"2024-03-26T15:54:50.071372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"TARGETS = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\nsub = pd.DataFrame({'eeg_id': test_df.eeg_id.values})\n#CV = 0.6640102859457346  to use  [0.4228774756141195, 0.5771225243858805]\n\nfor i in range(len(TARGETS)):\n    sub[f'{TARGETS[i]}']= wavenet_gru_preds[:, i]\n    \nsub.to_csv(f'submission.csv',index=False)\nprint(f'Submission shape: {sub.shape}')\nsub.head()\n# SANITY CHECK TO CONFIRM PREDICTIONS SUM TO ONE\nsub.iloc[:,-6:].sum(axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:50.074350Z","iopub.execute_input":"2024-03-26T15:54:50.074766Z","iopub.status.idle":"2024-03-26T15:54:50.097522Z","shell.execute_reply.started":"2024-03-26T15:54:50.074726Z","shell.execute_reply":"2024-03-26T15:54:50.096574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub","metadata":{"execution":{"iopub.status.busy":"2024-03-26T15:54:50.098961Z","iopub.execute_input":"2024-03-26T15:54:50.099611Z","iopub.status.idle":"2024-03-26T15:54:50.115907Z","shell.execute_reply.started":"2024-03-26T15:54:50.099578Z","shell.execute_reply":"2024-03-26T15:54:50.114990Z"},"trusted":true},"execution_count":null,"outputs":[]}]}