{"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"}],"dockerImageVersionId":30840,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ---------------------------\n# 1. Import Libraries\n# ---------------------------\nimport os\nimport sys\nimport tensorflow as tf\nfrom tensorflow import keras\nimport numpy as np\nimport pandas as pd\nimport joblib\nimport matplotlib.pyplot as plt\nfrom glob import glob\nfrom tqdm import tqdm\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom tensorflow.keras.layers import RandomFlip, RandomRotation, RandomZoom, GlobalAveragePooling2D, Dense, Dropout\nfrom tensorflow.keras.metrics import Precision, Recall, AUC\nfrom tensorflow.keras.callbacks import ModelCheckpoint, ReduceLROnPlateau, EarlyStopping\nfrom sklearn.metrics import classification_report, precision_recall_curve, auc\nfrom tensorflow.keras.layers import BatchNormalization","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T13:23:07.467854Z","iopub.execute_input":"2025-01-30T13:23:07.468114Z","iopub.status.idle":"2025-01-30T13:23:19.899295Z","shell.execute_reply.started":"2025-01-30T13:23:07.468088Z","shell.execute_reply":"2025-01-30T13:23:19.898628Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 2. Configuration\n# ---------------------------\nclass CFG:\n    verbose = 1\n    seed = 42\n    image_size = [400, 300]\n    epochs = 30  # Increased epochs\n    batch_size = 64\n    lr_mode = \"cos\"\n    drop_remainder = True\n    num_classes = 6\n    fold = 5\n    class_names = ['Seizure', 'LPD', 'GPD', 'LRDA', 'GRDA', 'Other']\n    label2name = dict(enumerate(class_names))\n    name2label = {v: k for k, v in label2name.items()}\n\ntf.random.set_seed(CFG.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T13:23:19.900851Z","iopub.execute_input":"2025-01-30T13:23:19.901396Z","iopub.status.idle":"2025-01-30T13:23:19.906127Z","shell.execute_reply.started":"2025-01-30T13:23:19.901372Z","shell.execute_reply":"2025-01-30T13:23:19.905227Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 3. Data Preparation\n# ---------------------------\nBASE_PATH = \"/kaggle/input/hms-harmful-brain-activity-classification\"\nSPEC_DIR = \"/tmp/dataset/hms-hbac\"\nos.makedirs(SPEC_DIR + '/train_spectrograms', exist_ok=True)\nos.makedirs(SPEC_DIR + '/test_spectrograms', exist_ok=True)\n\ndf = pd.read_csv(f'{BASE_PATH}/train.csv')\ndf['eeg_path'] = f'{BASE_PATH}/train_eegs/' + df['eeg_id'].astype(str) + '.parquet'\ndf['spec_path'] = f'{BASE_PATH}/train_spectrograms/' + df['spectrogram_id'].astype(str) + '.parquet'\ndf['spec2_path'] = f'{SPEC_DIR}/train_spectrograms/' + df['spectrogram_id'].astype(str) + '.npy'\ndf['class_name'] = df.expert_consensus.copy()\ndf['class_label'] = df.expert_consensus.map(CFG.name2label)\n\ntest_df = pd.read_csv(f'{BASE_PATH}/test.csv')\ntest_df['eeg_path'] = f'{BASE_PATH}/test_eegs/' + test_df['eeg_id'].astype(str) + '.parquet'\ntest_df['spec_path'] = f'{BASE_PATH}/test_spectrograms/' + test_df['spectrogram_id'].astype(str) + '.parquet'\ntest_df['spec2_path'] = f'{SPEC_DIR}/test_spectrograms/' + test_df['spectrogram_id'].astype(str) + '.npy'\n\ndef process_spec(spec_id, split=\"train\"):\n    spec_path = f\"{BASE_PATH}/{split}_spectrograms/{spec_id}.parquet\"\n    npy_file_path = f\"{SPEC_DIR}/{split}_spectrograms/{spec_id}.npy\"\n    \n    if os.path.exists(npy_file_path):\n        return\n    \n    if not os.path.exists(spec_path):\n        return\n    \n    try:\n        spec = pd.read_parquet(spec_path)\n        spec = spec.fillna(0).values[:, 1:].T\n        spec = spec.astype(\"float32\")\n        np.save(npy_file_path, spec)\n    except Exception as e:\n        return\n\nspec_ids = df[\"spectrogram_id\"].unique()\n_ = joblib.Parallel(n_jobs=-1)(\n    joblib.delayed(process_spec)(spec_id, \"train\") for spec_id in tqdm(spec_ids, desc=\"Processing train spectrograms\")\n)\n\ntest_spec_ids = test_df[\"spectrogram_id\"].unique()\n_ = joblib.Parallel(n_jobs=-1)(\n    joblib.delayed(process_spec)(spec_id, \"test\") for spec_id in tqdm(test_spec_ids, desc=\"Processing test spectrograms\")\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T13:23:19.908057Z","iopub.execute_input":"2025-01-30T13:23:19.908467Z","iopub.status.idle":"2025-01-30T13:25:54.852329Z","shell.execute_reply.started":"2025-01-30T13:23:19.908367Z","shell.execute_reply":"2025-01-30T13:25:54.851422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 4. Data Augmentation\n# ---------------------------\ndef build_augmenter():\n    augmentation = keras.Sequential([\n        RandomFlip(\"horizontal_and_vertical\"),\n        RandomRotation(0.3),\n        RandomZoom(0.3),\n        keras.layers.RandomContrast(0.2),\n    ])\n    \n    def augment(img, label):\n        img = augmentation(img, training=True)\n        return img, label\n    \n    return augment","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T13:25:54.853524Z","iopub.execute_input":"2025-01-30T13:25:54.853883Z","iopub.status.idle":"2025-01-30T13:25:54.858334Z","shell.execute_reply.started":"2025-01-30T13:25:54.853846Z","shell.execute_reply":"2025-01-30T13:25:54.857578Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 5. Data Decoder\n# ---------------------------\ndef build_decoder(with_labels=True, target_size=CFG.image_size, dtype=32):\n    def decode_signal(path, offset=None):\n        file_bytes = tf.io.read_file(path)\n        sig = tf.io.decode_raw(file_bytes, tf.float32)\n        sig = sig[1024 // dtype:]\n        sig = tf.reshape(sig, [400, -1])\n        \n        if offset is not None:\n            offset = offset // 2\n            sig = sig[:, offset:offset + 300]\n            pad_size = tf.maximum(0, 300 - tf.shape(sig)[1])\n            sig = tf.pad(sig, [[0, 0], [0, pad_size]])\n            sig = tf.reshape(sig, [400, 300])\n        \n        sig = tf.clip_by_value(sig, tf.exp(-4.0), tf.exp(8.0))\n        sig = tf.math.log(sig)\n        sig = (sig - tf.reduce_mean(sig)) / (tf.math.reduce_std(sig) + 1e-6)\n        sig = tf.tile(sig[..., None], [1, 1, 3])\n        return sig\n    \n    def decode_label(label):\n        return tf.one_hot(label, CFG.num_classes)\n    \n    def decode_with_labels(path, offset=None, label=None):\n        sig = decode_signal(path, offset)\n        label = decode_label(label)\n        return (sig, label)\n    \n    return decode_with_labels if with_labels else decode_signal","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T13:25:54.859218Z","iopub.execute_input":"2025-01-30T13:25:54.859487Z","iopub.status.idle":"2025-01-30T13:25:54.882248Z","shell.execute_reply.started":"2025-01-30T13:25:54.859461Z","shell.execute_reply":"2025-01-30T13:25:54.881428Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 6. Dataset Building\n# ---------------------------\ndef build_dataset(paths, offsets=None, labels=None, batch_size=32, cache=True,\n                  decode_fn=None, augment_fn=None, augment=False, repeat=True, shuffle=1024, \n                  drop_remainder=False):\n    if decode_fn is None:\n        decode_fn = build_decoder(labels is not None)\n    \n    if augment_fn is None:\n        augment_fn = build_augmenter()\n    \n    AUTO = tf.data.AUTOTUNE\n    slices = (paths, offsets) if labels is None else (paths, offsets, labels)\n    \n    valid_paths = [path for path in paths if os.path.exists(path)]\n    ds = tf.data.Dataset.from_tensor_slices(slices)\n    ds = ds.map(decode_fn, num_parallel_calls=AUTO)\n    \n    if repeat:\n        ds = ds.repeat()\n    \n    if shuffle:\n        ds = ds.shuffle(shuffle, seed=CFG.seed)\n        \n        # ✅ Correct way to set options\n        options = tf.data.Options()\n        options.experimental_deterministic = False\n        ds = ds.with_options(options)\n    \n    if augment:\n        ds = ds.map(augment_fn, num_parallel_calls=AUTO)\n    \n    ds = ds.batch(batch_size, drop_remainder=drop_remainder)\n    ds = ds.prefetch(AUTO)\n    \n    return ds\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T13:25:54.88432Z","iopub.execute_input":"2025-01-30T13:25:54.884599Z","iopub.status.idle":"2025-01-30T13:25:54.896878Z","shell.execute_reply.started":"2025-01-30T13:25:54.884571Z","shell.execute_reply":"2025-01-30T13:25:54.896137Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 7. Stratified Group K-Fold\n# ---------------------------\nsgkf = StratifiedGroupKFold(n_splits=6, shuffle=True, random_state=CFG.seed)\ndf[\"fold\"] = -1\nfor fold, (_, valid_idx) in enumerate(sgkf.split(df, df[\"class_label\"], groups=df[\"patient_id\"])):\n    df.loc[valid_idx, \"fold\"] = fold\n\ndef filter_valid_samples(df):\n    valid_mask = [os.path.exists(path) for path in df.spec2_path]\n    return df[valid_mask]\n\nsample_df = df.groupby(\"spectrogram_id\").head(1)\ntrain_df = filter_valid_samples(sample_df[sample_df.fold != CFG.fold])\nvalid_df = filter_valid_samples(sample_df[sample_df.fold == CFG.fold])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T13:25:54.897759Z","iopub.execute_input":"2025-01-30T13:25:54.898002Z","iopub.status.idle":"2025-01-30T13:25:55.902925Z","shell.execute_reply.started":"2025-01-30T13:25:54.897983Z","shell.execute_reply":"2025-01-30T13:25:55.902234Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 8. Build ResNet18 Model\n# ---------------------------\ndef resnet_block(x, filters, kernel_size=3, stride=1, conv_shortcut=True):\n    if conv_shortcut:\n        shortcut = keras.layers.Conv2D(filters, 1, strides=stride)(x)\n        shortcut = BatchNormalization()(shortcut)\n    else:\n        shortcut = x\n\n    x = keras.layers.Conv2D(filters, kernel_size, strides=stride, padding='same')(x)\n    x = BatchNormalization()(x)\n    x = keras.layers.Activation('relu')(x)\n    x = keras.layers.Conv2D(filters, kernel_size, padding='same')(x)\n    x = BatchNormalization()(x)\n    x = keras.layers.Add()([shortcut, x])\n    x = keras.layers.Activation('relu')(x)\n    return x\n\ndef build_resnet18(input_shape=(400, 300, 3), num_classes=CFG.num_classes):\n    inputs = keras.Input(shape=input_shape)\n    x = keras.layers.Conv2D(64, 7, strides=2, padding='same')(inputs)\n    x = BatchNormalization()(x)\n    x = keras.layers.Activation('relu')(x)\n    x = keras.layers.MaxPooling2D(3, strides=2, padding='same')(x)\n    \n    x = resnet_block(x, 64, conv_shortcut=True)\n    x = resnet_block(x, 64, conv_shortcut=False)\n    \n    x = resnet_block(x, 128, stride=2)\n    x = resnet_block(x, 128, conv_shortcut=False)\n    \n    x = resnet_block(x, 256, stride=2)\n    x = resnet_block(x, 256, conv_shortcut=False)\n    \n    x = resnet_block(x, 512, stride=2)\n    x = resnet_block(x, 512, conv_shortcut=False)\n    \n    x = GlobalAveragePooling2D()(x)\n    x = Dense(512, activation='relu', kernel_regularizer=keras.regularizers.l2(0.0005))(x)\n    x = BatchNormalization()(x)\n    x = Dropout(0.5)(x)\n    outputs = Dense(num_classes, activation='softmax')(x)\n    \n    model = keras.Model(inputs, outputs)\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T13:25:55.903763Z","iopub.execute_input":"2025-01-30T13:25:55.903983Z","iopub.status.idle":"2025-01-30T13:25:55.913012Z","shell.execute_reply.started":"2025-01-30T13:25:55.903963Z","shell.execute_reply":"2025-01-30T13:25:55.912184Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 9. Compile and Train Model\n# ---------------------------\nif __name__ == \"__main__\":\n    train_ds = build_dataset(\n        train_df.spec2_path.values,\n        train_df.spectrogram_label_offset_seconds.values.astype(int),\n        train_df.class_label.values,\n        batch_size=CFG.batch_size,\n        repeat=True,\n        shuffle=True,\n        augment=True\n    )\n    \n    valid_ds = build_dataset(\n        valid_df.spec2_path.values,\n        valid_df.spectrogram_label_offset_seconds.values.astype(int),\n        valid_df.class_label.values,\n        batch_size=CFG.batch_size,\n        repeat=False,\n        shuffle=False\n    )\n    \n    steps_per_epoch = (len(train_df) + CFG.batch_size - 1) // CFG.batch_size\n    validation_steps = (len(valid_df) + CFG.batch_size - 1) // CFG.batch_size\n\n    lr_schedule = keras.optimizers.schedules.CosineDecayRestarts(\n        initial_learning_rate=0.001,\n        first_decay_steps=steps_per_epoch * 5,  # Adjust decay steps\n        t_mul=2.0,\n        m_mul=0.9\n    )\n\n    \n    model = build_resnet18()\n    model.compile(\n        optimizer=keras.optimizers.Adam(lr_schedule),\n        loss=keras.losses.CategoricalCrossentropy(label_smoothing=0.1),\n        metrics=['accuracy', Precision(), Recall(), AUC(curve='PR', name='aupr')]\n    )\n    \n    callbacks = [\n        ModelCheckpoint(\"best_model.weights.h5\", save_best_only=True, monitor='val_accuracy', mode='max', save_weights_only=True),\n        EarlyStopping(monitor='val_aupr', patience=5, mode='max', restore_best_weights=True)\n    ]\n\n    \n    history = model.fit(\n        train_ds,\n        validation_data=valid_ds,\n        epochs=CFG.epochs,\n        steps_per_epoch=steps_per_epoch,\n        validation_steps=validation_steps,\n        callbacks=callbacks,\n        verbose=CFG.verbose\n    )\n    \n    ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-30T13:25:55.913768Z","iopub.execute_input":"2025-01-30T13:25:55.914043Z","iopub.status.idle":"2025-01-30T14:57:33.565404Z","shell.execute_reply.started":"2025-01-30T13:25:55.914013Z","shell.execute_reply":"2025-01-30T14:57:33.563335Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Evaluation\nmodel = keras.models.load_model(\"best_model_weights.h5\")\ny_true = valid_df.class_label.values\ny_pred = model.predict(valid_ds)\ny_pred_labels = np.argmax(y_pred, axis=1)\n\nprint(\"\\nClassification Report:\")\nreport = classification_report(y_true, y_pred_labels, target_names=CFG.class_names)\nprint(report)\n\nprint(\"\\nClass-wise AUPR:\")\naupr_scores = []\nfor i, name in enumerate(CFG.class_names):\n    precision, recall, _ = precision_recall_curve((y_true == i), y_pred[:, i])\n    aupr = auc(recall, precision)\n    aupr_scores.append(aupr)\n    print(f\"{name}: {aupr:.4f}\")\n\nprint(f\"\\nMean AUPR: {np.mean(aupr_scores):.4f}\")\n\n# Plot Metrics\nplt.figure(figsize=(12, 5))\n\nplt.subplot(1, 2, 1)\nplt.plot(history.history['accuracy'], label='Train Accuracy')\nplt.plot(history.history['val_accuracy'], label='Validation Accuracy')\nplt.title('Accuracy')\nplt.legend()\n\nplt.subplot(1, 2, 2)\nplt.plot(history.history['aupr'], label='Train AUPR')\nplt.plot(history.history['val_aupr'], label='Validation AUPR')\nplt.title('AUPR')\nplt.legend()\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-30T14:59:02.389266Z","iopub.execute_input":"2025-01-30T14:59:02.389568Z","iopub.status.idle":"2025-01-30T14:59:02.449329Z","shell.execute_reply.started":"2025-01-30T14:59:02.389544Z","shell.execute_reply":"2025-01-30T14:59:02.448202Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}