{"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"}],"dockerImageVersionId":30840,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"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.applications import ResNet50\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\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-24T18:05:21.264901Z","iopub.execute_input":"2025-01-24T18:05:21.265222Z","iopub.status.idle":"2025-01-24T18:05:21.271173Z","shell.execute_reply.started":"2025-01-24T18:05:21.265197Z","shell.execute_reply":"2025-01-24T18:05:21.270156Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"###### ---------------------------\n# 2. Configuration\n# ---------------------------\nclass CFG:\n    verbose = 1  # Verbosity\n    seed = 42  # Random seed\n    image_size = [400, 300]  # Input image size\n    epochs = 20  # Increased epochs for better training\n    batch_size = 64  # Batch size\n    lr_mode = \"cos\"  # Learning rate scheduler mode\n    drop_remainder = True  # Drop incomplete batches\n    num_classes = 6  # Number of classes in the dataset\n    fold = 6  # Which fold to set as validation data\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\n# Set random seed\ntf.random.set_seed(CFG.seed)\n","metadata":{"execution":{"iopub.status.busy":"2025-01-24T18:05:21.272266Z","iopub.execute_input":"2025-01-24T18:05:21.272539Z","iopub.status.idle":"2025-01-24T18:05:21.36894Z","shell.execute_reply.started":"2025-01-24T18:05:21.272519Z","shell.execute_reply":"2025-01-24T18:05:21.367956Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 3. Data Preparation\n# ---------------------------\n# Define paths for data\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\n# Load train and test data\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\n# ---------------------------\n# 3. Data Preparation (Updated with Debugging)\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    # Check if .npy file already exists\n    if os.path.exists(npy_file_path):\n        tqdm.write(f\"Spectrogram {spec_id} already exists, skipping.\")\n        return\n    \n    # Check if the .parquet file exists\n    if not os.path.exists(spec_path):\n        tqdm.write(f\"Spectrogram {spec_id} not found at {spec_path}, skipping.\")\n        return\n    \n    try:\n        # Read and process the spectrogram\n        spec = pd.read_parquet(spec_path)\n        spec = spec.fillna(0).values[:, 1:].T  # Fill NaN values and transpose\n        spec = spec.astype(\"float32\")\n        \n        # Save as .npy file\n        np.save(npy_file_path, spec)\n        tqdm.write(f\"Processed and saved spectrogram: {spec_id}\")\n    except Exception as e:\n        tqdm.write(f\"Error processing spectrogram {spec_id}: {e}\")\n# Process train spectrograms\nspec_ids = df[\"spectrogram_id\"].unique()\n_ = joblib.Parallel(n_jobs=-1, backend=\"loky\")(\n    joblib.delayed(process_spec)(spec_id, \"train\") for spec_id in tqdm(spec_ids, total=len(spec_ids), desc=\"Processing train spectrograms\")\n)\n\n# Process test spectrograms\ntest_spec_ids = test_df[\"spectrogram_id\"].unique()\n_ = joblib.Parallel(n_jobs=-1, backend=\"loky\")(\n    joblib.delayed(process_spec)(spec_id, \"test\") for spec_id in tqdm(test_spec_ids, total=len(test_spec_ids), desc=\"Processing test spectrograms\")\n)\n\n# Verify that the files were created\nprint(f\"Number of train spectrograms processed: {len(os.listdir(f'{SPEC_DIR}/train_spectrograms'))}\")\nprint(f\"Number of test spectrograms processed: {len(os.listdir(f'{SPEC_DIR}/test_spectrograms'))}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-24T18:05:21.376065Z","iopub.execute_input":"2025-01-24T18:05:21.37634Z","iopub.status.idle":"2025-01-24T18:05:22.478083Z","shell.execute_reply.started":"2025-01-24T18:05:21.376318Z","shell.execute_reply":"2025-01-24T18:05:22.477084Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 4. Data Augmentation\n# ---------------------------\ndef build_augmenter(dim=CFG.image_size):\n    augmentation = tf.keras.Sequential([\n        RandomFlip(\"horizontal\"),\n        RandomRotation(0.2),\n        RandomZoom(0.2),\n        tf.keras.layers.GaussianNoise(0.1),\n    ])\n    \n    def augment(img, label):\n        if tf.random.uniform([]) < 0.5:\n            img = augmentation(img, training=True)\n        return img, label\n    \n    return augment\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-24T18:05:22.479172Z","iopub.execute_input":"2025-01-24T18:05:22.479424Z","iopub.status.idle":"2025-01-24T18:05:22.484043Z","shell.execute_reply.started":"2025-01-24T18:05:22.479402Z","shell.execute_reply":"2025-01-24T18:05:22.483183Z"}},"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:]  # Remove header tag\n        sig = tf.reshape(sig, [400, -1])\n        \n        if offset is not None:\n            offset = offset // 2  # Only odd values are given\n            sig = sig[:, offset:offset + 300]\n            pad_size = tf.math.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.math.exp(-4.0), tf.math.exp(8.0))  # Avoid 0 in log\n        sig = tf.math.log(sig)\n        \n        sig -= tf.math.reduce_mean(sig)\n        sig /= tf.math.reduce_std(sig) + 1e-6\n        \n        sig = tf.tile(sig[..., None], [1, 1, 3])  # Mono to 3 channels\n        return sig\n    \n    def decode_label(label):\n        label = tf.one_hot(label, CFG.num_classes)\n        label = tf.cast(label, tf.float32)\n        label = tf.reshape(label, [CFG.num_classes])\n        return label\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\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-24T18:05:22.484804Z","iopub.execute_input":"2025-01-24T18:05:22.485017Z","iopub.status.idle":"2025-01-24T18:05:22.499569Z","shell.execute_reply.started":"2025-01-24T18:05:22.484998Z","shell.execute_reply":"2025-01-24T18:05:22.498709Z"}},"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                  cache_dir=\"\", drop_remainder=False):\n    if cache_dir != \"\" and cache is True:\n        os.makedirs(cache_dir, exist_ok=True)\n    \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.experimental.AUTOTUNE\n    slices = (paths, offsets) if labels is None else (paths, offsets, labels)\n    \n    # Filter out paths that do not exist\n    valid_paths = [path for path in paths if os.path.exists(path)]\n    if len(valid_paths) < len(paths):\n        print(f\"Warning: {len(paths) - len(valid_paths)} files are missing and will be skipped.\")\n    \n    ds = tf.data.Dataset.from_tensor_slices((valid_paths, offsets) if labels is None else (valid_paths, offsets, labels))\n    ds = ds.map(decode_fn, num_parallel_calls=AUTO)\n    ds = ds.cache(cache_dir) if cache else ds\n    ds = ds.repeat() if repeat else ds\n    if shuffle:\n        ds = ds.shuffle(shuffle, seed=CFG.seed)\n        opt = tf.data.Options()\n        opt.experimental_deterministic = False\n        ds = ds.with_options(opt)\n    ds = ds.batch(batch_size, drop_remainder=drop_remainder)\n    ds = ds.map(augment_fn, num_parallel_calls=AUTO) if augment else ds\n    ds = ds.prefetch(AUTO)\n    return ds\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-24T18:05:22.500596Z","iopub.execute_input":"2025-01-24T18:05:22.500866Z","iopub.status.idle":"2025-01-24T18:05:22.517001Z","shell.execute_reply.started":"2025-01-24T18:05:22.500843Z","shell.execute_reply":"2025-01-24T18:05:22.51622Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 7. Stratified Group K-Fold (Updated)\n# ---------------------------\nsgkf = StratifiedGroupKFold(n_splits=10, shuffle=True, random_state=CFG.seed)\n\ndf[\"fold\"] = -1\ndf.reset_index(drop=True, inplace=True)\nfor fold, (train_idx, valid_idx) in enumerate(sgkf.split(df, y=df[\"class_label\"], groups=df[\"patient_id\"])):\n    df.loc[valid_idx, \"fold\"] = fold\n\ndef filter_valid_samples(df):\n    \"\"\"Filter dataframe to only include valid existing paths\"\"\"\n    valid_mask = [os.path.exists(path) for path in df.spec2_path]\n    return df[valid_mask].reset_index(drop=True)\n\n# Sample and filter data\nsample_df = df.groupby(\"spectrogram_id\").head(1).reset_index(drop=True)\ntrain_df = filter_valid_samples(sample_df[sample_df.fold != CFG.fold])\nvalid_df = filter_valid_samples(sample_df[sample_df.fold == CFG.fold])\n\n# Verify we have valid data\nassert len(train_df) > 0, \"No valid training samples found!\"\nassert len(valid_df) > 0, \"No valid validation samples found!\"\nprint(f\"\\nTraining samples: {len(train_df)}\")\nprint(f\"Validation samples: {len(valid_df)}\\n\")\n\n# Create datasets with validated paths\ntrain_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\nvalid_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    augment=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-24T18:05:22.519186Z","iopub.execute_input":"2025-01-24T18:05:22.519396Z","iopub.status.idle":"2025-01-24T18:05:24.609094Z","shell.execute_reply.started":"2025-01-24T18:05:22.519375Z","shell.execute_reply":"2025-01-24T18:05:24.608215Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 8. Build ResNet50 Model\n# ---------------------------\ndef build_resnet50(input_shape=(400, 300, 3), num_classes=CFG.num_classes):\n    base_model = ResNet50(weights='imagenet', include_top=False, input_shape=input_shape)\n    base_model.trainable = True  # Fine-tune the entire model\n    \n    inputs = keras.Input(shape=input_shape)\n    x = base_model(inputs, training=True)\n    x = GlobalAveragePooling2D()(x)\n    x = Dense(256, activation='relu', kernel_regularizer=tf.keras.regularizers.l2(0.01))(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\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-24T18:05:24.61091Z","iopub.execute_input":"2025-01-24T18:05:24.611283Z","iopub.status.idle":"2025-01-24T18:05:24.616825Z","shell.execute_reply.started":"2025-01-24T18:05:24.61125Z","shell.execute_reply":"2025-01-24T18:05:24.615903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 9. Compile and Train Model (Updated)\n# ---------------------------\nif __name__ == \"__main__\":\n    \n    # Build and compile model\n    model = build_resnet50(input_shape=(CFG.image_size[0], CFG.image_size[1], 3), num_classes=CFG.num_classes)\n    model.compile(\n        optimizer=keras.optimizers.Adam(learning_rate=0.0005),\n        loss='categorical_crossentropy',\n        metrics=['accuracy', Precision(name='precision'), Recall(name='recall'), AUC(name='aupr', curve='PR')]\n    )\n    model.summary()\n\n    # Calculate steps using CEILING DIVISION\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    print(f\"\\nSteps per epoch: {steps_per_epoch}\")\n    print(f\"Validation steps: {validation_steps}\\n\")\n\n    # Define callbacks\n    callbacks = [\n        ModelCheckpoint(\"best_resnet_model.keras\", save_best_only=True, monitor=\"val_aupr\", mode='max'),\n        ReduceLROnPlateau(monitor=\"val_aupr\", factor=0.3, patience=3, verbose=1, mode='max'),\n        EarlyStopping(monitor=\"val_aupr\", patience=12, verbose=1, mode='max')\n    ]\n\n    # Train the model\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     # Evaluate and plot results\n    loss, accuracy, precision, recall, aupr = model.evaluate(valid_ds)\n    print(f\"Validation Loss: {loss}\")\n    print(f\"Validation Accuracy: {accuracy}\")\n    print(f\"Validation AUPR: {aupr}\")\n\n    # Calculate additional metrics\n    print(\"\\nCalculating detailed metrics...\")\n    \n    # Get predictions and true labels\n    y_true = valid_df.class_label.values\n    y_pred = model.predict(valid_ds)\n    y_pred_labels = np.argmax(y_pred, axis=1)\n    \n    # Weighted F1 Score\n    report = classification_report(\n        y_true, \n        y_pred_labels,\n        target_names=CFG.class_names,\n        output_dict=True\n    )\n    weighted_f1 = report['weighted avg']['f1-score']\n    print(f\"\\nWeighted F1 Score: {weighted_f1:.4f}\")\n    \n    # Per-class AUPR\n    print(\"\\nClass-wise AUPR Scores:\")\n    aupr_scores = []\n    for i, class_name in enumerate(CFG.class_names):\n        precision, recall, _ = precision_recall_curve(\n            (y_true == i).astype(int),\n            y_pred[:, i]\n        )\n        class_aupr = auc(recall, precision)\n        aupr_scores.append(class_aupr)\n        print(f\"{class_name}: {class_aupr:.4f}\")\n    \n    # Mean AUPR\n    mean_aupr = np.mean(aupr_scores)\n    print(f\"\\nMean AUPR: {mean_aupr:.4f}\")\n\n    # Plot training and validation metrics\n    def plot_metrics(history):\n        # [Previous plotting code remains the same]\n        plt.figure(figsize=(12, 6))\n        \n        # Plot loss\n        plt.subplot(1, 2, 1)\n        plt.plot(history.history['loss'], label='Train Loss')\n        plt.plot(history.history['val_loss'], label='Validation Loss')\n        plt.title('Loss')\n        plt.xlabel('Epoch')\n        plt.ylabel('Loss')\n        plt.legend()\n        \n        # Plot AUPR\n        plt.subplot(1, 2, 2)\n        plt.plot(history.history['aupr'], label='Train AUPR')\n        plt.plot(history.history['val_aupr'], label='Validation AUPR')\n        plt.title('AUPR')\n        plt.xlabel('Epoch')\n        plt.ylabel('AUPR')\n        plt.legend()\n        \n        plt.tight_layout()\n        plt.show()\n    \n    plot_metrics(history)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-24T18:05:24.6177Z","iopub.execute_input":"2025-01-24T18:05:24.617928Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Weighted F1 Score\nreport = classification_report(\n    y_true, \n    y_pred_labels,\n    target_names=CFG.class_names,\n    output_dict=True\n)\nweighted_f1 = report['weighted avg']['f1-score']\nprint(f\"\\nWeighted F1 Score: {weighted_f1:.4f}\")\n\n# Per-class AUPR\nprint(\"\\nClass-wise AUPR Scores:\")\naupr_scores = []\nfor i, class_name in enumerate(CFG.class_names):\n    precision, recall, _ = precision_recall_curve(\n        (y_true == i).astype(int),\n        y_pred[:, i]\n    )\n    class_aupr = auc(recall, precision)\n    aupr_scores.append(class_aupr)\n    print(f\"{class_name}: {class_aupr:.4f}\")\n    \n# Mean AUPR\nmean_aupr = np.mean(aupr_scores)\nprint(f\"\\nMean AUPR: {mean_aupr:.4f}\")\n\n# Plot training and validation metrics\ndef plot_metrics(history):\n    \n    # [Previous plotting code remains the same]\n    plt.figure(figsize=(12, 6))\n      \n    # Plot loss\n    plt.subplot(1, 2, 1)\n    plt.plot(history.history['loss'], label='Train Loss')\n    plt.plot(history.history['val_loss'], label='Validation Loss')\n    plt.title('Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n        \n    # Plot AUPR\n    plt.subplot(1, 2, 2)\n    plt.plot(history.history['aupr'], label='Train AUPR')\n    plt.plot(history.history['val_aupr'], label='Validation AUPR')\n    plt.title('AUPR')\n    plt.xlabel('Epoch')\n    plt.ylabel('AUPR')\n    plt.legend()\n     \n    plt.tight_layout()\n    plt.show()\n    \nplot_metrics(history)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-24T19:02:24.688288Z","iopub.execute_input":"2025-01-24T19:02:24.688595Z","iopub.status.idle":"2025-01-24T19:02:25.111685Z","shell.execute_reply.started":"2025-01-24T19:02:24.688572Z","shell.execute_reply":"2025-01-24T19:02:25.11086Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 10. Prepare Test Dataset\n# ---------------------------\n# Load test data\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\n# Process test spectrograms\ntest_spec_ids = test_df[\"spectrogram_id\"].unique()\n_ = joblib.Parallel(n_jobs=-1, backend=\"loky\")(\n    joblib.delayed(process_spec)(spec_id, \"test\") for spec_id in tqdm(test_spec_ids, total=len(test_spec_ids), desc=\"Processing test spectrograms\")\n)\n\n# Build test dataset\ntest_paths = test_df.spec2_path.values\ntest_offsets = np.zeros(len(test_df), dtype=int)  # No offset for test data\ntest_ds = build_dataset(test_paths, test_offsets, batch_size=CFG.batch_size, repeat=False, shuffle=False, augment=False, cache=True)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------\n# 11. Generate Predictions\n# ---------------------------\n# Load the best model\nmodel = keras.models.load_model(\"best_resnet_model.keras\")\n\n# Make predictions\npredictions = model.predict(test_ds, verbose=CFG.verbose)\n\n# Format predictions for submission\nsubmission_df = pd.DataFrame({\n    'eeg_id': test_df['eeg_id'],\n    'seizure_vote': predictions[:, 0],\n    'lpd_vote': predictions[:, 1],\n    'gpd_vote': predictions[:, 2],\n    'lrda_vote': predictions[:, 3],\n    'grda_vote': predictions[:, 4],\n    'other_vote': predictions[:, 5]\n})\n\n# Save submission file\nsubmission_df.to_csv(\"submission.csv\", index=False)\nprint(\"Submission file saved as submission.csv\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}