{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392733,"sourceType":"datasetVersion","datasetId":4297749},{"sourceId":7465251,"sourceType":"datasetVersion","datasetId":4317718},{"sourceId":7517324,"sourceType":"datasetVersion","datasetId":4378712},{"sourceId":7820073,"sourceType":"datasetVersion","datasetId":4581705,"isSourceIdPinned":true},{"sourceId":7932633,"sourceType":"datasetVersion","datasetId":4662876}],"dockerImageVersionId":30636,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Directory settings","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# directory settings\n# ====================================================\n\nimport os\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n    \nPOP_2_DIR = OUTPUT_DIR + 'pop_2_weight_oof/'\nif not os.path.exists(POP_2_DIR):\n    os.makedirs(POP_2_DIR)\n    \nPOP_1_DIR = OUTPUT_DIR + 'pop_1_weight_oof/'\nif not os.path.exists(POP_1_DIR):\n    os.makedirs(POP_1_DIR)","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:58:51.749757Z","iopub.execute_input":"2024-04-02T08:58:51.750127Z","iopub.status.idle":"2024-04-02T08:58:51.761903Z","shell.execute_reply.started":"2024-04-02T08:58:51.750097Z","shell.execute_reply":"2024-04-02T08:58:51.761011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Library","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=5","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:58:51.763363Z","iopub.execute_input":"2024-04-02T08:58:51.763628Z","iopub.status.idle":"2024-04-02T08:59:07.404062Z","shell.execute_reply.started":"2024-04-02T08:58:51.763599Z","shell.execute_reply":"2024-04-02T08:59:07.403258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# CFG\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    filter_signal=False\n    scheduler='OneCycleLR' # ['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 = 'stage1_quad_features_wavenet_gru'\n    optimizer='AdamW'\n    epochs = 15\n    factor = 0.9\n    patience = 2\n    eps = 1e-6\n    lr = 8e-3\n    min_lr = 1e-6\n    in_channels = 8\n    batch_size = 64\n    weight_decay = 1e-2\n    sampling_rate = 200\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_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 = 5\n    trn_fold = [0, 1, 2, 3, 4]\n    PATH = '/kaggle/input/hms-harmful-brain-activity-classification/'\n    data_root = \"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/\"\n    eegs_path = '/kaggle/input/hms-eeg-preprocessed-path-v3/EEG/Preprocessed_eeg_version2/'#'/kaggle/input/hms-eeg-preprocessed-path-v1/EEG/Preprocessed_eeg_version1/'","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:07.405116Z","iopub.execute_input":"2024-04-02T08:59:07.405435Z","iopub.status.idle":"2024-04-02T08:59:07.416505Z","shell.execute_reply.started":"2024-04-02T08:59:07.405411Z","shell.execute_reply":"2024-04-02T08:59:07.415719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"def init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()\n\ndef get_score(preds, targets):\n    oof = pd.DataFrame(preds.copy())\n    oof['id'] = np.arange(len(oof))\n\n    true = pd.DataFrame(targets.copy())\n    true['id'] = np.arange(len(true))\n\n    cv = score(solution=true, submission=oof, row_id_column_name='id')\n    return cv\n\ndef butter_bandpass(lowcut, highcut, fs, order=5):\n    return butter(order, [lowcut, highcut], fs=fs, btype='band')\n\ndef butter_bandpass_filter(data, lowcut, highcut, fs, order=5):\n    b, a = butter_bandpass(lowcut, highcut, fs, order=order)\n    y = lfilter(b, a, data)\n    return y\n\ndef _apply_notch_filter( data, freq, fs, Q):\n            \"\"\"Apply a notch filter to filter out the specified frequency range.\n            Args:\n                data: Input EEG data\n                freq: Frequency to be filtered out\n                fs: Sampling rate\n                Q: Quality factor\n            Returns:\n                filtered_data: Filtered EEG data\n            \"\"\"\n            b, a = iirnotch(freq, Q, fs)\n            filtered_data = filtfilt(b, a, data, axis=0)\n            return filtered_data    \n        \ndef denoise_filter(x):\n    # Sample rate and desired cutoff frequencies (in Hz).\n    fs = 200.0\n    lowcut = 1.0\n    highcut = 25.0\n    \n    # Filter a noisy signal.\n    T = 50\n    nsamples = T * fs\n    t = np.arange(0, nsamples) / fs\n    y = butter_bandpass_filter(x, lowcut, highcut, fs, order=6)\n    y = (y + np.roll(y,-1)+ np.roll(y,-2)+ np.roll(y,-3))/4\n    y = y[0:-1:4]\n    \n    return y\n\nclass KLDivLossWithLogits(nn.KLDivLoss):\n\n    def __init__(self):\n        super().__init__(reduction=\"batchmean\")\n\n    def forward(self, y, t):\n        y = nn.functional.log_softmax(y,  dim=1)\n        loss = super().forward(y, t)\n\n        return loss\n\n\n\ndef seed_torch(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\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_torch(seed=CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:07.419027Z","iopub.execute_input":"2024-04-02T08:59:07.419329Z","iopub.status.idle":"2024-04-02T08:59:07.453345Z","shell.execute_reply.started":"2024-04-02T08:59:07.419305Z","shell.execute_reply":"2024-04-02T08:59:07.452500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load train data","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\nTARGETS = train.columns[-6:]\nprint('Train shape:', train.shape )\nprint('Targets', list(TARGETS))\n\ntrain['total_evaluators'] = train[['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']].sum(axis=1)\n\n\nprint(f'There are {train.patient_id.nunique()} patients in the training data.')\nprint(f'There are {train.eeg_id.nunique()} EEG IDs in the training data.')\nprint(f'There are {train.shape[0]} unique eeg_id + votes in the training data.')","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:07.454463Z","iopub.execute_input":"2024-04-02T08:59:07.454729Z","iopub.status.idle":"2024-04-02T08:59:07.778677Z","shell.execute_reply.started":"2024-04-02T08:59:07.454706Z","shell.execute_reply":"2024-04-02T08:59:07.777744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eeg_df = pd.read_parquet(CFG.data_root + \"100261680.parquet\")\neeg_features = eeg_df.columns\nprint(f'There are {len(eeg_features)} raw eeg features')\nprint(list(eeg_features))\neeg_features = ['Fp1','T3','C3','O1','Fp2','C4','T4','O2']\neeg_features = [\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# eeg_features =[\n#         ['Fp1', 'F7', 'T3', 'T5', 'O1'],\n#         ['Fp2', 'F8', 'T4', 'T6', 'O2'],\n#         ['Fp1', 'F3', 'C3', 'P3', 'O1'],\n#         ['Fp2', 'F4', 'C4', 'P4', 'O2']\n#     ]\nfeature_to_index = {x:y for x,y in zip(eeg_features, range(len(eeg_features)))}\n\ndel eeg_df\n_ = gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:07.781008Z","iopub.execute_input":"2024-04-02T08:59:07.781424Z","iopub.status.idle":"2024-04-02T08:59:08.343194Z","shell.execute_reply.started":"2024-04-02T08:59:07.781387Z","shell.execute_reply":"2024-04-02T08:59:08.342123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Deduplicate Train EEG Id","metadata":{}},{"cell_type":"code","source":"# Create a new identifier combining multiple columns\nid_cols = ['eeg_id', 'spectrogram_id', 'seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\ntrain['new_id'] = train[id_cols].astype(str).agg('_'.join, axis=1)\ntarget_cols = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\n    \n\n    \n# Group the data by the new identifier and aggregate various features\nagg_functions = {\n        'eeg_id': 'first',\n        'eeg_label_offset_seconds': ['min', 'max'],\n        'spectrogram_label_offset_seconds': ['min', 'max', 'first'],\n        'spectrogram_id': 'first',\n        'patient_id': 'first',\n        'label_id': 'first',\n        'expert_consensus': 'first',\n        **{col: 'sum' for col in target_cols},\n        'total_evaluators': 'first',\n}\ntrain = train.groupby('new_id').agg(agg_functions).reset_index()\n\ndisplay(train)\n\n# Flatten the MultiIndex columns and adjust column names\ntrain.columns = [f\"{col[0]}_{col[1]}\" if col[1] else col[0] for col in train.columns]\ntrain.columns = train.columns.str.replace('_first', '').str.replace('_sum', '')\n    \n# Normalize the class columns\ny_data = train[target_cols].values\ny_data_normalized = y_data / y_data.sum(axis=1, keepdims=True)\ntrain[target_cols] = y_data_normalized\n\n\n\nplt.figure(figsize=(10, 6))\nplt.hist(train['total_evaluators'], bins=10, color='blue', edgecolor='black')\nplt.title('Histogram of Total Evaluators')\nplt.xlabel('Total Evaluators')\nplt.ylabel('Frequency')\nplt.grid(True)\nplt.show()\n\ndel y_data\n_ = gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:08.344456Z","iopub.execute_input":"2024-04-02T08:59:08.344723Z","iopub.status.idle":"2024-04-02T08:59:10.499342Z","shell.execute_reply.started":"2024-04-02T08:59:08.344701Z","shell.execute_reply":"2024-04-02T08:59:10.498270Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV Scheme","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedGroupKFold\nimport numpy as np\n\nbin_edges = [0, 5, 10, 15, 20, np.inf]\nnum_bins = len(bin_edges) - 1\ntrain['evaluator_bin'] = pd.cut(train['total_evaluators'], bins=bin_edges, labels=False)\npatient_groups = train.groupby('patient_id').ngroup()\nstratified_kfold = StratifiedGroupKFold(n_splits=CFG.n_fold)\ntrain_indices = []\nvalid_indices = []\nfor fold, (train_idx, valid_idx) in enumerate(stratified_kfold.split(X=train, y=train['evaluator_bin'], groups=patient_groups)):\n    train_indices.append(train_idx)\n    valid_indices.append(valid_idx) \n    train.loc[valid_idx, \"fold\"] = fold\n\nplt.figure(figsize=(8, 6))\nplt.hist(train['evaluator_bin'], bins=num_bins, color='green', edgecolor='black', align='left')\nplt.title('Distribution of Stratified Bins')\nplt.xlabel('Evaluator Bin')\nplt.ylabel('Frequency')\nplt.grid(True)\nplt.xticks(range(num_bins), [f'{bin_edges[i]}-{bin_edges[i+1]}' for i in range(num_bins)])\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:10.500538Z","iopub.execute_input":"2024-04-02T08:59:10.500827Z","iopub.status.idle":"2024-04-02T08:59:11.742784Z","shell.execute_reply.started":"2024-04-02T08:59:10.500803Z","shell.execute_reply":"2024-04-02T08:59:11.741856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Parquet to EEG Signals Numpy Processing","metadata":{}},{"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-02T08:59:11.744141Z","iopub.execute_input":"2024-04-02T08:59:11.744526Z","iopub.status.idle":"2024-04-02T08:59:11.754882Z","shell.execute_reply.started":"2024-04-02T08:59:11.744494Z","shell.execute_reply":"2024-04-02T08:59:11.754077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"def quantize_data(data, classes):\n    mu_x = mu_law_encoding(data, classes)\n    return mu_x  # quantized\n\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\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\n\ndef butter_bandpass(lowcut, highcut, fs, order=5):\n    return butter(order, [lowcut, highcut], fs=fs, btype=\"band\")\n\n\ndef butter_bandpass_filter(data, lowcut, highcut, fs, order=5):\n    b, a = butter_bandpass(lowcut, highcut, fs, order=order)\n    y = lfilter(b, a, data)\n    return y\n\n\ndef butter_lowpass_filter(\n    data, cutoff_freq=20, sampling_rate=CFG.sampling_rate, order=4\n):\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\n\nclass EEGDataset(Dataset):\n    def __init__(\n        self, df: pd.DataFrame, config, mode: str = 'train',\n        eegs_path = CFG.eegs_path, downsample: int = 5\n    ): \n        self.df = df\n        self.config = config\n        self.mode = mode\n        self.eegs_path = config.eegs_path\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_prob = 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_prob, dtype=torch.float32)\n        }\n        return output\n                        \n    def __data_generation(self, index):\n        row = self.df.iloc[index]\n        data = np.load(self.eegs_path + str(row['label_id']) + '.npy')#np.load\n          \n        if self.mode != 'test':\n            y_prob = row[self.config.target_cols].values.astype(np.float32)\n        return data, y_prob","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:11.758753Z","iopub.execute_input":"2024-04-02T08:59:11.759128Z","iopub.status.idle":"2024-04-02T08:59:11.773483Z","shell.execute_reply.started":"2024-04-02T08:59:11.759102Z","shell.execute_reply":"2024-04-02T08:59:11.772630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = EEGDataset(train, CFG, mode=\"train\")\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=CFG.batch_size,\n    shuffle=False,\n    num_workers=CFG.num_workers, pin_memory=True, drop_last=True\n)\noutput = train_dataset[0]\nX, y = output[\"eeg\"], output[\"labels\"]\nprint(f\"X shape: {X.shape}\")\nprint(f\"y shape: {y.shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:11.774456Z","iopub.execute_input":"2024-04-02T08:59:11.774756Z","iopub.status.idle":"2024-04-02T08:59:11.839668Z","shell.execute_reply.started":"2024-04-02T08:59:11.774723Z","shell.execute_reply":"2024-04-02T08:59:11.838912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eeg_ids = train.eeg_id.unique()\nif CFG.visualize:\n    for batch in train_loader:\n        X = batch.pop(\"eeg\")\n        y = batch.pop(\"labels\")\n        for item in range(4):\n            plt.figure(figsize=(20,4))\n            offset = 0\n            for col in range(X.shape[-1]):\n                if col != 0:\n                    offset -= X[item,:,col].min()\n                plt.plot(range(2000), X[item,:,col]+offset,label=f'feature {col+1}')\n                offset += X[item,:,col].max()\n            tt = f'{y[col][0]:0.1f}'\n            for t in y[col][1:]:\n                tt += f', {t:0.1f}'\n            plt.title(f'EEG_Id = {eeg_ids[item]}\\nTarget = {tt}',size=14)\n            plt.legend()\n            plt.show()\n        break","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:11.840589Z","iopub.execute_input":"2024-04-02T08:59:11.840823Z","iopub.status.idle":"2024-04-02T08:59:17.729345Z","shell.execute_reply.started":"2024-04-02T08:59:11.840803Z","shell.execute_reply":"2024-04-02T08:59:17.728502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"# Copyright (c) 2022, Kwanhyung Lee. All rights reserved.\n#\n# Licensed under the MIT License; \n# you may not use this file except in compliance with the License.\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport platform\nimport numpy as np\n\nimport torch\nimport torch.nn as nn\nfrom torch import Tensor, FloatTensor\nimport matplotlib.pyplot as plt\nfrom scipy import signal as sci_sig\n\n\nclass PSD_FEATURE1(nn.Module):\n    def __init__(self,\n            sample_rate: int = 200,\n            frame_length: int = 16,\n            frame_shift: int = 8,\n            feature_extract_by: str = 'kaldi'):\n        super(PSD_FEATURE1, self).__init__()\n\n        self.sample_rate = sample_rate\n        self.feature_extract_by = feature_extract_by.lower()\n        self.freq_resolution = 1\n\n        if self.feature_extract_by == 'kaldi':\n            assert platform.system().lower() == 'linux' or platform.system().lower() == 'darwin'\n            import torchaudio\n\n            self.transforms = torchaudio.transforms.Spectrogram(n_fft=self.freq_resolution*self.sample_rate,\n                                                                win_length=frame_length,\n                                                                hop_length=frame_shift)\n\n        else:\n            self.n_fft = self.freq_resolution*self.sample_rate\n            self.hop_length = frame_shift\n            self.frame_length = frame_length\n        \n    def psd(self, amp, begin, end):\n        return torch.mean(amp[begin*self.freq_resolution:end*self.freq_resolution], 0)\n        \n    def forward(self, batch):\n        psds_batch = []\n\n        for signals in batch:\n            psd_sample = []\n            for signal in signals:\n                if self.feature_extract_by == 'kaldi':\n                    stft = self.transforms(signal)\n                    amp = (torch.log(torch.abs(stft) + 1e-10))\n                    \n                else:\n                    stft = torch.stft(\n                        signal, self.n_fft, hop_length=self.hop_length,\n                        win_length=self.frame_length, window=torch.hamming_window(self.frame_length),\n                        center=False, normalized=False, onesided=True\n                    )\n                    amp = (torch.log(torch.abs(stft) + 1e-10))\n                # http://citeseerx.ist.psu.edu/viewdoc/download?doi=10.1.1.641.3620&rep=rep1&type=pdf\n                psd1 = self.psd(amp,0,4)\n                psd2 = self.psd(amp,4,7)\n                psd3 = self.psd(amp,7,13)\n                psd4 = self.psd(amp,13,15)\n                psd5 = self.psd(amp,14,30)\n                psd6 = self.psd(amp,31,45)\n                psd7 = self.psd(amp,55,100)\n                \n                psds = torch.stack((psd1, psd2, psd3, psd4, psd5, psd6, psd7))\n                psd_sample.append(psds)\n\n            psds_batch.append(torch.stack(psd_sample))\n\n        return torch.stack(psds_batch)\n\n\nclass PSD_FEATURE2(nn.Module):\n    def __init__(self,\n            sample_rate: int = 200,\n            frame_length: int = 16,\n            frame_shift: int = 8,\n            feature_extract_by: str = 'kaldi'):\n        super(PSD_FEATURE2, self).__init__()\n\n        self.sample_rate = sample_rate\n        self.feature_extract_by = feature_extract_by.lower()\n        self.freq_resolution = 1\n\n        if self.feature_extract_by == 'kaldi':\n            assert platform.system().lower() == 'linux' or platform.system().lower() == 'darwin'\n            import torchaudio\n\n            self.transforms = torchaudio.transforms.Spectrogram(n_fft=self.freq_resolution*self.sample_rate,\n                                                                win_length=frame_length,\n                                                                hop_length=frame_shift)\n\n        else:\n            self.n_fft = self.freq_resolution*self.sample_rate\n            self.hop_length = frame_shift\n            self.frame_length = frame_length\n        \n    def psd(self, amp, begin, end):\n        return torch.mean(amp[begin*self.freq_resolution:end*self.freq_resolution], 0)\n        \n    def forward(self, batch):\n        psds_batch = []\n\n        for signals in batch:\n            psd_sample = []\n            for signal in signals:\n                if self.feature_extract_by == 'kaldi':\n                    stft = self.transforms(signal)\n                    amp = (torch.log(torch.abs(stft) + 1e-10))\n                    \n                else:\n                    stft = torch.stft(\n                        signal, self.n_fft, hop_length=self.hop_length,\n                        win_length=self.frame_length, window=torch.hamming_window(self.frame_length),\n                        center=False, normalized=False, onesided=True\n                    )\n                    amp = (torch.log(torch.abs(stft) + 1e-10))\n                # https://ieeexplore.ieee.org/stamp/stamp.jsp?arnumber=8910555\n                psd1 = self.psd(amp,0,4)\n                psd2 = self.psd(amp,4,8)\n                psd3 = self.psd(amp,8,12)\n                psd4 = self.psd(amp,12,30)\n                psd5 = self.psd(amp,30,50)\n                psd6 = self.psd(amp,50,70)\n                psd7 = self.psd(amp,70,100)\n                \n                psds = torch.stack((psd1, psd2, psd3, psd4, psd5, psd6, psd7))\n                psd_sample.append(psds)\n\n            psds_batch.append(torch.stack(psd_sample))\n\n        return torch.stack(psds_batch)","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:17.730653Z","iopub.execute_input":"2024-04-02T08:59:17.730973Z","iopub.status.idle":"2024-04-02T08:59:17.757064Z","shell.execute_reply.started":"2024-04-02T08:59:17.730946Z","shell.execute_reply":"2024-04-02T08:59:17.756140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"Sinc-based convolution\nMirco Ravanelli, Yoshua Bengio,\n\"Speaker Recognition from raw waveform with SincNet\".\nhttps://arxiv.org/abs/1808.00158\n\"\"\"\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nimport torch.nn as nn\nimport sys\nfrom torch.autograd import Variable\nimport math\nimport matplotlib.pyplot as plt\nfrom torch import Tensor\n\ndef flip(x, dim):\n    xsize = x.size()\n    dim = x.dim() + dim if dim < 0 else dim\n    x = x.contiguous()\n    x = x.view(-1, *xsize[dim:])\n    x = x.view(x.size(0), x.size(1), -1)[:, getattr(torch.arange(x.size(1)-1, \n                      -1, -1), ('cpu','cuda')[x.is_cuda])().long(), :]\n    return x.view(xsize)\n\n# class LayerNorm(nn.Module):\n#     \"\"\" Wrapper class of torch.nn.LayerNorm \"\"\"\n#     def __init__(self, dim: int, eps: float = 1e-6) -> None:\n#         super(LayerNorm, self).__init__()\n#         self.gamma = nn.Parameter(torch.ones(dim))\n#         self.beta = nn.Parameter(torch.zeros(dim))\n#         self.eps = eps\n\n#     def forward(self, z: Tensor) -> Tensor:\n#         print(\"0: \", z.shape)\n#         mean = z.mean(dim=-1, keepdim=True)\n#         std = z.std(dim=-1, keepdim=True)\n#         output = (z - mean) / (std + self.eps)\n#         print(\"1: \", output.shape)\n#         print(\"2: \", self.gamma.shape)\n#         print(\"3: \", self.beta.shape)\n#         output = self.gamma * output + self.beta\n\n#         return output\n\nclass SincConv_fast(nn.Module):\n    \"\"\"Sinc-based convolution\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=200, in_channels=1,\n                 stride=1, padding=0, normalize=None, slice_len=280, dilation=1, bias=False, groups=1, min_low_hz=0, min_band_hz=2):\n        super(SincConv_fast, self).__init__()\n\n        self.normalize = normalize\n        if self.normalize == \"layernorm\":\n            # self.ln0=nn.LayerNorm([1,slice_len])\n            self.ln0=nn.LayerNorm(1)\n            # self.ln0=LayerNorm(1)\n        elif self.normalize == \"batchnorm\":\n            self.bn0=nn.BatchNorm1d([1],momentum=0.05)\n\n        if in_channels != 1:\n            msg = \"SincConv only support one input channel (here, in_channels = {%i})\" % (in_channels)\n            raise ValueError(msg)\n\n        self.out_channels = out_channels\n        self.kernel_size = kernel_size\n        \n        # Forcing the filters to be odd (i.e, perfectly symmetrics)\n        if kernel_size%2==0:\n            self.kernel_size=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\n        low_hz = 0\n        high_hz = self.sample_rate / 2 - (self.min_low_hz + self.min_band_hz)\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\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        # Hamming window\n        #self.window_ = torch.hamming_window(self.kernel_size)\n        n_lin=torch.linspace(0, (self.kernel_size/2)-1, steps=int((self.kernel_size/2))) # computing only half of the window\n        self.window_=0.54-0.46*torch.cos(2*math.pi*n_lin/self.kernel_size)\n\n        # (1, kernel_size/2)\n        n = (self.kernel_size - 1) / 2.0\n        self.n_ = 2*math.pi*torch.arange(-n, 0).view(1, -1) / self.sample_rate # Due to symmetry, I only need half of the time axes\n\n\n\n    def forward(self, waveforms):\n        \"\"\"\n        Parameters\n        ----------\n        eeg waveforms : `torch.Tensor` (batch_size, 1, n_samples)\n\n        Returns\n        -------\n        features : `torch.Tensor` (batch_size, out_channels, n_samples_out)\n            Batch of sinc filters activations.\n        \"\"\"\n\n        waveforms = waveforms.unsqueeze(1)\n        if self.normalize == \"layernorm\":\n            waveforms = self.ln0(waveforms)\n        elif self.normalize == \"batchnorm\":\n            waveforms = self.bn0(waveforms)\n\n        self.n_ = self.n_.to(waveforms.device)\n\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 + 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        band_pass_left=((torch.sin(f_times_t_high)-torch.sin(f_times_t_low))/(self.n_/2))*self.window_ # Equivalent of Eq.4 of the reference paper (SPEAKER RECOGNITION FROM RAW WAVEFORM WITH SINCNET). I just have expanded the sinc and simplified the terms. This way I avoid several useless computations. \n\n        band_pass_center = 2*band.view(-1,1)\n\n        band_pass_right= torch.flip(band_pass_left,dims=[1])\n        \n        band_pass=torch.cat([band_pass_left,band_pass_center,band_pass_right],dim=1)\n   \n        # import matplotlib.pyplot as plt\n        # plt.figure()\n        # plt.subplot(3,1,1)\n        # plt.plot(band_pass[0].detach().cpu().numpy())\n        # plt.subplot(3,1,2)\n        # plt.plot(band_pass[5].detach().cpu().numpy())\n        # plt.subplot(3,1,3)\n        # plt.plot(band_pass[10].detach().cpu().numpy())\n        # plt.show()\n        # exit(1)\n\n        band_pass = band_pass / (2*band[:,None])\n\n        self.filters = (band_pass).view(\n            self.out_channels, 1, self.kernel_size)\n\n        return F.conv1d(waveforms, self.filters, stride=self.stride,\n                        padding=self.padding, dilation=self.dilation,\n                         bias=None, groups=1) \n\nclass sincnet_conv_layers(nn.Module):\n    def __init__(self, args):\n        super(sincnet_conv_layers, self).__init__()\n        self.cnn_layer_num = args.sincnet_layer_num\n        self.conv  = nn.ModuleList([])\n        self.ln  = nn.ModuleList([])\n        self.act = nn.ModuleList([])\n        self.maxpool = nn.ModuleList([])\n\n        slice_len = args.window_size_sig + (args.sincnet_kernel_size -1)\n        cnn_channel_sizes = args.cnn_channel_sizes\n        cnn_kernel_sizes = [args.sincnet_kernel_size, 5, 5]\n        self.stride_list = [args.sincnet_stride, 2, 2]\n        current_input = args.window_size_sig + (args.sincnet_kernel_size -1)\n        normalize = args.sincnet_input_normalize\n        \n        for i in range(self.cnn_layer_num):\n            if i==0:\n                self.conv.append(SincConv_fast(out_channels=cnn_channel_sizes[i], \n                                            kernel_size=cnn_kernel_sizes[i], \n                                            stride=self.stride_list[i], \n                                            padding=0, \n                                            normalize=normalize,\n                                            slice_len=slice_len))\n            else:\n                self.conv.append(nn.Conv1d(in_channels=cnn_channel_sizes[i-1], \n                                            out_channels=cnn_channel_sizes[i], \n                                            kernel_size=cnn_kernel_sizes[i], \n                                            stride=self.stride_list[i]))\n            \n            # self.maxpool.append(nn.MaxPool1d(self.max_pool_list[i]))\n            layernorm_output = int((current_input-cnn_kernel_sizes[i]+1)/self.stride_list[i])\n            # self.ln.append(nn.LayerNorm([cnn_channel_sizes[i], layernorm_output]))\n            self.ln.append(nn.LayerNorm(cnn_channel_sizes[i]))\n            # self.ln.append(LayerNorm(cnn_channel_sizes[i]))\n            self.act.append(nn.LeakyReLU(0.2))\n\n            current_input = layernorm_output\n\n    def forward(self, x):\n        # print(\"x0: \", x.shape)\n        for i in range(self.cnn_layer_num):\n            if i==0:\n                x = torch.abs(self.conv[i](x))\n                x = x.permute(0,2,1)\n                # x = self.act[i](self.ln[i](x.permute(0,2,1)).permute(0,2,1))\n                x = self.act[i](self.ln[i](x))\n            else:\n                x = self.act[i](self.conv[i](x))\n        return x\n\n\nclass SINCNET_FEATURE(nn.Module):\n    def __init__(self, args, num_eeg_channel):\n        super(SINCNET_FEATURE, self).__init__()\n        \n        # self.conv_list = nn.ModuleList()\n        # for _ in range(num_eeg_channel):\n        #     self.conv_list.append(sincnet_conv_layers(args))\n        self.sinc_conv = sincnet_conv_layers(args)\n        \n    def forward(self, waveforms):\n        output_list = []\n        waveforms = waveforms.permute(1,0,2)\n        # for idx, conv in enumerate(self.conv_list):\n        #     output_list.append(conv(waveforms[idx]))\n        for waveform in waveforms:\n            print(waveform.shape)\n            output_list.append(self.sinc_conv(waveform))\n\n        return torch.stack(output_list).permute(1,0,3,2)","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:17.758537Z","iopub.execute_input":"2024-04-02T08:59:17.759057Z","iopub.status.idle":"2024-04-02T08:59:17.797622Z","shell.execute_reply.started":"2024-04-02T08:59:17.759016Z","shell.execute_reply":"2024-04-02T08:59:17.796925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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        #print(x.shape)\n        #x = x.permute(0, 2,1)\n        #x = self.sincnet(x)\n        # Start first stack\n        #x_temp = x[:,0,:,:].permute(0,2,1)\n        y = self.extract_features(x)\n        y = self.head(y)\n        \n        return y","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:17.798680Z","iopub.execute_input":"2024-04-02T08:59:17.798948Z","iopub.status.idle":"2024-04-02T08:59:17.842305Z","shell.execute_reply.started":"2024-04-02T08:59:17.798907Z","shell.execute_reply":"2024-04-02T08:59:17.841384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\niot = torch.randn(2, 2000, 15)#.cuda()\nmodel = CustomModel()\noutput = model(iot)\nprint(output.shape)\n\ndel iot, model\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:17.843374Z","iopub.execute_input":"2024-04-02T08:59:17.843613Z","iopub.status.idle":"2024-04-02T08:59:20.512393Z","shell.execute_reply.started":"2024-04-02T08:59:17.843593Z","shell.execute_reply":"2024-04-02T08:59:20.511410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nimport torch\nfrom torch.optim.optimizer import Optimizer\n\n\nclass Adan(Optimizer):\n    \"\"\"\n    Implements a pytorch variant of Adan\n    Adan was proposed in\n    Adan: Adaptive Nesterov Momentum Algorithm for Faster Optimizing Deep Models[J]. arXiv preprint arXiv:2208.06677, 2022.\n    https://arxiv.org/abs/2208.06677\n    Arguments:\n        params (iterable): iterable of parameters to optimize or dicts defining parameter groups.\n        lr (float, optional): learning rate. (default: 1e-3)\n        betas (Tuple[float, float, flot], optional): coefficients used for computing \n            running averages of gradient and its norm. (default: (0.98, 0.92, 0.99))\n        eps (float, optional): term added to the denominator to improve \n            numerical stability. (default: 1e-8)\n        weight_decay (float, optional): decoupled weight decay (L2 penalty) (default: 0)\n        max_grad_norm (float, optional): value used to clip \n            global grad norm (default: 0.0 no clip)\n        no_prox (bool): how to perform the decoupled weight decay (default: False)\n    \"\"\"\n\n    def __init__(self, params, lr=1e-3, betas=(0.98, 0.92, 0.99), eps=1e-8,\n                 weight_decay=0.2, max_grad_norm=0.0, no_prox=False):\n        if not 0.0 <= max_grad_norm:\n            raise ValueError(\"Invalid Max grad norm: {}\".format(max_grad_norm))\n        if not 0.0 <= lr:\n            raise ValueError(\"Invalid learning rate: {}\".format(lr))\n        if not 0.0 <= eps:\n            raise ValueError(\"Invalid epsilon value: {}\".format(eps))\n        if not 0.0 <= betas[0] < 1.0:\n            raise ValueError(\"Invalid beta parameter at index 0: {}\".format(betas[0]))\n        if not 0.0 <= betas[1] < 1.0:\n            raise ValueError(\"Invalid beta parameter at index 1: {}\".format(betas[1]))\n        if not 0.0 <= betas[2] < 1.0:\n            raise ValueError(\"Invalid beta parameter at index 2: {}\".format(betas[2]))\n        defaults = dict(lr=lr, betas=betas, eps=eps,\n                        weight_decay=weight_decay,\n                        max_grad_norm=max_grad_norm, no_prox=no_prox)\n        super(Adan, self).__init__(params, defaults)\n\n    def __setstate__(self, state):\n        super(Adan, self).__setstate__(state)\n        for group in self.param_groups:\n            group.setdefault('no_prox', False)\n\n    @torch.no_grad()\n    def restart_opt(self):\n        for group in self.param_groups:\n            group['step'] = 0\n            for p in group['params']:\n                if p.requires_grad:\n                    state = self.state[p]\n                    # State initialization\n\n                    # Exponential moving average of gradient values\n                    state['exp_avg'] = torch.zeros_like(p)\n                    # Exponential moving average of squared gradient values\n                    state['exp_avg_sq'] = torch.zeros_like(p)\n                    # Exponential moving average of gradient difference\n                    state['exp_avg_diff'] = torch.zeros_like(p)\n\n    @torch.no_grad()\n    def step(self):\n        \"\"\"\n            Performs a single optimization step.\n        \"\"\"\n        if self.defaults['max_grad_norm'] > 0:\n            device = self.param_groups[0]['params'][0].device\n            global_grad_norm = torch.zeros(1, device=device)\n\n            max_grad_norm = torch.tensor(self.defaults['max_grad_norm'], device=device)\n            for group in self.param_groups:\n\n                for p in group['params']:\n                    if p.grad is not None:\n                        grad = p.grad\n                        global_grad_norm.add_(grad.pow(2).sum())\n\n            global_grad_norm = torch.sqrt(global_grad_norm)\n\n            clip_global_grad_norm = torch.clamp(max_grad_norm / (global_grad_norm + group['eps']), max=1.0)\n        else:\n            clip_global_grad_norm = 1.0\n\n        for group in self.param_groups:\n            beta1, beta2, beta3 = group['betas']\n            # assume same step across group now to simplify things\n            # per parameter step can be easily support by making it tensor, or pass list into kernel\n            if 'step' in group:\n                group['step'] += 1\n            else:\n                group['step'] = 1\n\n            bias_correction1 = 1.0 - beta1 ** group['step']\n\n            bias_correction2 = 1.0 - beta2 ** group['step']\n\n            bias_correction3 = 1.0 - beta3 ** group['step']\n\n            for p in group['params']:\n                if p.grad is None:\n                    continue\n\n                state = self.state[p]\n                if len(state) == 0:\n                    state['exp_avg'] = torch.zeros_like(p)\n                    state['exp_avg_sq'] = torch.zeros_like(p)\n                    state['exp_avg_diff'] = torch.zeros_like(p)\n\n                grad = p.grad.mul_(clip_global_grad_norm)\n                if 'pre_grad' not in state or group['step'] == 1:\n                    state['pre_grad'] = grad\n\n                copy_grad = grad.clone()\n\n                exp_avg, exp_avg_sq, exp_avg_diff = state['exp_avg'], state['exp_avg_sq'], state['exp_avg_diff']\n                diff = grad - state['pre_grad']\n\n                update = grad + beta2 * diff\n                exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1)  # m_t\n                exp_avg_diff.mul_(beta2).add_(diff, alpha=1 - beta2)  # diff_t\n                exp_avg_sq.mul_(beta3).addcmul_(update, update, value=1 - beta3)  # n_t\n\n                denom = ((exp_avg_sq).sqrt() / math.sqrt(bias_correction3)).add_(group['eps'])\n                update = ((exp_avg / bias_correction1 + beta2 * exp_avg_diff / bias_correction2)).div_(denom)\n\n                if group['no_prox']:\n                    p.data.mul_(1 - group['lr'] * group['weight_decay'])\n                    p.add_(update, alpha=-group['lr'])\n                else:\n                    p.add_(update, alpha=-group['lr'])\n                    p.data.div_(1 + group['lr'] * group['weight_decay'])\n\n                state['pre_grad'] = copy_grad","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:20.513611Z","iopub.execute_input":"2024-04-02T08:59:20.513898Z","iopub.status.idle":"2024-04-02T08:59:20.538399Z","shell.execute_reply.started":"2024-04-02T08:59:20.513875Z","shell.execute_reply":"2024-04-02T08:59:20.537526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\n\ndef train_fn(fold, train_loader, model, criterion, optimizer, epoch, scheduler, device):\n    model.train()\n    scaler = torch.cuda.amp.GradScaler(enabled=CFG.apex)\n    losses = AverageMeter()\n    start = end = time.time()\n    global_step = 0\n    for step, batch in enumerate(train_loader):\n        eegs = batch['eeg'].to(device)\n        labels = batch['labels'].to(device)\n        batch_size = labels.size(0)\n        with torch.cuda.amp.autocast(enabled=CFG.apex):\n            y_preds= model(eegs)\n            loss = criterion(F.log_softmax(y_preds, dim=1), labels)\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        losses.update(loss.item(), batch_size)\n        scaler.scale(loss).backward()\n        grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.max_grad_norm)\n        if (step + 1) % CFG.gradient_accumulation_steps == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            global_step += 1\n            if CFG.batch_scheduler:\n                scheduler.step()\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(train_loader)-1):\n            print('Epoch: [{0}][{1}/{2}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  'Grad: {grad_norm:.4f}  '\n                  'LR: {lr:.8f}  '\n                  .format(epoch+1, step, len(train_loader), \n                          remain=timeSince(start, float(step+1)/len(train_loader)),\n                          loss=losses,\n                          grad_norm=grad_norm,\n                          lr=scheduler.get_lr()[0]))\n        if CFG.wandb:\n            wandb.log({f\"[fold{fold}] loss\": losses.val,\n                       f\"[fold{fold}] lr\": scheduler.get_lr()[0]})\n    return losses.avg\n\n\ndef valid_fn(valid_loader, model, criterion, device):\n    losses = AverageMeter()\n    model.eval()\n    preds = []\n    targets = []\n    start = end = time.time()\n    for step, batch in enumerate(valid_loader):\n        eegs = batch['eeg'].to(device)\n        labels = batch['labels'].to(device)\n        batch_size = labels.size(0)\n        with torch.no_grad():\n            y_preds = model(eegs)\n            loss = criterion(F.log_softmax(y_preds, dim=1), labels)\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        losses.update(loss.item(), batch_size)\n        preds.append(nn.Softmax(dim=1)(y_preds).to('cpu').numpy())\n        targets.append(labels.to('cpu').numpy())\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(valid_loader)-1):\n            print('EVAL: [{0}/{1}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  .format(step, len(valid_loader),\n                          loss=losses,\n                          remain=timeSince(start, float(step+1)/len(valid_loader))))\n    predictions = np.concatenate(preds)\n    targets = np.concatenate(targets)\n    return losses.avg, predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:20.540055Z","iopub.execute_input":"2024-04-02T08:59:20.540394Z","iopub.status.idle":"2024-04-02T08:59:20.562468Z","shell.execute_reply.started":"2024-04-02T08:59:20.540363Z","shell.execute_reply":"2024-04-02T08:59:20.561541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train Loop","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# train loop\n# ====================================================\ndef train_loop(folds, fold, directory):\n    \n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n\n    # ====================================================\n    # loader\n    # ====================================================\n    if CFG.stage1_pop1:\n        train_folds = folds[folds['fold'] != fold ].reset_index(drop=True)\n    else:\n        train_folds = folds[(folds['fold'] != fold) & (folds['total_evaluators'] >= 10)].reset_index(drop=True)\n        #valid_folds = folds[(folds['fold'] == fold) & (folds['total_evaluators'] >= 10)].reset_index(drop=True)\n    valid_folds = folds[folds['fold'] == fold ].reset_index(drop=True)    \n    valid_labels = valid_folds[ CFG.target_cols].values\n    \n    train_dataset = EEGDataset(train_folds, CFG, mode=\"train\")\n    valid_dataset = EEGDataset(valid_folds, CFG, mode=\"train\")\n\n    train_loader = DataLoader(train_dataset,\n                              batch_size=CFG.batch_size,\n                              shuffle=True,\n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\n    valid_loader = DataLoader(valid_dataset,\n                              batch_size=CFG.batch_size * 2,\n                              shuffle=False,\n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n\n    # ====================================================\n    # model & optimizer\n    # ====================================================\n    model = CustomModel()\n    if CFG.stage2_pop2:\n        model_weight = POP_1_DIR + f\"{CFG.model_name}_fold{fold}_best_version{VERSION}_stage1.pth\"\n        checkpoint = torch.load(model_weight, map_location=device)\n        model.load_state_dict(checkpoint[\"model\"])\n    model.to(device)\n    # CPMP: wrap the model to use all GPUs\n    model = nn.DataParallel(model)\n    \n    def build_optimizer(cfg, model, device):\n        lr = cfg.lr\n        if cfg.optimizer == \"SAM\":\n            base_optimizer = torch.optim.SGD  # define an optimizer for the \"sharpness-aware\" update\n            optimizer_model = SAM(model.parameters(), base_optimizer, lr=lr, momentum=0.9, weight_decay=cfg.weight_decay, adaptive=True)\n        elif cfg.optimizer == \"Ranger21\":\n            optimizer_model = Ranger21(model.parameters(), lr=lr, weight_decay=cfg.weight_decay, \n            num_epochs=cfg.epochs, num_batches_per_epoch=len(train_loader))\n        elif cfg.optimizer == \"SGD\":\n            optimizer_model = torch.optim.SGD(model.parameters(), lr=lr, weight_decay=cfg.weight_decay, momentum=0.9)\n        elif cfg.optimizer == \"Adam\":\n            optimizer_model = Adam(model.parameters(), lr=lr, weight_decay=CFG.weight_decay)\n        elif cfg.optimizer == \"AdamW\":\n            optimizer_model = AdamW(model.parameters(), lr=lr, weight_decay=CFG.weight_decay)\n        elif cfg.optimizer == \"Lion\":\n            optimizer_model = Lion(model.parameters(), lr=lr, weight_decay=cfg.weight_decay)\n        elif cfg.optimizer == \"Adan\":\n            optimizer_model = Adan(model.parameters(), lr=lr, weight_decay=cfg.weight_decay)\n    \n        return optimizer_model\n    \n    optimizer = build_optimizer(CFG, model, device)\n    \n    # ====================================================\n    # scheduler\n    # ====================================================\n    # ====================================================\n\n    def get_scheduler(optimizer):\n        if CFG.scheduler=='ReduceLROnPlateau':\n            scheduler = ReduceLROnPlateau(optimizer, **CFG.reduce_params)\n        elif CFG.scheduler=='CosineAnnealingLR':\n            scheduler = CosineAnnealingLR(optimizer, **CFG.cosanneal_params)\n        elif CFG.scheduler=='CosineAnnealingWarmRestarts':\n            scheduler = CosineAnnealingWarmRestarts(optimizer, **CFG.cosanneal_res_params)\n        elif CFG.scheduler=='OneCycleLR':\n            scheduler = OneCycleLR(optimizer=optimizer, epochs=CFG.epochs, pct_start=0.0, steps_per_epoch=len(train_loader),\n        max_lr=CFG.lr, div_factor=25, final_div_factor=4.0e-01)\n        return scheduler\n    \n    scheduler = get_scheduler(optimizer)\n\n    # ====================================================\n    # loop\n    # ====================================================\n    criterion = nn.KLDivLoss(reduction=\"batchmean\")\n\n    \n    best_score = np.inf\n\n    for epoch in range(CFG.epochs):\n\n        start_time = time.time()\n\n        # train\n        avg_loss = train_fn(fold, train_loader, model, criterion, optimizer, epoch, scheduler, device)\n\n        # eval\n        avg_val_loss, predictions = valid_fn(valid_loader, model, criterion, device)\n\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  avg_val_loss: {avg_val_loss:.4f}  time: {elapsed:.0f}s')\n        if CFG.wandb:\n            wandb.log({f\"[fold{fold}] epoch\": epoch+1, \n                       f\"[fold{fold}] avg_train_loss\": avg_loss, \n                       f\"[fold{fold}] avg_val_loss\": avg_val_loss,\n                       f\"[fold{fold}] score\": score})\n        \n        if best_score > avg_val_loss:\n            best_score = avg_val_loss\n            LOGGER.info(f'Epoch {epoch+1} - Save Best valid loss: {avg_val_loss:.4f} Model')\n            # CPMP: save the original model. It is stored as the module attribute of the DP model.\n            if CFG.stage1_pop1:\n                \n                torch.save({'model': model.module.state_dict(),\n                            'predictions': predictions},\n                             directory+f\"{CFG.model_name}_fold{fold}_best_version{VERSION}_stage1.pth\")\n            else:\n                \n                torch.save({'model': model.module.state_dict(),\n                            'predictions': predictions},\n                             directory+f\"{CFG.model_name}_fold{fold}_best_version{VERSION}_stage2.pth\")\n                \n    if CFG.stage1_pop1:\n        predictions = torch.load(directory+f\"{CFG.model_name}_fold{fold}_best_version{VERSION}_stage1.pth\", \n                             map_location=torch.device('cpu'))['predictions']\n    else:\n        predictions = torch.load(directory+f\"{CFG.model_name}_fold{fold}_best_version{VERSION}_stage2.pth\", \n                             map_location=torch.device('cpu'))['predictions']\n    valid_folds[[f\"pred_{c}\" for c in CFG.target_cols]] = predictions\n    valid_folds[CFG.target_cols] = valid_labels \n    torch.cuda.empty_cache()\n    _ = gc.collect()\n    \n    return valid_folds, best_score","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:20.565026Z","iopub.execute_input":"2024-04-02T08:59:20.565383Z","iopub.status.idle":"2024-04-02T08:59:20.590653Z","shell.execute_reply.started":"2024-04-02T08:59:20.565349Z","shell.execute_reply":"2024-04-02T08:59:20.589744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## STAGE 1","metadata":{}},{"cell_type":"code","source":"if __name__ == '__main__':\n    \n    if CFG.train:\n        oof_df = pd.DataFrame()\n        scores = []\n        for fold in range(CFG.n_fold):\n            if fold in CFG.trn_fold:\n                _oof_df, score = train_loop(train, fold, POP_1_DIR)\n                oof_df = pd.concat([oof_df, _oof_df])\n                scores.append(score)\n                LOGGER.info(f\"========== fold: {fold} result ==========\")\n                LOGGER.info(f'Score with best loss weights stage1: {score}')\n        oof_df = oof_df.reset_index(drop=True)\n        LOGGER.info(f\"========== CV ==========\")\n        LOGGER.info(f'Score with best loss weights: {np.mean(scores)}')\n        oof_df.to_csv(POP_1_DIR+f'{CFG.model_name}_oof_df_version{VERSION}_stage1.csv', index=False)\n        \n    if CFG.wandb:\n        wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2024-04-02T08:59:20.591836Z","iopub.execute_input":"2024-04-02T08:59:20.592609Z","iopub.status.idle":"2024-04-02T09:07:20.818377Z","shell.execute_reply.started":"2024-04-02T08:59:20.592578Z","shell.execute_reply":"2024-04-02T09:07:20.817564Z"},"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\n# === Pre-process OOF ===\nlabel_cols = CFG.target_cols\ngt = oof_df[[\"eeg_id\"] + CFG.target_cols]\ngt.sort_values(by=\"eeg_id\", inplace=True)\ngt.reset_index(inplace=True, drop=True)\n\npreds = oof_df[[\"eeg_id\"] + CFG.pred_cols]\npreds.columns = [\"eeg_id\"] + CFG.target_cols\npreds.sort_values(by=\"eeg_id\", inplace=True)\npreds.reset_index(inplace=True, drop=True)\n\ny_trues = gt[CFG.target_cols]\ny_preds = preds[CFG.target_cols]\n\noof = pd.DataFrame(y_preds.copy())\noof['id'] = np.arange(len(oof))\n\ntrue = pd.DataFrame(y_trues.copy())\ntrue['id'] = np.arange(len(true))\n\ncv = score(solution=true, submission=oof, row_id_column_name='id')\nprint(f'CV Stage1 Score with {CFG.model_name} Raw EEG =',cv)","metadata":{"execution":{"iopub.status.busy":"2024-04-02T09:07:20.819919Z","iopub.execute_input":"2024-04-02T09:07:20.820684Z","iopub.status.idle":"2024-04-02T09:07:20.869558Z","shell.execute_reply.started":"2024-04-02T09:07:20.820644Z","shell.execute_reply":"2024-04-02T09:07:20.868713Z"},"trusted":true},"execution_count":null,"outputs":[]}]}