{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","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}],"dockerImageVersionId":31041,"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\"\"\"\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:\nSun, Y., Si, X., He, R. et al. An Automated Classifier of Harmful Brain Activities for Clinical Usage Based on a Vision-Inspired Pre-trained Framework. npj Digit. Med. 8, 768 (2025).\nhttps://doi.org/10.1038/s41746-025-02154-4\n\nIf you have any questions or concerns regarding this code or the related manuscript, please contact syuri@tju.edu.cn.\n\n19 December 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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-19T02:25:23.444765Z","iopub.execute_input":"2025-12-19T02:25:23.445386Z","iopub.status.idle":"2025-12-19T02:25:40.078942Z","shell.execute_reply.started":"2025-12-19T02:25:23.445355Z","shell.execute_reply":"2025-12-19T02:25:40.078352Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}