{"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":"import 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 VGG19\nfrom tensorflow.keras.layers import RandomFlip, RandomRotation, RandomZoom\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom sklearn.metrics import f1_score, average_precision_score, classification_report","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:19:27.084543Z","iopub.execute_input":"2025-01-29T18:19:27.084855Z","iopub.status.idle":"2025-01-29T18:19:39.200075Z","shell.execute_reply.started":"2025-01-29T18:19:27.084828Z","shell.execute_reply":"2025-01-29T18:19:39.199412Z"}},"outputs":[],"execution_count":1},{"cell_type":"code","source":"# Add at the very top of your script\n# tf.keras.mixed_precision.set_global_policy('mixed_float16')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:19:39.201191Z","iopub.execute_input":"2025-01-29T18:19:39.201957Z","iopub.status.idle":"2025-01-29T18:19:39.205274Z","shell.execute_reply.started":"2025-01-29T18:19:39.201924Z","shell.execute_reply":"2025-01-29T18:19:39.204398Z"}},"outputs":[],"execution_count":2},{"cell_type":"code","source":"class CFG:\n    verbose = 1\n    seed = 42\n    preset = \"vgg19\"\n    image_size = [400, 300]\n    epochs = 20  # Increased epochs for better convergence\n    batch_size = 48  # More stable updates  # Reduced batch size for better memory management\n    lr_mode = \"cosine_warmup\"\n    label_smoothing = 0.1  # Added to prevent overconfidence\n    mixup_alpha = 0.2  # New regularization\n    drop_remainder = True\n    num_classes = 6\n    fold = 0\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-29T18:19:39.207036Z","iopub.execute_input":"2025-01-29T18:19:39.207341Z","iopub.status.idle":"2025-01-29T18:19:39.234796Z","shell.execute_reply.started":"2025-01-29T18:19:39.207314Z","shell.execute_reply":"2025-01-29T18:19:39.234039Z"}},"outputs":[],"execution_count":3},{"cell_type":"code","source":"# Data loading and processing functions (same as original)\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\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'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:19:39.23617Z","iopub.execute_input":"2025-01-29T18:19:39.236448Z","iopub.status.idle":"2025-01-29T18:19:39.623024Z","shell.execute_reply.started":"2025-01-29T18:19:39.236426Z","shell.execute_reply":"2025-01-29T18:19:39.622382Z"}},"outputs":[],"execution_count":4},{"cell_type":"code","source":"# Function to process spectrograms and check if .npy file already exists\ndef process_spec(spec_id, split=\"train\"):\n    spec_path = f\"{BASE_PATH}/{split}_spectrograms/{spec_id}.parquet\"\n    spec = pd.read_parquet(spec_path)\n    spec = spec.fillna(0).values[:, 1:].T  # fill NaN values with 0, transpose for (Time, Freq) -> (Freq, Time)\n    spec = spec.astype(\"float32\")\n    \n    npy_file_path = f\"{SPEC_DIR}/{split}_spectrograms/{spec_id}.npy\"\n    \n    # Check if .npy file already exists\n    if not os.path.exists(npy_file_path):\n        np.save(npy_file_path, spec)\n        tqdm.write(f\"Processed and saved spectrogram: {spec_id}\")  # Use tqdm.write to log without interfering with the progress bar\n    else:\n        tqdm.write(f\"Spectrogram {spec_id} already exists, skipping.\")  # Use tqdm.write to avoid print flood\n# Parallelize the processing of spectrograms\nspec_ids = df[\"spectrogram_id\"].unique()\n\n# Using joblib for parallel processing with tqdm progress\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# Repeat the same for the test dataset\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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:19:39.62382Z","iopub.execute_input":"2025-01-29T18:19:39.624127Z","iopub.status.idle":"2025-01-29T18:21:57.761837Z","shell.execute_reply.started":"2025-01-29T18:19:39.624097Z","shell.execute_reply":"2025-01-29T18:21:57.760995Z"}},"outputs":[{"name":"stderr","text":"Processing train spectrograms: 100%|██████████| 11138/11138 [02:17<00:00, 80.88it/s]\nProcessing test spectrograms: 100%|██████████| 1/1 [00:00<00:00, 697.66it/s]\n","output_type":"stream"}],"execution_count":5},{"cell_type":"code","source":"def build_augmenter():\n    augmentation = tf.keras.Sequential([\n        RandomFlip(\"horizontal_and_vertical\"),\n        RandomRotation(0.25),\n        RandomZoom(0.25),\n        tf.keras.layers.RandomContrast(0.15),\n        tf.keras.layers.GaussianNoise(0.1),\n        tf.keras.layers.RandomBrightness(0.1),\n    ])\n    \n    def augment(img, label):\n        # Remove explicit training=True argument\n        img = augmentation(img)\n        return img, label\n    \n    return augment","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:28:03.519815Z","iopub.execute_input":"2025-01-29T18:28:03.520135Z","iopub.status.idle":"2025-01-29T18:28:03.524971Z","shell.execute_reply.started":"2025-01-29T18:28:03.520077Z","shell.execute_reply":"2025-01-29T18:28:03.524059Z"}},"outputs":[],"execution_count":12},{"cell_type":"code","source":"# Decoder function for loading and processing spectrogram data\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        sig = tf.image.resize(sig, [400, 300]) \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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:28:04.449962Z","iopub.execute_input":"2025-01-29T18:28:04.450223Z","iopub.status.idle":"2025-01-29T18:28:04.458012Z","shell.execute_reply.started":"2025-01-29T18:28:04.450199Z","shell.execute_reply":"2025-01-29T18:28:04.457304Z"}},"outputs":[],"execution_count":13},{"cell_type":"code","source":"def build_dataset(paths, offsets=None, labels=None, batch_size=32, cache=True,\n                  decode_fn=None, augment_fn=None, augment=False, repeat=True, \n                  shuffle=256, cache_dir=\"\", drop_remainder=False):\n    \n    AUTO = tf.data.AUTOTUNE\n    slices = (paths, offsets) if labels is None else (paths, offsets, labels)\n    \n    ds = tf.data.Dataset.from_tensor_slices(slices)\n    \n    if decode_fn is None:\n        decode_fn = build_decoder(labels is not None)\n    ds = ds.map(decode_fn, num_parallel_calls=AUTO)\n    \n    if cache:\n        ds = ds.cache(cache_dir if cache_dir else \"\")\n    \n    if repeat:\n        ds = ds.repeat()\n    \n    if shuffle:\n        ds = ds.shuffle(shuffle, seed=CFG.seed)\n        options = tf.data.Options()\n        options.experimental_deterministic = False\n        ds = ds.with_options(options)\n    \n    ds = ds.batch(batch_size, drop_remainder=drop_remainder)\n    \n    if augment:\n        if augment_fn is None:\n            # ✅ Fixed: Create proper augmentation wrapper\n            augmentation_model = tf.keras.Sequential([\n                RandomFlip(\"horizontal_and_vertical\"),\n                RandomRotation(0.25),\n                RandomZoom(0.25),\n                tf.keras.layers.RandomContrast(0.15),\n                tf.keras.layers.GaussianNoise(0.1),\n                tf.keras.layers.RandomBrightness(0.1),\n            ])\n            \n            def augment_wrapper(images, labels):\n                return augmentation_model(images, training=True), labels\n            \n            augment_fn = augment_wrapper\n            \n        # ✅ Apply augmentations in correct order\n        ds = ds.map(\n            lambda x, y: (augment_fn(x, y)[0], y),  # Only augment images\n            num_parallel_calls=AUTO\n        )\n        ds = ds.map(mixup_images, num_parallel_calls=AUTO)\n    \n    ds = ds.prefetch(AUTO)\n    return ds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:39:33.825664Z","iopub.execute_input":"2025-01-29T18:39:33.825988Z","iopub.status.idle":"2025-01-29T18:39:33.833613Z","shell.execute_reply.started":"2025-01-29T18:39:33.825962Z","shell.execute_reply":"2025-01-29T18:39:33.832762Z"}},"outputs":[],"execution_count":16},{"cell_type":"code","source":"# Stratified group K-fold\nsgkf = StratifiedGroupKFold(n_splits=5, 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\n# Sample data\nsample_df = df.groupby(\"spectrogram_id\").head(1).reset_index(drop=True)\ntrain_df = sample_df[sample_df.fold != CFG.fold]\nvalid_df = sample_df[sample_df.fold == CFG.fold]\n\n# Train and Validation datasets\ntrain_paths = train_df.spec2_path.values\ntrain_offsets = train_df.spectrogram_label_offset_seconds.values.astype(int)\ntrain_labels = train_df.class_label.values\ntrain_ds = build_dataset(train_paths, train_offsets, train_labels, batch_size=CFG.batch_size,\n                         repeat=True, shuffle=True, augment=True, cache=False)\n\nvalid_paths = valid_df.spec2_path.values\nvalid_offsets = valid_df.spectrogram_label_offset_seconds.values.astype(int)\nvalid_labels = valid_df.class_label.values\nvalid_ds = build_dataset(valid_paths, valid_offsets, valid_labels, batch_size=CFG.batch_size,\n                         repeat=False, shuffle=False, augment=False, cache=False)\n\ntrain_ds = train_ds.repeat().shuffle(256)\nvalid_ds = valid_ds.repeat()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:39:37.324876Z","iopub.execute_input":"2025-01-29T18:39:37.325206Z","iopub.status.idle":"2025-01-29T18:39:38.553716Z","shell.execute_reply.started":"2025-01-29T18:39:37.325176Z","shell.execute_reply":"2025-01-29T18:39:38.553013Z"}},"outputs":[],"execution_count":17},{"cell_type":"code","source":"# # Augmented training images display\n# imgs, tars = next(iter(train_ds))\n# num_imgs = 8\n# plt.figure(figsize=(4*4, num_imgs // 4 * 5))\n# for i in range(num_imgs):\n#     plt.subplot(num_imgs // 4, 4, i + 1)\n#     img = imgs[i].numpy()[..., 0]\n#     img -= img.min()\n#     img /= img.max() + 1e-4\n#     tar = CFG.label2name[np.argmax(tars[i].numpy())]\n#     plt.imshow(img)\n#     plt.title(f\"Target: {tar}\")\n#     plt.axis('off')\n# plt.tight_layout()\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:22:01.842062Z","iopub.status.idle":"2025-01-29T18:22:01.842376Z","shell.execute_reply":"2025-01-29T18:22:01.842258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_vgg19(input_shape=(400, 300, 3), num_classes=6):\n    inputs = keras.Input(shape=input_shape)\n    \n    # Improved preprocessing\n    x = keras.layers.Rescaling(1./255)(inputs)\n    x = keras.layers.RandomContrast(0.1)(x)  # Added contrast augmentation\n    \n    # Base model with modified unfreezing\n    base_model = VGG19(\n        include_top=False,\n        weights='imagenet',\n        input_shape=input_shape,\n        pooling='avg'  # Better for feature extraction\n    )\n    \n    # Progressive unfreezing strategy\n    base_model.trainable = True\n    for layer in base_model.layers[:-15]:  # Unfreeze more layers\n        layer.trainable = False\n    \n    # Add intermediate layers\n    x = base_model(x)\n    x = keras.layers.BatchNormalization()(x)\n    x = keras.layers.Dense(1024, activation='swish', kernel_regularizer=keras.regularizers.l2(1e-4))(x)\n    x = keras.layers.Dropout(0.5)(x)  # Reduced from 0.7\n    x = keras.layers.Dense(512, activation='swish', kernel_regularizer=keras.regularizers.l2(1e-4))(x)\n    x = keras.layers.Dropout(0.3)(x)\n    outputs = keras.layers.Dense(num_classes, activation='softmax')(x)\n    \n    return keras.Model(inputs, outputs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:39:41.514859Z","iopub.execute_input":"2025-01-29T18:39:41.515183Z","iopub.status.idle":"2025-01-29T18:39:41.521561Z","shell.execute_reply.started":"2025-01-29T18:39:41.515155Z","shell.execute_reply":"2025-01-29T18:39:41.520696Z"}},"outputs":[],"execution_count":18},{"cell_type":"code","source":"# Custom learning rate schedule\ndef lr_scheduler(epoch):\n    warmup_epochs = 10\n    base_lr = 1e-4  # Increased from 1e-5\n    decay_factor = 0.1\n    \n    if epoch < warmup_epochs:\n        return base_lr * (epoch + 1) / warmup_epochs\n    return base_lr * decay_factor ** ((epoch - warmup_epochs) // 5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:39:44.259842Z","iopub.execute_input":"2025-01-29T18:39:44.260162Z","iopub.status.idle":"2025-01-29T18:39:44.264367Z","shell.execute_reply.started":"2025-01-29T18:39:44.260131Z","shell.execute_reply":"2025-01-29T18:39:44.263512Z"}},"outputs":[],"execution_count":19},{"cell_type":"code","source":"# Focal loss for class imbalance\ndef focal_loss(y_true, y_pred):\n    gamma = 2.0\n    alpha = tf.constant([0.3, 0.2, 0.15, 0.15, 0.1, 0.1], dtype=tf.float32)  # Class weights\n    epsilon = keras.backend.epsilon()\n    y_pred = keras.backend.clip(y_pred, epsilon, 1. - epsilon)\n    pt = tf.where(keras.backend.equal(y_true, 1), y_pred, 1 - y_pred)\n    return -keras.backend.sum(alpha * keras.backend.pow(1. - pt, gamma) * keras.backend.log(pt), axis=-1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:39:45.814514Z","iopub.execute_input":"2025-01-29T18:39:45.814809Z","iopub.status.idle":"2025-01-29T18:39:45.819611Z","shell.execute_reply.started":"2025-01-29T18:39:45.814783Z","shell.execute_reply":"2025-01-29T18:39:45.818812Z"}},"outputs":[],"execution_count":20},{"cell_type":"code","source":"# Custom metrics\nclass F1Score(tf.keras.metrics.Metric):\n    def __init__(self, name='f1_score', **kwargs):\n        super().__init__(name=name, **kwargs)\n        self.precision = tf.keras.metrics.Precision()\n        self.recall = tf.keras.metrics.Recall()\n\n    def update_state(self, y_true, y_pred, sample_weight=None):\n        y_pred = tf.argmax(y_pred, axis=1)\n        y_true = tf.argmax(y_true, axis=1)\n        self.precision.update_state(y_true, y_pred)\n        self.recall.update_state(y_true, y_pred)\n\n    def result(self):\n        p = self.precision.result()\n        r = self.recall.result()\n        return 2 * ((p * r) / (p + r + 1e-6))\n\n    def reset_state(self):\n        self.precision.reset_state()\n        self.recall.reset_state()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:41:24.92524Z","iopub.execute_input":"2025-01-29T18:41:24.925533Z","iopub.status.idle":"2025-01-29T18:41:24.931387Z","shell.execute_reply.started":"2025-01-29T18:41:24.92551Z","shell.execute_reply":"2025-01-29T18:41:24.93047Z"}},"outputs":[],"execution_count":21},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    # Build and compile model\n    model = build_vgg19()\n    # Update model compilation\n    model.compile(\n        optimizer=keras.optimizers.AdamW(learning_rate=1e-4, weight_decay=1e-4),\n        loss=focal_loss,\n        metrics=['accuracy', F1Score()]\n    )\n\n    model.summary()\n\n    # Compute class weights\n    train_labels = train_df.class_label.values\n    class_weights = compute_class_weight('balanced', classes=np.unique(train_labels), y=train_labels)\n    class_weights = class_weights * 2  # More aggressive reweighting\n    class_weights_dict = {0: 4, 1: 6, 2: 7, 3: 7, 4: 5, 5: 2}  # More conservative weights\n\n\n    # Callbacks\n    # Enhanced callbacks\n    callbacks = [\n        keras.callbacks.ModelCheckpoint(\"best_model.keras\", \n                                      monitor='val_accuracy',  # Changed focus\n                                      mode='max',\n                                      save_best_only=True),\n        keras.callbacks.EarlyStopping(\n            monitor='val_accuracy',\n            patience=15,\n            min_delta=0.001,\n            baseline=0.4,\n            mode='max'\n        ),\n        keras.callbacks.ReduceLROnPlateau(\n            monitor='val_loss',\n            factor=0.5,\n            patience=5,\n            min_lr=1e-7\n        ),\n        keras.callbacks.TerminateOnNaN()\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=len(train_df) // CFG.batch_size,\n        validation_steps=len(valid_df) // CFG.batch_size,\n        class_weight=class_weights_dict,\n        callbacks=callbacks,\n        verbose=CFG.verbose\n    )\n\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:41:27.405018Z","iopub.execute_input":"2025-01-29T18:41:27.405328Z"}},"outputs":[{"name":"stdout","text":"Downloading data from https://storage.googleapis.com/tensorflow/keras-applications/vgg19/vgg19_weights_tf_dim_ordering_tf_kernels_notop.h5\n\u001b[1m80134624/80134624\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m0s\u001b[0m 0us/step\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"\u001b[1mModel: \"functional_2\"\u001b[0m\n","text/html":"<pre style=\"white-space:pre;overflow-x:auto;line-height:normal;font-family:Menlo,'DejaVu Sans Mono',consolas,'Courier New',monospace\"><span style=\"font-weight: bold\">Model: \"functional_2\"</span>\n</pre>\n"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"┏━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━┓\n┃\u001b[1m \u001b[0m\u001b[1mLayer (type)                        \u001b[0m\u001b[1m \u001b[0m┃\u001b[1m \u001b[0m\u001b[1mOutput Shape               \u001b[0m\u001b[1m \u001b[0m┃\u001b[1m \u001b[0m\u001b[1m        Param #\u001b[0m\u001b[1m \u001b[0m┃\n┡━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━━━┩\n│ input_layer_2 (\u001b[38;5;33mInputLayer\u001b[0m)           │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m400\u001b[0m, \u001b[38;5;34m300\u001b[0m, \u001b[38;5;34m3\u001b[0m)         │               \u001b[38;5;34m0\u001b[0m │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ rescaling (\u001b[38;5;33mRescaling\u001b[0m)                │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m400\u001b[0m, \u001b[38;5;34m300\u001b[0m, \u001b[38;5;34m3\u001b[0m)         │               \u001b[38;5;34m0\u001b[0m │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ random_contrast_2 (\u001b[38;5;33mRandomContrast\u001b[0m)   │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m400\u001b[0m, \u001b[38;5;34m300\u001b[0m, \u001b[38;5;34m3\u001b[0m)         │               \u001b[38;5;34m0\u001b[0m │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ vgg19 (\u001b[38;5;33mFunctional\u001b[0m)                   │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m512\u001b[0m)                 │      \u001b[38;5;34m20,024,384\u001b[0m │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ batch_normalization                  │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m512\u001b[0m)                 │           \u001b[38;5;34m2,048\u001b[0m │\n│ (\u001b[38;5;33mBatchNormalization\u001b[0m)                 │                             │                 │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ dense (\u001b[38;5;33mDense\u001b[0m)                        │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m1024\u001b[0m)                │         \u001b[38;5;34m525,312\u001b[0m │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ dropout (\u001b[38;5;33mDropout\u001b[0m)                    │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m1024\u001b[0m)                │               \u001b[38;5;34m0\u001b[0m │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ dense_1 (\u001b[38;5;33mDense\u001b[0m)                      │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m512\u001b[0m)                 │         \u001b[38;5;34m524,800\u001b[0m │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ dropout_1 (\u001b[38;5;33mDropout\u001b[0m)                  │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m512\u001b[0m)                 │               \u001b[38;5;34m0\u001b[0m │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ dense_2 (\u001b[38;5;33mDense\u001b[0m)                      │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m6\u001b[0m)                   │           \u001b[38;5;34m3,078\u001b[0m │\n└──────────────────────────────────────┴─────────────────────────────┴─────────────────┘\n","text/html":"<pre style=\"white-space:pre;overflow-x:auto;line-height:normal;font-family:Menlo,'DejaVu Sans Mono',consolas,'Courier New',monospace\">┏━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━┓\n┃<span style=\"font-weight: bold\"> Layer (type)                         </span>┃<span style=\"font-weight: bold\"> Output Shape                </span>┃<span style=\"font-weight: bold\">         Param # </span>┃\n┡━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━━━┩\n│ input_layer_2 (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">InputLayer</span>)           │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">400</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">300</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">3</span>)         │               <span style=\"color: #00af00; text-decoration-color: #00af00\">0</span> │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ rescaling (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Rescaling</span>)                │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">400</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">300</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">3</span>)         │               <span style=\"color: #00af00; text-decoration-color: #00af00\">0</span> │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ random_contrast_2 (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">RandomContrast</span>)   │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">400</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">300</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">3</span>)         │               <span style=\"color: #00af00; text-decoration-color: #00af00\">0</span> │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ vgg19 (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Functional</span>)                   │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">512</span>)                 │      <span style=\"color: #00af00; text-decoration-color: #00af00\">20,024,384</span> │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ batch_normalization                  │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">512</span>)                 │           <span style=\"color: #00af00; text-decoration-color: #00af00\">2,048</span> │\n│ (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">BatchNormalization</span>)                 │                             │                 │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ dense (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Dense</span>)                        │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">1024</span>)                │         <span style=\"color: #00af00; text-decoration-color: #00af00\">525,312</span> │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ dropout (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Dropout</span>)                    │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">1024</span>)                │               <span style=\"color: #00af00; text-decoration-color: #00af00\">0</span> │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ dense_1 (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Dense</span>)                      │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">512</span>)                 │         <span style=\"color: #00af00; text-decoration-color: #00af00\">524,800</span> │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ dropout_1 (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Dropout</span>)                  │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">512</span>)                 │               <span style=\"color: #00af00; text-decoration-color: #00af00\">0</span> │\n├──────────────────────────────────────┼─────────────────────────────┼─────────────────┤\n│ dense_2 (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Dense</span>)                      │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">6</span>)                   │           <span style=\"color: #00af00; text-decoration-color: #00af00\">3,078</span> │\n└──────────────────────────────────────┴─────────────────────────────┴─────────────────┘\n</pre>\n"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"\u001b[1m Total params: \u001b[0m\u001b[38;5;34m21,079,622\u001b[0m (80.41 MB)\n","text/html":"<pre style=\"white-space:pre;overflow-x:auto;line-height:normal;font-family:Menlo,'DejaVu Sans Mono',consolas,'Courier New',monospace\"><span style=\"font-weight: bold\"> Total params: </span><span style=\"color: #00af00; text-decoration-color: #00af00\">21,079,622</span> (80.41 MB)\n</pre>\n"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"\u001b[1m Trainable params: \u001b[0m\u001b[38;5;34m20,523,270\u001b[0m (78.29 MB)\n","text/html":"<pre style=\"white-space:pre;overflow-x:auto;line-height:normal;font-family:Menlo,'DejaVu Sans Mono',consolas,'Courier New',monospace\"><span style=\"font-weight: bold\"> Trainable params: </span><span style=\"color: #00af00; text-decoration-color: #00af00\">20,523,270</span> (78.29 MB)\n</pre>\n"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"\u001b[1m Non-trainable params: \u001b[0m\u001b[38;5;34m556,352\u001b[0m (2.12 MB)\n","text/html":"<pre style=\"white-space:pre;overflow-x:auto;line-height:normal;font-family:Menlo,'DejaVu Sans Mono',consolas,'Courier New',monospace\"><span style=\"font-weight: bold\"> Non-trainable params: </span><span style=\"color: #00af00; text-decoration-color: #00af00\">556,352</span> (2.12 MB)\n</pre>\n"},"metadata":{}},{"name":"stdout","text":"Epoch 1/20\n","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"# Evaluation\nprint(\"\\nGenerating validation predictions...\")\ny_true = []\ny_pred_probs = []\n\nfor x, y in tqdm(valid_ds, total=len(valid_df) // CFG.batch_size + 1):\n    y_true.append(y.numpy())\n    y_pred_probs.append(model.predict(x, verbose=0))\n\ny_true = np.concatenate(y_true)\ny_pred_probs = np.concatenate(y_pred_probs)\ny_pred_labels = np.argmax(y_pred_probs, axis=1)\ny_true_labels = np.argmax(y_true, axis=1)\n\n# Calculate metrics\nweighted_f1 = f1_score(y_true_labels, y_pred_labels, average='weighted')\nprint(f\"\\nWeighted F1 Score: {weighted_f1:.4f}\")\n\nprint(\"\\nAUPR Scores:\")\naupr_scores = {}\nfor i, class_name in enumerate(CFG.class_names):\n    ap = average_precision_score(y_true[:, i], y_pred_probs[:, i])\n    aupr_scores[class_name] = ap\n    print(f\"{class_name}: {ap:.4f}\")\n\nprint(\"\\nClassification Report:\")\nprint(classification_report(y_true_labels, y_pred_labels, \n                            target_names=CFG.class_names, digits=4))\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:22:01.848462Z","iopub.status.idle":"2025-01-29T18:22:01.848791Z","shell.execute_reply":"2025-01-29T18:22:01.848623Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 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['loss'], label='Train Loss')\nplt.plot(history.history['val_loss'], label='Validation Loss')\nplt.title('Loss')\nplt.legend()\n\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-29T18:22:01.849716Z","iopub.status.idle":"2025-01-29T18:22:01.850025Z","shell.execute_reply":"2025-01-29T18:22:01.849909Z"}},"outputs":[],"execution_count":null}]}