{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":11830157,"sourceType":"datasetVersion","datasetId":7431921}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!/usr/bin/env python3\n# -*- coding: utf-8 -*-\n\"\"\"\n\nWelcome to validate and use this algorithm\n\nThis code is associated with a manuscript \"An Automated Classifier of Harmful Brain Activities for Clinical Usage Based on a Vision-Inspired Pre-trained Framework\".\n\nCite as:\n@misc{sun2025automatedclassifierharmfulbrain,\n      title={An Automated Classifier of Harmful Brain Activities for Clinical Usage Based on a Vision-Inspired Pre-trained Framework}, \n      author={Yulin Sun and Xiaopeng Si and Runnan He and Xiao Hu and Peter Smielewski and Wenlong Wang and Xiaoguang Tong and Wei Yue and Meijun Pang and Kuo Zhang and Xizi Song and Dong Ming and Xiuyun Liu},\n      year={2025},\n      eprint={2507.08874},\n      archivePrefix={arXiv},\n      primaryClass={cs.LG},\n      url={https://arxiv.org/abs/2507.08874}, \n}\n\nThe work presented in this code is currently under review and should be cited accordingly once published.\nIf you have any questions or concerns regarding this code or the related manuscript, please contact syuri@tju.edu.cn.\n\nJune 20, 2025\n\"\"\"\n\n# =============================================================================\n# GLOBAL VARIABLES AND CONFIGURATION\n# =============================================================================\n\n# Flag to determine if model training is needed\nNEEDTRAIN = True        \n# Path to trained model weights for testing\nLOAD_MODELS_FROM = 'modelsxxxxxxx'       \n\nimport os\n# Set Keras backend to TensorFlow\nos.environ[\"KERAS_BACKEND\"] = \"tensorflow\"\n\n# Determine platform (local or Kaggle)\nif os.getcwd().split(os.sep)[1] == 'home':\n    PLATFORM = 'local'  # Local training or online testing\n    # Find the correct models directory in local input\n    for dir_name in os.listdir('./input/'):\n        if dir_name[:6] == 'models':\n            LOAD_MODELS_FROM = dir_name\nelif os.getcwd().split(os.sep)[1] == 'kaggle':\n    PLATFORM = 'kaggle' # Kaggle platform\n    NEEDTRAIN = False\n    # Find the correct models directory in Kaggle input\n    for dir_name in os.listdir('/kaggle/input/'):\n        if dir_name[:6] == 'models':\n            LOAD_MODELS_FROM = dir_name\n\n# Data type used (EEG in this case)\nDATATYPE = ['eeg']    \nprint(DATATYPE)\n\n# Set paths based on platform\nif PLATFORM == 'local':\n    LOAD_MODELS_FROM = f'./input/{LOAD_MODELS_FROM}'\n    LOAD_DATA_FROM = './input/hms-harmful-brain-activity-classification'\nelif PLATFORM == 'kaggle':\n    LOAD_MODELS_FROM = f'/kaggle/input/{LOAD_MODELS_FROM}'\n    LOAD_DATA_FROM = '/kaggle/input/hms-harmful-brain-activity-classification'\n\n# EEG sampling configuration\nSFREQ = 200             \nRSFREQ = 200            \nEEG_LENGTH = 50         \nEEG_LENGTH_USED = 50\nEEG_CHANNEL_USED = 16 \n\n# EEG filtering configuration\nfilter_range = [0.5, 45]   \nSEED = 2024             \nBATCHSIZE = 16        \nLEARN_RATE = 1e-3\nEPOCHS = 15\nSPLITS = 5\n\n# Flag for preprocessing EEG files\nREAD_EEG_FILES = False      \n# Dictionaries to store preprocessed EEG data\neegs = {}               \neegs_test = {}             \n\n# Brain channels configuration\nBRAIN = [\n         'Fp1-F7', 'F7-T3', 'T3-T5', 'T5-O1',   \n         'Fp1-F3', 'F3-C3', 'C3-P3', 'P3-O1',   \n         'Fz-Cz', 'Cz-Pz',\n         'Fp2-F4', 'F4-C4', 'C4-P4', 'P4-O2',   \n         'Fp2-F8', 'F8-T4', 'T4-T6', 'T6-O2',   \n        ]\n\nTEST_BATCHSIZE = 128\n\n# =============================================================================\n# IMPORTS\n# =============================================================================\n\n# Suppress warnings and set environment variables\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'\nos.environ['CUDA_VISIBLE_DEVICES']='0, 1'\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Import necessary libraries\nimport pandas as pd, numpy as np\nfrom sklearn.metrics import confusion_matrix\nfrom tensorflow.keras import optimizers\nimport matplotlib.pyplot as plt\nfrom scipy import signal\nimport time\nimport gc\n\n# Set random seeds for reproducibility\nnp.random.seed(SEED)\nos.environ['PYTHONHASHSEED'] = str(SEED)\nos.environ['TF_DETERMINISTIC_OPS'] = '1'\n\n# Configure TensorFlow\nimport tensorflow as tf\nprint(tf.version.VERSION)\nprint(tf.config.list_physical_devices('GPU'))\ngpus = tf.config.list_physical_devices('GPU')\nif gpus:\n    try:\n        for gpu in gpus:\n            tf.config.experimental.set_memory_growth(gpu, True)  \n    except RuntimeError as e:\n        print(e)\n        \ntf.random.set_seed(SEED)\ntf.keras.utils.set_random_seed(SEED)\ntf.config.experimental.enable_op_determinism()\n\n# Configure mixed precision training\nMIX = True\nif MIX:\n    policy = tf.keras.mixed_precision.Policy('mixed_float16')\n    tf.keras.mixed_precision.set_global_policy(policy)\nelse:\n    print('Using full precision')\n    \n# Import itertools if needed for training\nif NEEDTRAIN:\n    import itertools\n\n# =============================================================================\n# LOAD TRAIN DATAFRAME\n# =============================================================================\n\n# Load training data\ndf = pd.read_csv(os.path.join(LOAD_DATA_FROM, 'train.csv'))\nTARGETS = df.columns[-6:]\nprint('Train shape:', df.shape)\nprint('Targets', list(TARGETS))\n\n# =============================================================================\n# CREATE NON-OVERLAPPING EEG ID TRAIN DATAFRAME\n# =============================================================================\n\nif NEEDTRAIN:\n    TARGETS_RAW = list()\n    for i in TARGETS:\n        TARGETS_RAW.append(i + '_raw')\n        \n    if READ_EEG_FILES:\n        # Create a non-overlapping EEG ID dataframe\n        train = df.drop_duplicates(['eeg_id', 'seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']).reset_index(drop=True)\n        \n        train['sign_id'] = train.index.values\n        df['sign_id'] = df.index.values\n        \n        y_data = train[TARGETS].values\n        train[TARGETS_RAW] = y_data\n        y_data = y_data / y_data.sum(axis=1,keepdims=True)\n        train[TARGETS] = y_data\n\n        train.to_csv('train.csv', index=False)\n    else:\n        # Load preprocessed training dataframe\n        train = pd.read_csv('train.csv')\n\n# =============================================================================\n# READ TRAIN EEGS\n# =============================================================================\n\nif filter_range != None:\n    # Configure EEG filtering\n    b, a = signal.butter(3, np.float32(filter_range)*2/RSFREQ, 'bandpass')\n    \nif NEEDTRAIN:\n    PATH = os.path.join(LOAD_DATA_FROM, 'train_eegs') + '/'\n    if READ_EEG_FILES:\n        # Read and preprocess EEG files\n        time_start_time = time.time()\n        \n        for i, eeg_id in enumerate(train.eeg_id.unique()):\n            \n            if i%200==0:\n                gc.collect()\n                xx = (time.time() - time_start_time)\n                yy = xx / (i+1) * len(train.eeg_id.unique())\n                print(i, f'time: {round(xx / 60, 2)} min / {round(yy / 60, 2)} min')\n            eeg_default = pd.read_parquet(os.path.join(PATH, (str(eeg_id) + '.parquet')))\n            \n            eeg = list()\n            for channel in BRAIN:\n                eeg_temp = (eeg_default.loc[:, channel.split('-')[0]] - eeg_default.loc[:, channel.split('-')[1]]).values\n                eeg_temp[np.isnan(eeg_temp)] = 0\n                eeg.append(np.reshape(eeg_temp, (1, -1)))\n            eeg = np.concatenate(eeg, axis=0)\n            \n            if SFREQ != RSFREQ:\n                eeg = signal.resample_poly(eeg, RSFREQ, SFREQ, axis=1)\n\n            eeg = np.clip(eeg, a_min=-1024, a_max=1024)\n            \n            if filter_range != None:\n                eeg = signal.filtfilt(b, a, eeg, axis=1)\n                \n            eeg = np.array(eeg, dtype=np.float32)\n            \n            if 'eeg' in DATATYPE:\n                eegs[eeg_id] = eeg\n\n        # Save preprocessed EEG data\n        if not os.path.exists('./input/preprocess'):\n            os.makedirs('./input/preprocess')\n        if 'eeg' in DATATYPE:\n            np.save('./input/preprocess/eegs.npy', eegs, allow_pickle=True)\n\n    else:\n        # Load preprocessed EEG data\n        if PLATFORM == 'local':\n            datapath = './' + os.path.join('input', 'preprocess')\n        elif PLATFORM == 'kaggle':\n            datapath = '/kaggle/' + os.path.join('input', 'preprocess')\n\n        eegs = np.load(os.path.join(datapath, 'eegs.npy'), allow_pickle=True).item()\n\n# =============================================================================\n# DATA GENERATOR\n# =============================================================================\n\nclass DataGenerator(tf.keras.utils.Sequence):\n    def __init__(self, dataframe, batch_size=32, shuffle=False, sample_weights=False, mode='train',\n                 eegs=None, stage=2): \n\n        self.dataframe = dataframe\n        self.batch_size = batch_size\n        self.shuffle = shuffle\n        self.sample_weights = sample_weights\n        self.mode = mode\n        self.eegs = eegs\n        self.stage = stage\n        self.on_epoch_end()\n        \n    def __len__(self):\n        # Calculate number of batches\n        ct = int( np.ceil( len(self.dataframe) / 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, sample_weights = self.__data_generation(indexes)\n        return x, y, sample_weights\n\n    def on_epoch_end(self):\n        # Update indexes after each epoch\n        self.nan = 0\n        self.indexes = np.arange( len(self.dataframe) )\n        if self.shuffle: np.random.shuffle(self.indexes)\n                        \n    def __data_generation(self, indexes):\n        # Generate data for a batch\n        x_eeg = np.zeros((len(indexes), EEG_CHANNEL_USED, round(EEG_LENGTH_USED * RSFREQ)),dtype='float32')\n        y = np.zeros((len(indexes), len(TARGETS)),dtype='float32')\n        sample_weights = np.zeros((len(indexes), 1),dtype='float32')\n            \n        for j, i in enumerate(indexes):\n            row = self.dataframe.iloc[i]\n            if self.mode != 'test':\n                sample_weight = sum(row[TARGETS_RAW].values)/20\n            \n            # Process EEG data based on mode\n            if self.mode == 'test':\n                r_eeg = 0\n            else:\n                rows = df.loc[(df.eeg_id == row.eeg_id) * (df.seizure_vote == row.seizure_vote_raw) * (df.lpd_vote == row.lpd_vote_raw) * (df.gpd_vote == row.gpd_vote_raw) * (df.lrda_vote == row.lrda_vote_raw) * (df.grda_vote == row.grda_vote_raw), :].reset_index(drop=True)\n                if self.mode == 'train':\n                    rows = rows.iloc[np.random.permutation(len(rows))].reset_index(drop=True)\n                    row = rows.loc[0, :]\n                elif self.mode == 'valid':\n                    row = rows.sort_values(by='eeg_sub_id').reset_index(drop=True).iloc[len(rows)//2]\n                r_eeg = row.eeg_label_offset_seconds\n                if (self.mode == 'train'):\n                    r_eeg = r_eeg + np.random.random() * 10 - 5\n                    r_eeg = max(0, r_eeg)\n                    r_eeg = min(r_eeg, self.eegs[row.eeg_id].shape[1] / RSFREQ - 50)\n\n            eeg = self.eegs[row.eeg_id][:, round(r_eeg * RSFREQ):round((r_eeg + 50) * RSFREQ)]\n            eeg = np.concatenate((eeg[0:round(EEG_CHANNEL_USED/2), :], eeg[-round(EEG_CHANNEL_USED/2):, :]), axis=0)\n\n            eeg = eeg[:, round((EEG_LENGTH - EEG_LENGTH_USED) * RSFREQ / 2):round((EEG_LENGTH + EEG_LENGTH_USED) * RSFREQ / 2)]\n           \n            if self.mode=='train':\n                # Apply data augmentation for training\n                if (self.stage == 2) and (np.random.rand() > 0):\n                    eeg2 = eeg.copy()\n                    eeg[4:8, :] = eeg2[12:16, :]\n                    eeg[8:12, :] = eeg2[4:8, :]\n                    eeg[12:16, :] = eeg2[8:12, :]\n                else:\n                    if np.random.rand() > 0.5:\n                        mask = round(np.random.rand() * eeg.shape[1])\n                        eeg[:, mask:round(mask + np.random.rand() * eeg.shape[1] * 0.02)] = 0\n                        \n                    if np.random.rand() > 0.5:\n                        mask = round(np.random.rand() * eeg.shape[1])\n                        eeg[:, mask:round(mask + np.random.rand() * eeg.shape[1] * 0.02)] = 0\n                        \n                    if np.random.rand() > 0.5:\n                        mask = round(np.random.rand() * eeg.shape[1])\n                        eeg[:, mask:round(mask + np.random.rand() * eeg.shape[1] * 0.02)] = 0\n                    \n                    \n                    if np.random.rand() > 0.5:\n                        eeg[np.random.permutation(eeg.shape[0])[0], :] = 0\n            \n                    if np.random.rand() > 0.5:\n                        eeg[np.random.permutation(eeg.shape[0])[0], :] = 0\n                    \n                    eeg[0:round(EEG_CHANNEL_USED/2), :] = eeg[0:round(EEG_CHANNEL_USED/2), :][np.random.permutation(8), :]\n                    eeg[-round(EEG_CHANNEL_USED/2):, :] = eeg[-round(EEG_CHANNEL_USED/2):, :][np.random.permutation(8), :]\n                    \n                    eeg2 = eeg.copy()\n                    eeg[4:8, :] = eeg2[12:16, :]\n                    eeg[8:12, :] = eeg2[4:8, :]\n                    eeg[12:16, :] = eeg2[8:12, :]\n                    \n                    if np.random.rand() > 0.5:\n                        eeg = eeg[::-1, :]\n\n                    if np.random.rand() > 0.5:\n                        eeg = -eeg\n                    \n                    if np.random.rand() > 0.5:\n                        eeg = eeg[:, ::-1]\n            else:\n                # Process validation/test data\n                eeg2 = eeg.copy()\n                eeg[4:8, :] = eeg2[12:16, :]\n                eeg[8:12, :] = eeg2[4:8, :]\n                eeg[12:16, :] = eeg2[8:12, :]\n\n            # Normalize EEG data\n            eeg = np.clip(eeg, a_min=-1024, a_max=1024)\n            eeg = eeg + 1024\n            eeg = eeg / 2048 * 255\n            \n            x_eeg[j] = eeg\n            \n            if self.mode!='test':\n                y[j] = row[TARGETS].values / sum(row[TARGETS].values)\n                \n                if self.sample_weights:\n                    sample_weights[j] = sample_weight\n                else:\n                    sample_weights[j] = 1\n        \n        return x_eeg, y, sample_weights\n\n# =============================================================================\n# MODEL BUILDING\n# =============================================================================\n\nclass CosineAnnealingLRScheduler(optimizers.schedules.LearningRateSchedule):\n    def __init__(self, total_step, lr_max, lr_min=0, warmth_rate=0):\n        super(CosineAnnealingLRScheduler, self).__init__()\n        self.total_step = total_step\n\n        if warmth_rate == 0:\n            self.warm_step = 1\n        else:\n            self.warm_step = int(warmth_rate)\n\n        self.lr_max = lr_max\n        self.lr_min = lr_min\n        \n        self.begin = 1\n        \n    def __call__(self, step):\n        # Implement cosine annealing learning rate schedule\n        if step == self.total_step:\n            self.begin = 0\n            self.lr_max = self.lr_max * 0.5\n            self.lr_min = self.lr_min * 0.1\n            \n        step = step % self.total_step\n        step = step + 1\n\n        if (self.begin==1) and (step < self.warm_step):\n            lr = self.lr_max / self.warm_step * step\n        else:\n            if self.begin==1:\n                if self.total_step == 1:\n                    lr = self.lr_max\n                else:\n                    lr = self.lr_min + 0.5 * (self.lr_max - self.lr_min) * (1.0 + tf.cos((step - self.warm_step) / (self.total_step-self.warm_step) * np.pi))\n            else:\n                lr = self.lr_min + 0.5 * (self.lr_max - self.lr_min) * (1.0 + tf.cos(step / 10 * np.pi))\n        \n        return np.float32(lr)\n\nclass IniToOne(tf.keras.initializers.Initializer):\n    def __init__(self):\n        super(IniToOne, self).__init__()\n\n    def __call__(self, shape, dtype=None):\n       # Initialize weights with ones\n       assert len(shape) == 3\n       filter_length, input_channel, filter_count = shape\n       \n       kernel = np.zeros(shape, dtype=np.float32)\n       for i in range(filter_count):\n           kernel[i%filter_length, 0, i] = 1.0\n       kernel = tf.convert_to_tensor(kernel, dtype=dtype)\n       return kernel\n\n    def get_config(self):\n        return {}\n\nclass SumToOne(tf.keras.constraints.Constraint):\n    def __init__(self):\n        super(SumToOne, self).__init__()\n\n    def __call__(self, w):\n        # Constrain weights to sum to one\n        w = tf.abs(w)\n        w_normed = w / tf.reduce_sum(w, axis=[0, 1], keepdims=True)\n        return w_normed\n\n    def get_config(self):\n        return {}\n\nclass IniToOneAtten(tf.keras.initializers.Initializer):\n    def __init__(self):\n        super(IniToOneAtten, self).__init__()\n\n    def __call__(self, shape, dtype=None):\n       # Initialize attention weights\n       assert len(shape) == 3\n       filter_length, input_channel, filter_count = shape\n       \n       kernel = np.zeros(shape, dtype=np.float32)\n       kernel[(filter_length-1)//2:(filter_length)//2+1, :, :] = 1/((filter_length)//2+1 - (filter_length-1)//2)\n       kernel = tf.convert_to_tensor(kernel, dtype=dtype)\n       return kernel\n\n    def get_config(self):\n        return {}\n\nclass SumToOneAtten(tf.keras.constraints.Constraint):\n    def __init__(self):\n        super(SumToOneAtten, self).__init__()\n\n    def __call__(self, w):\n        # Constrain attention weights\n        w = tf.abs(w)\n        w_normed = w / tf.reduce_sum(w, axis=[0, 1], keepdims=True)\n        return w_normed\n\n    def get_config(self):\n        return {}\n\ndef build_model():\n    # Build the neural network model\n    inp_eeg = tf.keras.Input(shape=(EEG_CHANNEL_USED, round(EEG_LENGTH_USED * RSFREQ)), name='eeg')\n    x_eeg_raw = tf.keras.layers.Reshape((inp_eeg.shape[1], inp_eeg.shape[2], 1))(inp_eeg)\n\n    strides = 10\n    if PLATFORM == 'local':\n        # Configure EEG embedding layer with custom initialization\n        eeg_embed = tf.keras.layers.Conv1D(filters=strides*3, kernel_size=strides, strides=strides,\n                                           padding='same', use_bias=False, activation=None,\n                                           kernel_initializer = IniToOne(),\n                                           kernel_constraint = SumToOne(),\n                                           input_shape=(None, 1)\n                                           )\n    else:\n        # Configure standard EEG embedding layer\n        eeg_embed = tf.keras.layers.Conv1D(filters=strides*3, kernel_size=strides, strides=strides,\n                                            padding='same', use_bias=False, activation=None)\n    \n    x_eeg = tf.keras.layers.TimeDistributed(eeg_embed)(x_eeg_raw)\n    \n    # Reshape and prepare data for EfficientNet\n    x_eeg = tf.keras.layers.Concatenate(axis=-1)([tf.keras.layers.Reshape((x_eeg.shape[1], x_eeg.shape[2], -1, 1))(x_eeg[:, :, :, 0*strides:1*strides]),\n                                                  tf.keras.layers.Reshape((x_eeg.shape[1], x_eeg.shape[2], -1, 1))(x_eeg[:, :, :, 1*strides:2*strides]),\n                                                  tf.keras.layers.Reshape((x_eeg.shape[1], x_eeg.shape[2], -1, 1))(x_eeg[:, :, :, 2*strides:3*strides])\n                                                  ])\n    x_eeg = tf.keras.layers.Permute([4, 2, 1, 3])(x_eeg)\n    x_eeg = tf.keras.layers.Reshape((x_eeg.shape[1], x_eeg.shape[2], -1))(x_eeg)\n    x_eeg = tf.keras.layers.Permute((3, 2, 1))(x_eeg)\n    \n    # Load pre-trained EfficientNetV2B3\n    base_model_eeg = tf.keras.applications.EfficientNetV2B3(include_top=False, weights=None,\n                                                           include_preprocessing=True)\n    \n    if NEEDTRAIN:\n        if PLATFORM == 'local':\n            # Load local pre-trained weights\n            base_model_eeg.load_weights(f'./input/pre-trained-weights/{base_model_eeg.name}_notop.h5')\n        if PLATFORM == 'kaggle':\n            # Load Kaggle pre-trained weights\n            base_model_eeg.load_weights(f'/kaggle/input/pre-trained-weights/{base_model_eeg.name}_notop.h5')\n\n    base_model_eeg.name = 'eeg_extractor'\n\n    x_eeg = base_model_eeg(x_eeg)\n    \n    # Process output features\n    x_eeg = x_eeg[:, :, (x_eeg.shape[2]-1)//2:(x_eeg.shape[2])//2+1, :]\n    \n    x_eeg = tf.keras.layers.GlobalAveragePooling2D()(x_eeg)\n    x_eeg = tf.keras.layers.Dropout(0.5)(x_eeg)\n\n    # Output layer with softmax activation\n    y = tf.keras.layers.Dense(len(TARGETS), activation='softmax', dtype='float32')(x_eeg)\n \n    model = tf.keras.Model(inputs=inp_eeg, outputs=y)\n        \n    return model\n\n# =============================================================================\n# TRAINING FUNCTION\n# =============================================================================\n\ndef train_fold(i, stage, train_index, valid_index, df_train_stage1, df_valid_stage1, df_train_stage2, df_valid_stage2, \n               build_model, BATCHSIZE, EPOCHS, LEARN_RATE, TARGETS, TARGETS_RAW):\n    \n    print('#'*25)\n    print(f'### Fold {i+1}')\n    \n    # Build model\n    model = build_model()\n    loss = tf.keras.losses.KLDivergence()\n\n    # Configure data generators based on stage\n    if stage == 1:\n        train_gen_stage = DataGenerator(df_train_stage1, shuffle=True, sample_weights=True, batch_size=BATCHSIZE, eegs=eegs, stage=stage)\n        valid_gen_stage = DataGenerator(df_valid_stage1, shuffle=False, sample_weights=True, batch_size=BATCHSIZE*2, mode='valid', eegs=eegs, stage=stage)\n        opt = tf.keras.optimizers.AdamW(learning_rate=LEARN_RATE)\n        # Configure callbacks for stage 1 training\n        callbacks_stage = [\n            tf.keras.callbacks.LearningRateScheduler(CosineAnnealingLRScheduler(EPOCHS, LEARN_RATE, LEARN_RATE * 0.1 * 0.1, 5)),\n            tf.keras.callbacks.ModelCheckpoint(filepath=os.path.join('models', f'fold{i}_stage1.weights.h5'),\n                                               monitor='val_loss', mode='min',\n                                               save_weights_only=True, save_best_only=True)\n        ]\n    elif stage == 2:\n        train_gen_stage = DataGenerator(df_train_stage2, shuffle=True, sample_weights=False, batch_size=BATCHSIZE * 2, eegs=eegs, stage=stage)\n        valid_gen_stage = DataGenerator(df_valid_stage2, shuffle=False, sample_weights=False, batch_size=BATCHSIZE*2 * 2, mode='valid', eegs=eegs, stage=stage)\n        opt = tf.keras.optimizers.AdamW(learning_rate=LEARN_RATE * 0.1 * 3)\n        model.load_weights(os.path.join('models', f'fold{i}_stage1.weights.h5'))\n        # Configure callbacks for stage 2 training\n        callbacks_stage = [\n            tf.keras.callbacks.LearningRateScheduler(CosineAnnealingLRScheduler(max(round(EPOCHS/3), 1), LEARN_RATE * 0.1 * 3, LEARN_RATE * 0.1 * 0.1 * 0.1, 0)),\n            tf.keras.callbacks.ModelCheckpoint(filepath=os.path.join('models', f'fold{i}_stage2.weights.h5'),\n                                               monitor='val_loss', mode='min',\n                                               save_weights_only=True, save_best_only=True)\n        ]\n\n    # Compile model\n    model.compile(loss=loss, optimizer=opt)\n    \n    # Train model\n    if stage == 1:\n        history = model.fit(train_gen_stage, verbose=1, validation_data=valid_gen_stage,\n                            epochs=EPOCHS, callbacks=callbacks_stage)\n    elif stage == 2:\n        history = model.fit(train_gen_stage, verbose=1, validation_data=valid_gen_stage,\n                            epochs=max(round(EPOCHS/3), 1), callbacks=callbacks_stage)\n\n    # Load best weights\n    model.load_weights(os.path.join('models', f'fold{i}_stage{stage}.weights.h5'))\n \n    # Plot training history\n    loss = history.history['loss']\n    val_loss = history.history['val_loss']\n    epochs = range(1, len(loss) + 1)\n    plt.figure()\n    plt.plot(epochs, loss, 'bo', label='loss')\n    plt.plot(epochs, val_loss, 'b', label='val_loss')\n    plt.title(f'loss: {round(min(loss), 4)}, val loss: {round(min(val_loss), 4)}', fontsize=12)\n    plt.legend()\n    plt.savefig(os.path.join('models', f'fold{i}_stage{stage}.svg'))\n    plt.close()\n\n    # Generate confusion matrix\n    if stage == 1:\n        valid_stage = df_valid_stage1[TARGETS].values\n    elif stage == 2:\n        valid_stage = df_valid_stage2[TARGETS].values\n    predict_stage = model.predict(valid_gen_stage)\n        \n    # Clean up resources\n    del train_gen_stage, valid_gen_stage, history, model\n    tf.keras.backend.clear_session()\n    gc.collect()\n\n    # Calculate and plot confusion matrix\n    cm = confusion_matrix(np.argmax(valid_stage, 1), np.argmax(predict_stage, 1))\n    cm = cm / np.sum(cm, 1, keepdims=True)\n        \n    plt.figure()\n    plt.imshow(cm, interpolation='nearest', cmap=plt.cm.Blues)\n    plt.title('Confusion Matrix')\n    plt.colorbar()\n    tick_marks = np.arange(6)\n    plt.xticks(tick_marks, [f'{TARGETS[i][:-5]}' for i in [0, 1, 2, 3, 4, 5]], fontsize=10)\n    plt.yticks(tick_marks, [f'{TARGETS[i][:-5]}' for i in [0, 1, 2, 3, 4, 5]], fontsize=10)\n    thresh = cm.max() / 2.\n    for ii, jj in itertools.product(range(cm.shape[0]), range(cm.shape[1])):\n        if cm[ii, jj] > -0.1:\n            plt.text(jj, ii, str(round(cm[ii, jj] * 1e4) * 1e-2)[:5], horizontalalignment=\"center\", color=\"white\" if cm[ii, jj] > thresh else \"black\", fontsize=10)\n    plt.xlabel('Predicted label')\n    plt.ylabel('True label')\n    plt.tight_layout()\n    plt.savefig(os.path.join('models', f'fold{i}_stage{stage}_cm.svg'))\n    plt.close()\n\n    # Clean up resources\n    del df_train_stage1, df_valid_stage1, df_train_stage2, df_valid_stage2\n    gc.collect()\n\n# =============================================================================\n# MAIN TRAINING AND INFERENCE\n# =============================================================================\n\nif __name__ == '__main__':\n    if NEEDTRAIN:\n        # Create models directory if needed\n        if not os.path.exists('models'):\n            os.makedirs('models')\n        \n        # Configure cross-validation\n        from sklearn.model_selection import GroupKFold\n        import multiprocessing as mp\n        mp.set_start_method('spawn')  # Set spawn method for multiprocessing\n        \n        gkf = GroupKFold(n_splits=SPLITS)\n        \n        # Perform k-fold cross-validation\n        for i, (train_index, valid_index) in enumerate(gkf.split(train, train.expert_consensus, train.patient_id)):  \n            print('#'*25)\n            print(f'### Fold {i+1}')\n            \n            df_train_stage1 = train.iloc[train_index].reset_index(drop=True)\n            df_valid_stage1 = train.iloc[valid_index].reset_index(drop=True)\n\n            df_train_stage2 = df_train_stage1[np.sum(df_train_stage1[TARGETS_RAW].values, 1) >= 10].reset_index(drop=True)\n            df_valid_stage2 = df_valid_stage1[np.sum(df_valid_stage1[TARGETS_RAW].values, 1) >= 10].reset_index(drop=True)\n\n            # Train two stages for each fold\n            for stage in [1, 2]:\n                p = mp.Process(target=train_fold, args=(\n                    i,\n                    stage,\n                    train_index,\n                    valid_index,\n                    df_train_stage1,\n                    df_valid_stage1,\n                    df_train_stage2,\n                    df_valid_stage2,\n                    build_model,\n                    BATCHSIZE,\n                    EPOCHS,\n                    LEARN_RATE,\n                    TARGETS,\n                    TARGETS_RAW,\n                ))\n                p.start()\n                p.join()\n    \n    # =============================================================================\n    # INFERENCE ON TEST DATA\n    # =============================================================================\n    \n    else:\n        # Load all trained models for inference\n        preds_all = []\n        models = list()\n        for model_i in range(999):\n            if os.path.exists(os.path.join(LOAD_MODELS_FROM, f'fold{model_i}_stage2.weights.h5')):\n                print(f'Fold {model_i+1}')\n                model = build_model()\n                model.load_weights(os.path.join(LOAD_MODELS_FROM, f'fold{model_i}_stage2.weights.h5'))\n                models.append(model)\n\n        # Load test data\n        test = pd.read_csv(os.path.join(LOAD_DATA_FROM, 'test.csv'))\n        test['sign_id'] = test.index.values\n        print('Test shape', test.shape)\n\n        PATH_test = os.path.join(LOAD_DATA_FROM, 'test_eegs') + '/'\n\n        # Process test EEG data\n        for i, eeg_id in enumerate(test.eeg_id):\n            if i%100==0: print(i,', ',end='')\n            eeg_default = pd.read_parquet(os.path.join(PATH_test, (str(eeg_id) + '.parquet')))\n            \n            eeg = list()\n            for channel in BRAIN:\n                eeg_temp = (eeg_default.loc[:, channel.split('-')[0]] - eeg_default.loc[:, channel.split('-')[1]]).values\n                eeg_temp[np.isnan(eeg_temp)] = 0\n                eeg.append(np.reshape(eeg_temp, (1, -1)))\n            eeg = np.concatenate(eeg, axis=0)\n            \n            if SFREQ != RSFREQ:\n                eeg = signal.resample_poly(eeg, RSFREQ, SFREQ, axis=1)\n\n            eeg = np.clip(eeg, a_min=-1024, a_max=1024)\n            eegshape = eeg.shape[1]\n            eeg = np.concatenate((eeg[:, ::-1], eeg, eeg[:, ::-1]), axis=1)\n            if filter_range != None:\n                eeg = signal.filtfilt(b, a, eeg, axis=1)\n            eeg = eeg[:, eegshape:eegshape*2]\n            \n            eeg = np.array(eeg, dtype=np.float32)\n            \n            eegs_test[eeg_id] = eeg\n\n            # Make predictions in batches\n            if ((i+1)%TEST_BATCHSIZE==0) or ((i+1)==len(test.eeg_id)):\n                preds = []\n                test_gen = DataGenerator(test.loc[max(i-TEST_BATCHSIZE+1, len(preds_all)):i, :], shuffle=False, sample_weights=False, batch_size=TEST_BATCHSIZE, mode='test', eegs=eegs_test, stage=2)\n                for model_i in range(len(models)):\n                    pred = models[model_i].predict(test_gen, verbose=1)\n                    preds.append(pred)\n                pred = np.mean(preds, axis=0)\n                del eegs_test\n                gc.collect()\n                eegs_test = {}\n                if len(preds_all) == 0:\n                    preds_all = pred.copy()\n                else:\n                    preds_all = np.concatenate((preds_all, pred), axis=0)\n\n        # Prepare submission file\n        sub = pd.DataFrame({'eeg_id': test.eeg_id.values})\n        sub[TARGETS] = preds_all\n        sub.to_csv('submission.csv', index=False)\n        print('Submission shape', sub.shape)\n        sub.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}