{"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":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392733,"sourceType":"datasetVersion","datasetId":4297749},{"sourceId":7465251,"sourceType":"datasetVersion","datasetId":4317718}],"dockerImageVersionId":30648,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### 导入库和数据集","metadata":{}},{"cell_type":"code","source":"import pandas as pd, numpy as np, os\nimport matplotlib.pyplot as plt\n\ntrain = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\nprint( train.shape )\ntrain","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:11:34.792219Z","iopub.execute_input":"2024-02-29T10:11:34.793317Z","iopub.status.idle":"2024-02-29T10:11:35.574552Z","shell.execute_reply.started":"2024-02-29T10:11:34.793262Z","shell.execute_reply":"2024-02-29T10:11:35.573372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 创建evalutors列和consensus列将相同EEGid的行合并","metadata":{}},{"cell_type":"code","source":"# 创建新列并计算值\ntrain['total_evaluators'] = train[['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']].sum(axis=1)\n# 将数据框按照 eeg_id 进行分组，并对每个分组进行合并操作\ntrain_grouped_eeg = train.groupby('eeg_id').agg({\n    'seizure_vote': 'sum',\n    'lpd_vote': 'sum',\n    'gpd_vote': 'sum',\n    'lrda_vote': 'sum',\n    'grda_vote': 'sum',\n    'other_vote': 'sum',\n    'eeg_sub_id': 'count',  # 计算每个 eeg_id 对应的行数\n    # 保留其他列的第一行值\n    'eeg_label_offset_seconds': 'first',\n    'spectrogram_id': 'first',\n    'spectrogram_sub_id': 'first',\n    'spectrogram_label_offset_seconds': 'first',\n    'label_id': 'first',\n    'patient_id': 'first',\n    'expert_consensus': 'first',\n    'total_evaluators': 'sum'\n}).reset_index()\n\n# 显示结果\ntrain_grouped_eeg","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:11:35.576596Z","iopub.execute_input":"2024-02-29T10:11:35.576939Z","iopub.status.idle":"2024-02-29T10:11:35.672720Z","shell.execute_reply.started":"2024-02-29T10:11:35.576907Z","shell.execute_reply":"2024-02-29T10:11:35.671647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_grouped_eeg['consensus'] = train_grouped_eeg[['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']].max(axis=1)\ntrain_grouped_eeg","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:11:35.674007Z","iopub.execute_input":"2024-02-29T10:11:35.674341Z","iopub.status.idle":"2024-02-29T10:11:35.703837Z","shell.execute_reply.started":"2024-02-29T10:11:35.674311Z","shell.execute_reply":"2024-02-29T10:11:35.702730Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 丢弃不需要的列","metadata":{}},{"cell_type":"code","source":"train_grouped_eeg.drop(columns=['eeg_sub_id', 'eeg_label_offset_seconds', 'spectrogram_sub_id', 'spectrogram_label_offset_seconds', 'label_id', 'spectrogram_id'], inplace=True)\ntrain_grouped_eeg","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:11:35.705487Z","iopub.execute_input":"2024-02-29T10:11:35.705931Z","iopub.status.idle":"2024-02-29T10:11:35.725629Z","shell.execute_reply.started":"2024-02-29T10:11:35.705896Z","shell.execute_reply":"2024-02-29T10:11:35.724319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 计算一致性eeg_agreements，丢弃低于0.4的数据","metadata":{}},{"cell_type":"code","source":"train_grouped_eeg['row_agreement'] = train_grouped_eeg['consensus']/train_grouped_eeg['total_evaluators']\ntrain_grouped_eeg = train_grouped_eeg[train_grouped_eeg['row_agreement'] >= 0.95]\ntrain_grouped_eeg","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:11:35.728681Z","iopub.execute_input":"2024-02-29T10:11:35.729128Z","iopub.status.idle":"2024-02-29T10:11:35.749915Z","shell.execute_reply.started":"2024-02-29T10:11:35.729098Z","shell.execute_reply":"2024-02-29T10:11:35.748625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 将所有的vote转化为频率的形式","metadata":{}},{"cell_type":"code","source":"# 计算每一列的比例\ntrain_grouped_eeg[['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']] = train_grouped_eeg[['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']].div(train_grouped_eeg['total_evaluators'], axis=0)\n# 显示结果\ntrain_grouped_eeg","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:11:35.751249Z","iopub.execute_input":"2024-02-29T10:11:35.751677Z","iopub.status.idle":"2024-02-29T10:11:35.781197Z","shell.execute_reply.started":"2024-02-29T10:11:35.751643Z","shell.execute_reply":"2024-02-29T10:11:35.779925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 丢弃后三列，将按照eegid聚合的数据集分为大小两个部分","metadata":{}},{"cell_type":"code","source":"# # 筛选出 total_evaluators 列小于 10 的行，存储到 small 数据框中\n# small_eeg = train_grouped_eeg[train_grouped_eeg['total_evaluators'] < 10]\n\n# # 筛选出 total_evaluators 列大于 9 的行，存储到 large 数据框中\n# large_eeg = train_grouped_eeg[train_grouped_eeg['total_evaluators'] > 9]\n\n# train_grouped_eeg = train_grouped_eeg.drop(columns=['total_evaluators', 'consensus', 'row_agreement'])","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:11:35.782758Z","iopub.execute_input":"2024-02-29T10:11:35.783105Z","iopub.status.idle":"2024-02-29T10:11:35.787339Z","shell.execute_reply.started":"2024-02-29T10:11:35.783074Z","shell.execute_reply":"2024-02-29T10:11:35.786421Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# CHOICE TO CREATE OR LOAD EEGS FROM NOTEBOOK VERSION 1\nCREATE_EEGS = False\nTRAIN_MODEL = False\ndf = pd.read_parquet('/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/1000913311.parquet')\nFEATS = df.columns\nprint(f'There are {len(FEATS)} raw eeg features')\nprint( list(FEATS) )\nprint('We will use the following subset of raw EEG features:')\nFEATS = ['Fp1','T3','C3','O1','Fp2','C4','T4','O2']\nFEAT2IDX = {x:y for x,y in zip(FEATS,range(len(FEATS)))}\nprint( list(FEATS) )\ndef eeg_from_parquet(parquet_path, display=False):\n    \n    # EXTRACT MIDDLE 50 SECONDS\n    eeg = pd.read_parquet(parquet_path, columns=FEATS)\n    rows = len(eeg)\n    offset = (rows-10_000)//2\n    eeg = eeg.iloc[offset:offset+10_000]\n    \n    if display: \n        plt.figure(figsize=(10,5))\n        offset = 0\n    \n    # CONVERT TO NUMPY\n    data = np.zeros((10_000,len(FEATS)))\n    for j,col in enumerate(FEATS):\n        \n        # FILL NAN\n        x = eeg[col].values.astype('float32')\n        m = np.nanmean(x)\n        if np.isnan(x).mean()<1: x = np.nan_to_num(x,nan=m)\n        else: x[:] = 0\n            \n        data[:,j] = x\n        \n        if display: \n            if j!=0: offset += x.max()\n            plt.plot(range(10_000),x-offset,label=col)\n            offset -= x.min()\n            \n    if display:\n        plt.legend()\n        name = parquet_path.split('/')[-1]\n        name = name.split('.')[0]\n        plt.title(f'EEG {name}',size=16)\n        plt.show()\n        \n    return data\nall_eegs = {}\nDISPLAY = 4\nEEG_IDS = train_grouped_eeg.eeg_id.unique()\nPATH = '/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/'\n\nfor i,eeg_id in enumerate(EEG_IDS):\n    if (i%100==0)&(i!=0): print(i,', ',end='') \n    \n    # SAVE EEG TO PYTHON DICTIONARY OF NUMPY ARRAYS\n    data = eeg_from_parquet(f'{PATH}{eeg_id}.parquet', display=i<DISPLAY)              \n    all_eegs[eeg_id] = data\n    \n    if i==DISPLAY:\n        if CREATE_EEGS:\n            print(f'Processing {train_grouped_eeg.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()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:11:35.789365Z","iopub.execute_input":"2024-02-29T10:11:35.789938Z","iopub.status.idle":"2024-02-29T10:13:26.164578Z","shell.execute_reply.started":"2024-02-29T10:11:35.789884Z","shell.execute_reply":"2024-02-29T10:13:26.162140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from scipy.signal import butter, lfilter\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\nFREQS = [1,2,4,8,16][::-1]\nx = [all_eegs[EEG_IDS[0]][:,0]]\nfor k in FREQS:\n    x.append( butter_lowpass_filter(x[0], cutoff_freq=k) )\n\nplt.figure(figsize=(20,20))\nplt.plot(range(10_000),x[0], label='without filter')\nfor k in range(1,len(x)):\n    plt.plot(range(10_000),x[k]-k*(x[0].max()-x[0].min()), label=f'with filter {FREQS[k-1]}Hz')\nplt.legend()\nplt.title('Butter Low-Pass Filter Examples',size=18)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:13:26.167599Z","iopub.execute_input":"2024-02-29T10:13:26.168059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\n\nclass DataGenerator(tf.keras.utils.Sequence):\n    'Generates data for Keras'\n    def __init__(self, data, batch_size=32, shuffle=False, eegs=all_eegs, mode='train',\n                 downsample=5): \n\n        self.data = data\n        self.batch_size = batch_size\n        self.shuffle = shuffle\n        self.eegs = eegs\n        self.mode = mode\n        self.downsample = downsample\n        self.on_epoch_end()\n        \n    def __len__(self):\n        'Denotes the number of batches per epoch'\n        ct = int( np.ceil( len(self.data) / self.batch_size ) )\n        return ct\n\n    def __getitem__(self, index):\n        'Generate one batch of data'\n        indexes = self.indexes[index*self.batch_size:(index+1)*self.batch_size]\n        X, y = self.__data_generation(indexes)\n        return X[:,::self.downsample,:], y\n\n    def on_epoch_end(self):\n        'Updates indexes after each epoch'\n        self.indexes = np.arange( len(self.data) )\n        if self.shuffle: np.random.shuffle(self.indexes)\n                        \n    def __data_generation(self, indexes):\n        'Generates data containing batch_size samples' \n    \n        X = np.zeros((len(indexes),10_000,8),dtype='float32')\n        y = np.zeros((len(indexes),6),dtype='float32')\n        \n        sample = np.zeros((10_000,X.shape[-1]))\n        for j,i in enumerate(indexes):\n            row = self.data.iloc[i]      \n            data = self.eegs[row.eeg_id]\n            \n            # FEATURE ENGINEER\n            sample[:,0] = data[:,FEAT2IDX['Fp1']] - data[:,FEAT2IDX['T3']]\n            sample[:,1] = data[:,FEAT2IDX['T3']] - data[:,FEAT2IDX['O1']]\n            \n            sample[:,2] = data[:,FEAT2IDX['Fp1']] - data[:,FEAT2IDX['C3']]\n            sample[:,3] = data[:,FEAT2IDX['C3']] - data[:,FEAT2IDX['O1']]\n            \n            sample[:,4] = data[:,FEAT2IDX['Fp2']] - data[:,FEAT2IDX['C4']]\n            sample[:,5] = data[:,FEAT2IDX['C4']] - data[:,FEAT2IDX['O2']]\n            \n            sample[:,6] = data[:,FEAT2IDX['Fp2']] - data[:,FEAT2IDX['T4']]\n            sample[:,7] = data[:,FEAT2IDX['T4']] - data[:,FEAT2IDX['O2']]\n            \n            # STANDARDIZE\n            sample = np.clip(sample,-1024,1024)\n            sample = np.nan_to_num(sample, nan=0) / 32.0\n            \n            # BUTTER LOW-PASS FILTER\n            sample = butter_lowpass_filter(sample)\n            \n            X[j,] = sample\n            if self.mode!='test':\n                y[j] = row[TARGETS]\n            \n        return X,y","metadata":{"execution":{"iopub.status.idle":"2024-02-29T10:13:43.463261Z","shell.execute_reply.started":"2024-02-29T10:13:27.696382Z","shell.execute_reply":"2024-02-29T10:13:43.462138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# LOAD TRAIN \nTARGETS = train_grouped_eeg.columns[1:7]","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:13:43.464919Z","iopub.execute_input":"2024-02-29T10:13:43.465705Z","iopub.status.idle":"2024-02-29T10:13:43.470925Z","shell.execute_reply.started":"2024-02-29T10:13:43.465671Z","shell.execute_reply":"2024-02-29T10:13:43.469774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gen = DataGenerator(train_grouped_eeg, shuffle=False)\n\nfor x,y in gen:\n    for k in range(4):\n        plt.figure(figsize=(20,4))\n        offset = 0\n        for j in range(x.shape[-1]):\n            if j!=0: offset -= x[k,:,j].min()\n            plt.plot(range(2_000),x[k,:,j]+offset,label=f'feature {j+1}')\n            offset += x[k,:,j].max()\n        tt = f'{y[k][0]:0.1f}'\n        for t in y[k][1:]:\n            tt += f', {t:0.1f}'\n        plt.title(f'EEG_Id = {EEG_IDS[k]}\\n',size=14)\n        plt.legend()\n        plt.show()\n    break","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:13:43.472407Z","iopub.execute_input":"2024-02-29T10:13:43.472907Z","iopub.status.idle":"2024-02-29T10:13:45.472828Z","shell.execute_reply.started":"2024-02-29T10:13:43.472865Z","shell.execute_reply":"2024-02-29T10:13:45.471261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 初始化GPU","metadata":{}},{"cell_type":"code","source":"import os\nos.environ[\"CUDA_VISIBLE_DEVICES\"]=\"0,1\"\nimport tensorflow as tf\nprint('TensorFlow version =',tf.__version__)\n\n# USE MULTIPLE GPUS\ngpus = tf.config.list_physical_devices('GPU')\nif len(gpus)<=1: \n    strategy = tf.distribute.OneDeviceStrategy(device=\"/gpu:0\")\n    print(f'Using {len(gpus)} GPU')\nelse: \n    strategy = tf.distribute.MirroredStrategy()\n    print(f'Using {len(gpus)} GPUs')","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:13:45.473991Z","iopub.execute_input":"2024-02-29T10:13:45.474327Z","iopub.status.idle":"2024-02-29T10:13:45.489415Z","shell.execute_reply.started":"2024-02-29T10:13:45.474298Z","shell.execute_reply":"2024-02-29T10:13:45.488197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# USE MIXED PRECISION\nMIX = True\nif MIX:\n    tf.config.optimizer.set_experimental_options({\"auto_mixed_precision\": True})\n    print('Mixed precision enabled')\nelse:\n    print('Using full precision')","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:13:45.493903Z","iopub.execute_input":"2024-02-29T10:13:45.494778Z","iopub.status.idle":"2024-02-29T10:13:45.501203Z","shell.execute_reply.started":"2024-02-29T10:13:45.494743Z","shell.execute_reply":"2024-02-29T10:13:45.499963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# TRAIN SCHEDULE\ndef lrfn(epoch):\n        return [1e-3][epoch]\nLR = tf.keras.callbacks.LearningRateScheduler(lrfn, verbose = True)\nEPOCHS = 1","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:13:45.503106Z","iopub.execute_input":"2024-02-29T10:13:45.503950Z","iopub.status.idle":"2024-02-29T10:13:45.510617Z","shell.execute_reply.started":"2024-02-29T10:13:45.503910Z","shell.execute_reply":"2024-02-29T10:13:45.509083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tensorflow.keras.layers import Input, Dense, Multiply, Add, Conv1D, Concatenate\n\ndef wave_block(x, filters, kernel_size, n):\n    dilation_rates = [2**i for i in range(n)]\n    x = Conv1D(filters = filters,\n               kernel_size = 1,\n               padding = 'same')(x)\n    res_x = x\n    for dilation_rate in dilation_rates:\n        tanh_out = Conv1D(filters = filters,\n                          kernel_size = kernel_size,\n                          padding = 'same', \n                          activation = 'tanh', \n                          dilation_rate = dilation_rate)(x)\n        sigm_out = Conv1D(filters = filters,\n                          kernel_size = kernel_size,\n                          padding = 'same',\n                          activation = 'sigmoid', \n                          dilation_rate = dilation_rate)(x)\n        x = Multiply()([tanh_out, sigm_out])\n        x = Conv1D(filters = filters,\n                   kernel_size = 1,\n                   padding = 'same')(x)\n        res_x = Add()([res_x, x])\n    return res_x","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:13:45.512207Z","iopub.execute_input":"2024-02-29T10:13:45.514595Z","iopub.status.idle":"2024-02-29T10:13:45.526086Z","shell.execute_reply.started":"2024-02-29T10:13:45.514520Z","shell.execute_reply":"2024-02-29T10:13:45.524982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_model():\n        \n    # INPUT \n    inp = tf.keras.Input(shape=(2_000,8))\n    \n    ############\n    # FEATURE EXTRACTION SUB MODEL\n    inp2 = tf.keras.Input(shape=(2_000,1))\n    x = wave_block(inp2, 8, 3, 12)\n    x = wave_block(x, 16, 3, 8)\n    x = wave_block(x, 32, 3, 4)\n    x = wave_block(x, 64, 3, 1)\n    model2 = tf.keras.Model(inputs=inp2, outputs=x)\n    ###########\n    \n    # LEFT TEMPORAL CHAIN\n    x1 = model2(inp[:,:,0:1])\n    x1 = tf.keras.layers.GlobalAveragePooling1D()(x1)\n    x2 = model2(inp[:,:,1:2])\n    x2 = tf.keras.layers.GlobalAveragePooling1D()(x2)\n    z1 = tf.keras.layers.Average()([x1,x2])\n    \n    # LEFT PARASAGITTAL CHAIN\n    x1 = model2(inp[:,:,2:3])\n    x1 = tf.keras.layers.GlobalAveragePooling1D()(x1)\n    x2 = model2(inp[:,:,3:4])\n    x2 = tf.keras.layers.GlobalAveragePooling1D()(x2)\n    z2 = tf.keras.layers.Average()([x1,x2])\n    \n    # RIGHT PARASAGITTAL CHAIN\n    x1 = model2(inp[:,:,4:5])\n    x1 = tf.keras.layers.GlobalAveragePooling1D()(x1)\n    x2 = model2(inp[:,:,5:6])\n    x2 = tf.keras.layers.GlobalAveragePooling1D()(x2)\n    z3 = tf.keras.layers.Average()([x1,x2])\n    \n    # RIGHT TEMPORAL CHAIN\n    x1 = model2(inp[:,:,6:7])\n    x1 = tf.keras.layers.GlobalAveragePooling1D()(x1)\n    x2 = model2(inp[:,:,7:8])\n    x2 = tf.keras.layers.GlobalAveragePooling1D()(x2)\n    z4 = tf.keras.layers.Average()([x1,x2])\n    \n    # COMBINE CHAINS\n    y = tf.keras.layers.Concatenate()([z1,z2,z3,z4])\n    y = tf.keras.layers.Dense(64, activation='relu')(y)\n    y = tf.keras.layers.Dense(6,activation='softmax', dtype='float32')(y)\n    \n    # COMPILE MODEL\n    model = tf.keras.Model(inputs=inp, outputs=y)\n    opt = tf.keras.optimizers.Adam(learning_rate = 1e-3)\n    loss = tf.keras.losses.KLDivergence()\n    model.compile(loss=loss, optimizer = opt)\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:13:45.528231Z","iopub.execute_input":"2024-02-29T10:13:45.529035Z","iopub.status.idle":"2024-02-29T10:13:45.545893Z","shell.execute_reply.started":"2024-02-29T10:13:45.528993Z","shell.execute_reply":"2024-02-29T10:13:45.544521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_grouped_eeg.rename(columns={'expert_consensus': 'target'}, inplace=True)\nVERBOSE = 1\nFOLDS_TO_TRAIN = 1\nif not os.path.exists('WaveNet_Model'):\n    os.makedirs('WaveNet_Model')\n\nfrom sklearn.model_selection import KFold, GroupKFold\nimport tensorflow.keras.backend as K, gc\n\nall_oof = []; all_oof2 = []; all_true = []\ngkf = GroupKFold(n_splits=5)\nfor i, (train_index, valid_index) in enumerate(gkf.split(train_grouped_eeg, train_grouped_eeg.target, train_grouped_eeg.patient_id)):   \n    \n    print('#'*25)\n    print(f'### Fold {i+1}')\n    train_gen = DataGenerator(train_grouped_eeg.iloc[train_index], shuffle=True, batch_size=32)\n    valid_gen = DataGenerator(train_grouped_eeg.iloc[valid_index], shuffle=False, batch_size=64, mode='valid')\n    print(f'### train size {len(train_index)}, valid size {len(valid_index)}')\n    print('#'*25)\n    \n    # TRAIN MODEL\n    K.clear_session()\n    with strategy.scope():\n        model = build_model()\n    if TRAIN_MODEL:\n        model.fit(train_gen, verbose=VERBOSE,\n              validation_data = valid_gen,\n              epochs=EPOCHS, callbacks = [LR])\n        model.save_weights(f'WaveNet_Model/WaveNet_fold{i}.h5')\n    else:\n        model.load_weights(f'/kaggle/input/brain-eegs/WaveNet_Model/WaveNet_fold{i}.h5')\n    \n    # WAVENET OOF\n    oof = model.predict(valid_gen, verbose=VERBOSE)\n    all_oof.append(oof)\n    all_true.append(train_grouped_eeg.iloc[valid_index][TARGETS].values)\n    \n    # TRAIN MEAN OOF\n    y_train = train_grouped_eeg.iloc[train_index][TARGETS].values\n    y_valid = train_grouped_eeg.iloc[valid_index][TARGETS].values\n    oof = y_valid.copy()\n    for j in range(6):\n        oof[:,j] = y_train[:,j].mean()\n    oof = oof / oof.sum(axis=1,keepdims=True)\n    all_oof2.append(oof)\n    \n    del model, oof, y_train, y_valid\n    gc.collect()\n    \n    if i==FOLDS_TO_TRAIN-1: break\n    \nall_oof = np.concatenate(all_oof)\nall_oof2 = np.concatenate(all_oof2)\nall_true = np.concatenate(all_true)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:13:45.549311Z","iopub.execute_input":"2024-02-29T10:13:45.550369Z","iopub.status.idle":"2024-02-29T10:16:51.907257Z","shell.execute_reply.started":"2024-02-29T10:13:45.550318Z","shell.execute_reply":"2024-02-29T10:16:51.905362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/kaggle-kl-div')\nfrom kaggle_kl_div import score\n\noof = pd.DataFrame(all_oof.copy())\noof['id'] = np.arange(len(oof))\n\ntrue = pd.DataFrame(all_true.copy())\ntrue['id'] = np.arange(len(true))\n\ncv = score(solution=true, submission=oof, row_id_column_name='id')\nprint('CV Score with WaveNet Raw EEG =',cv)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:16:51.909651Z","iopub.execute_input":"2024-02-29T10:16:51.911364Z","iopub.status.idle":"2024-02-29T10:16:51.981199Z","shell.execute_reply.started":"2024-02-29T10:16:51.911310Z","shell.execute_reply":"2024-02-29T10:16:51.979635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"oof = pd.DataFrame(all_oof2.copy())\noof['id'] = np.arange(len(oof))\n\ntrue = pd.DataFrame(all_true.copy())\ntrue['id'] = np.arange(len(true))\n\ncv = score(solution=true, submission=oof, row_id_column_name='id')\nprint('CV Score with Train Means =',cv)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:16:51.983430Z","iopub.execute_input":"2024-02-29T10:16:51.987876Z","iopub.status.idle":"2024-02-29T10:16:52.053959Z","shell.execute_reply.started":"2024-02-29T10:16:51.987818Z","shell.execute_reply":"2024-02-29T10:16:52.052092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del all_eegs, train_grouped_eeg; gc.collect()\ntest = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/test.csv')\nprint('Test shape:',test.shape)\ntest.head()\nall_eegs2 = {}\nDISPLAY = 1\nEEG_IDS2 = test.eeg_id.unique()\nPATH2 = '/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/'\n\nprint('Processing Test EEG parquets...'); print()\nfor i,eeg_id in enumerate(EEG_IDS2):\n        \n    # SAVE EEG TO PYTHON DICTIONARY OF NUMPY ARRAYS\n    data = eeg_from_parquet(f'{PATH2}{eeg_id}.parquet', i<DISPLAY)\n    all_eegs2[eeg_id] = data\n# INFER MLP ON TEST\npreds = []\nmodel = build_model()\ntest_gen = DataGenerator(test, shuffle=False, batch_size=64, eegs=all_eegs2, mode='test')\n\nprint('Inferring test... ',end='')\nfor i in range(FOLDS_TO_TRAIN):\n    print(f'fold {i+1}, ',end='')\n    if TRAIN_MODEL:\n        model.load_weights(f'WaveNet_Model/WaveNet_fold{i}.h5')\n    else:\n        model.load_weights(f'/kaggle/input/brain-eegs/WaveNet_Model/WaveNet_fold{i}.h5')\n    pred = model.predict(test_gen, verbose=0)\n    preds.append(pred)\npred = np.mean(preds,axis=0)\nprint()\nprint('Test preds shape',pred.shape)    \n# CREATE SUBMISSION.CSV\nfrom IPython.display import display\n\nsub = pd.DataFrame({'eeg_id':test.eeg_id.values})\nsub[TARGETS] = pred\nsub.to_csv('submission.csv',index=False)\nprint('Submission shape',sub.shape)\ndisplay( sub.head() )\n\n# SANITY CHECK TO CONFIRM PREDICTIONS SUM TO ONE\nprint('Sub row 0 sums to:',sub.iloc[0,-6:].sum())","metadata":{"execution":{"iopub.status.busy":"2024-02-29T10:16:52.057087Z","iopub.execute_input":"2024-02-29T10:16:52.057635Z","iopub.status.idle":"2024-02-29T10:17:41.503358Z","shell.execute_reply.started":"2024-02-29T10:16:52.057593Z","shell.execute_reply":"2024-02-29T10:17:41.502218Z"},"trusted":true},"execution_count":null,"outputs":[]}]}