{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":13964866,"sourceType":"datasetVersion","datasetId":8902183},{"sourceId":14168981,"sourceType":"datasetVersion","datasetId":9031621}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Data Exploration","metadata":{}},{"cell_type":"code","source":"import os #For file paths\nimport keras_cv #For audio spectrograms and data augmentation. This has some prebuilt models\nimport keras#Main deep learning framework you’re using to build models, layers, and training loops\nimport keras.backend as K #“Low-level” backend ops used inside/around Keras models\nimport tensorflow as tf #TensorFlow is one of the possible “backends” Keras can run on,\nimport tensorflow_io as tfio #For audio tasks,\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nfrom tensorflow.keras.optimizers.schedules import CosineDecay\n\nimport numpy as np \nimport pandas as pd\n\nfrom glob import glob\nfrom tqdm import tqdm\n\nimport librosa #Core audio processing library in Python.\nimport IPython.display as ipd #For Jupyter/Colab display utilities.\nimport librosa.display as lid #Plotting helpers for audio from librosa.\n\nimport matplotlib.pyplot as plt\nimport matplotlib as mpl\n\nimport random\nfrom sklearn.model_selection import train_test_split\nimport os\nfrom tqdm import tqdm\nos.environ[\"KERAS_BACKEND\"] = \"torch\"  # \"jax\" or \"tensorflow\" or \"torch\" \n\nfrom collections import defaultdict","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:02.890645Z","iopub.execute_input":"2025-12-16T17:52:02.891631Z","iopub.status.idle":"2025-12-16T17:52:02.89792Z","shell.execute_reply.started":"2025-12-16T17:52:02.891604Z","shell.execute_reply":"2025-12-16T17:52:02.896907Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_meta = pd.read_csv('/kaggle/input/birdclef-2024/train_metadata.csv')\neBird_Taxonomy = pd.read_csv('/kaggle/input/birdclef-2024/eBird_Taxonomy_v2021.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:02.898917Z","iopub.execute_input":"2025-12-16T17:52:02.899152Z","iopub.status.idle":"2025-12-16T17:52:03.067608Z","shell.execute_reply.started":"2025-12-16T17:52:02.899131Z","shell.execute_reply":"2025-12-16T17:52:03.066658Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Configuration class to avoid having configuration features accidentally modified\nclass CFG:\n    seed = 42\n    \n    # Input image size and batch size\n    img_size = [128, 512]\n    batch_size = 64\n    \n    # Audio duration, sample rate, and length\n    duration = 15 \n    sample_rate = 32000 #32 kHz were the downsampled dimensions\n    audio_len = duration * sample_rate #douration * sample rate (samples per second)\n    \n    # STFT parameters\n    nfft = 2028\n    window = 2048\n    hop_length = audio_len // (img_size[1] - 1)\n    fmin = 20\n    fmax = 16000\n    \n    # Number of epochs, model name\n    epochs = 10\n    preset = 'efficientnetv2_b2_imagenet'\n    \n    # Data augmentation parameters\n    augment=True\n\n    # Class Labels for BirdCLEF 24, these are actually bird species (common names)\n    class_names = sorted(os.listdir('/kaggle/input/birdclef-2024/train_audio/'))\n    num_classes = len(class_names)\n    class_labels = list(range(num_classes))\n    label2name = dict(zip(class_labels, class_names))\n    name2label = {specs:idx for idx,specs in label2name.items()}\n\nCFG = CFG()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.068437Z","iopub.execute_input":"2025-12-16T17:52:03.068682Z","iopub.status.idle":"2025-12-16T17:52:03.075244Z","shell.execute_reply.started":"2025-12-16T17:52:03.068665Z","shell.execute_reply":"2025-12-16T17:52:03.074532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BASE_PATH = '/kaggle/input/birdclef-2024'\ndf = pd.read_csv(f'{BASE_PATH}/train_metadata.csv')\ndf['filepath'] = BASE_PATH + '/train_audio/' + df.filename #Get filepath for the each record\ndf['target'] = df.primary_label.map(CFG.name2label) #number for each species\ndf['filename'] = df.filepath.map(lambda x: x.split('/')[-1])\ndf['xc_id'] = df.filepath.map(lambda x: x.split('/')[-1].split('.')[0])\n\n# Display rwos\ndf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.076887Z","iopub.execute_input":"2025-12-16T17:52:03.077105Z","iopub.status.idle":"2025-12-16T17:52:03.222899Z","shell.execute_reply.started":"2025-12-16T17:52:03.077087Z","shell.execute_reply":"2025-12-16T17:52:03.222172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Dictionary of species:\nlabels = df['primary_label'].unique()\nspecies = df['common_name'].unique()\nlabel2species = dict(zip(labels, species))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.223727Z","iopub.execute_input":"2025-12-16T17:52:03.22399Z","iopub.status.idle":"2025-12-16T17:52:03.230421Z","shell.execute_reply.started":"2025-12-16T17:52:03.223973Z","shell.execute_reply":"2025-12-16T17:52:03.229614Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Species Visualization:\nprint(label2species['pursun4'], len(label2species))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.231406Z","iopub.execute_input":"2025-12-16T17:52:03.2317Z","iopub.status.idle":"2025-12-16T17:52:03.240919Z","shell.execute_reply.started":"2025-12-16T17:52:03.231682Z","shell.execute_reply":"2025-12-16T17:52:03.24015Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Pre-process audio ","metadata":{}},{"cell_type":"code","source":"AUDIO_DIR = \"/kaggle/input/birdclef-2024/train_audio\"\nOUT_DIR = \"/kaggle/working/spec_train\"\nSPEC_DIR = '/kaggle/input/bird-call-spectrogram-10secs/spec_train'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.241665Z","iopub.execute_input":"2025-12-16T17:52:03.241823Z","iopub.status.idle":"2025-12-16T17:52:03.252226Z","shell.execute_reply.started":"2025-12-16T17:52:03.241811Z","shell.execute_reply":"2025-12-16T17:52:03.251489Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom sklearn.model_selection import train_test_split\nfrom tensorflow.keras import layers\n\n# 1. Gather all file paths and labels\nall_npy_paths = []\nall_labels = []\n\n# Walk through the output directory to get files and species labels\nfor species in os.listdir(SPEC_DIR):\n    species_dir = os.path.join(SPEC_DIR, species)\n    if os.path.isdir(species_dir):\n        \n        # Get label index from species name\n        if species in CFG.name2label:\n            label_idx = CFG.name2label[species]\n            for fname in os.listdir(species_dir):\n                if fname.endswith('.npy'):\n                    all_npy_paths.append(os.path.join(species_dir, fname))\n                    all_labels.append(label_idx)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.253035Z","iopub.execute_input":"2025-12-16T17:52:03.253267Z","iopub.status.idle":"2025-12-16T17:52:03.511209Z","shell.execute_reply.started":"2025-12-16T17:52:03.253251Z","shell.execute_reply":"2025-12-16T17:52:03.510656Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"path = \"/kaggle/input/bird-call-spectrogram-10secs/spec_train/ashdro1/XC114600.npy\"\n\nspec = np.load(path)\n\nprint(spec.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.513113Z","iopub.execute_input":"2025-12-16T17:52:03.513365Z","iopub.status.idle":"2025-12-16T17:52:03.51809Z","shell.execute_reply.started":"2025-12-16T17:52:03.51335Z","shell.execute_reply":"2025-12-16T17:52:03.51743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2. Split into Train and Validation\ntrain_paths, val_paths, train_labels, val_labels = train_test_split(\n    all_npy_paths, all_labels, test_size=0.2, random_state=CFG.seed, stratify=all_labels\n)\n\n# 3. Define the NPY loader function\ndef load_npy_data(path, label):\n    # Load numpy file\n    spec = np.load(path)\n    \n    # Add channel dimension (H, W, 1) -> ViT expects channels\n    spec = spec[..., np.newaxis] \n    \n    # Convert to tensor\n    spec = tf.convert_to_tensor(spec, dtype=tf.float32)\n    \n    # Resize to the target CFG.img_size [128, 384] (Height, Width)\n    # Note: Your generated spec is (256, Time). We resize to match CFG.\n    #spec = tf.image.resize(spec, CFG.img_size)\n    \n    # Normalize if not already done in generation (your gen code does 0-1 norm)\n    return spec, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.518765Z","iopub.execute_input":"2025-12-16T17:52:03.519041Z","iopub.status.idle":"2025-12-16T17:52:03.547102Z","shell.execute_reply.started":"2025-12-16T17:52:03.519023Z","shell.execute_reply":"2025-12-16T17:52:03.546529Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"idx = 100\nspec, label = load_npy_data(val_paths[idx], val_labels[idx])\nprint(f\"Spectrogram Shape: {spec.shape}\")\nprint(f\"Species class label/No.: {label} | codename: {CFG.label2name[label]} | Common Name: {label2species[CFG.label2name[label]]}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.547811Z","iopub.execute_input":"2025-12-16T17:52:03.54834Z","iopub.status.idle":"2025-12-16T17:52:03.553889Z","shell.execute_reply.started":"2025-12-16T17:52:03.548322Z","shell.execute_reply":"2025-12-16T17:52:03.553212Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset Creation","metadata":{}},{"cell_type":"code","source":"# Wrapper for tf.data.Dataset to use numpy load\ndef tf_load_npy(path, label):\n    [spec, label] = tf.numpy_function(load_npy_data, [path, label], [tf.float32, tf.int32])\n    spec.set_shape([CFG.img_size[0], CFG.img_size[1], 1]) \n    label.set_shape([])\n    return spec, spec\n\n# 4. Create Tensorflow Datasets\ndef create_dataset(paths, labels, shuffle=False):\n    ds = tf.data.Dataset.from_tensor_slices((paths, labels))\n    ds = ds.map(tf_load_npy, num_parallel_calls=tf.data.AUTOTUNE)\n    #ds = ds.enumerate()\n    if shuffle:\n        ds = ds.shuffle(1000)\n    ds = ds.batch(CFG.batch_size)\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n    return ds\n\ntrain_ds = create_dataset(train_paths, train_labels, shuffle=True)\nval_ds = create_dataset(val_paths, val_labels, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.554598Z","iopub.execute_input":"2025-12-16T17:52:03.554803Z","iopub.status.idle":"2025-12-16T17:52:03.651445Z","shell.execute_reply.started":"2025-12-16T17:52:03.554782Z","shell.execute_reply":"2025-12-16T17:52:03.650905Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# for idx, (x, y) in train_ds.take(1):\n#     print(idx.shape)  # (batch_size,)\n#     print(x.shape)    # (batch_size, H, W, 1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.65197Z","iopub.execute_input":"2025-12-16T17:52:03.65214Z","iopub.status.idle":"2025-12-16T17:52:03.655267Z","shell.execute_reply.started":"2025-12-16T17:52:03.652127Z","shell.execute_reply":"2025-12-16T17:52:03.654583Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Patch Creation","metadata":{}},{"cell_type":"code","source":"# ViT Hyperparameters\n# PATCH_SIZE = 16    # Size of the patches to be extracted from the input images\n# NUM_PATCHES = (CFG.img_size[0] // PATCH_SIZE) * (CFG.img_size[1] // PATCH_SIZE)\n# EMBEDDING_DIM  = [512, 256]\n# HIDDEN_UNITS = 4\n# NOISE_FACTOR = 0.5\n\nclass Patches(layers.Layer):\n    #Split into non overlapping sqaured patches.\n    def __init__(self, patch_size, **kwargs):\n        super().__init__(**kwargs)\n        self.patch_size = patch_size\n        \n    def call(self, images):\n        \"Input: patch of images\"\n        batch_size = tf.shape(images)[0]\n        # Function documentation: https://www.tensorflow.org/api_docs/python/tf/image/extract_patches\n        patches = tf.image.extract_patches(images=images,\n            sizes=[1, self.patch_size, self.patch_size, 1],\n            strides=[1, self.patch_size, self.patch_size, 1],\n            rates=[1, 1, 1, 1], #extract every pixel one by one in all directions\n            padding=\"VALID\", #Do not padd non fitting patches, just keep the ones that fit.\n        )\n        #each pathc is of size = 16 x 16 x 1\n        # patches.shape = batch size, # that fit in height, # that fit in length, patch dim (16x16)\n        \n        # Flatten\n        patch_dims = patches.shape[-1]\n        # (batch_size, num_patches_h, num_patches_w, patch_size * patch_size * channels)\n        patches = tf.reshape(patches, [batch_size, -1, patch_dims])\n        # (batch_size, num_patches_h * num_patches_w, patch_size * patch_size * channels)\n        \n        return patches\n    \n    def get_config(self):\n        config = super().get_config()\n        config.update({\"patch_size\": self.patch_size})\n        return config\n\n\nclass PatchReconstruction(layers.Layer):\n    def __init__(self, patch_size, img_height, img_width, channels=1, **kwargs):\n        super().__init__(**kwargs)\n        self.patch_size = patch_size\n        self.img_height = img_height\n        self.img_width = img_width\n        self.channels = channels\n        self.num_patches_h = img_height // patch_size\n        self.num_patches_w = img_width // patch_size\n    \n    def call(self, patches):\n        batch_size = tf.shape(patches)[0]\n        \n        # Reshape patches back to grid\n        # patches shape: (batch, num_patches_h, num_pathces_w, patch_size_h, patch_size_w, channels)\n        patches = tf.reshape(patches,\n                             [batch_size, self.num_patches_h, self.num_patches_w,\n                              self.patch_size, self.patch_size, self.channels])\n        \n        # Rearrange to form image\n        # (batch, num_patches_h, patch_size_h, num_pathces_w, patch_size_w, channels)\n        patches = tf.transpose(patches, [0, 1, 3, 2, 4, 5])\n        \n        # Merge patches into image\n        images = tf.reshape(patches, [batch_size, self.img_height, self.img_width, self.channels])\n        \n        return images\n    \n    def get_config(self):\n        config = super().get_config()\n        config.update({\n            \"patch_size\": self.patch_size,\n            \"img_height\": self.img_height,\n            \"img_width\": self.img_width,\n            \"channels\": self.channels\n        })\n        return config","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.655935Z","iopub.execute_input":"2025-12-16T17:52:03.656217Z","iopub.status.idle":"2025-12-16T17:52:03.666463Z","shell.execute_reply.started":"2025-12-16T17:52:03.656191Z","shell.execute_reply":"2025-12-16T17:52:03.66564Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PatchDecoder(layers.Layer):\n    \"\"\"Reconstruct image from patches\"\"\"\n    def __init__(self, patch_size, img_height, img_width, channels=1, **kwargs):\n        super().__init__(**kwargs)\n        self.patch_size = patch_size\n        self.img_height = img_height\n        self.img_width = img_width\n        self.channels = channels\n        self.num_patches_h = img_height // patch_size\n        self.num_patches_w = img_width // patch_size\n    \n    def call(self, patches):\n        batch_size = tf.shape(patches)[0]\n        \n        # Reshape patches back to grid\n        patches = tf.reshape(\n            patches,\n            [batch_size, self.num_patches_h, self.num_patches_w,\n             self.patch_size, self.patch_size, self.channels]\n        )\n        \n        # Rearrange to form image\n        patches = tf.transpose(patches, [0, 1, 3, 2, 4, 5])\n        \n        # Merge patches into image\n        images = tf.reshape(\n            patches,\n            [batch_size, self.img_height, self.img_width, self.channels]\n        )\n        \n        return images\n    \n    def get_config(self):\n        config = super().get_config()\n        config.update({\n            \"patch_size\": self.patch_size,\n            \"img_height\": self.img_height,\n            \"img_width\": self.img_width,\n            \"channels\": self.channels\n        })\n        return config","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.667368Z","iopub.execute_input":"2025-12-16T17:52:03.667715Z","iopub.status.idle":"2025-12-16T17:52:03.682129Z","shell.execute_reply.started":"2025-12-16T17:52:03.667693Z","shell.execute_reply":"2025-12-16T17:52:03.681478Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Patch Creation - Test","metadata":{}},{"cell_type":"code","source":"npy_path = \"/kaggle/input/bird-call-spectrogram/spec_train/asbfly/XC175797.npy\"\nimage = np.load(npy_path)\n#Shorter image\n\n# Visualize the spectrogram\nplt.figure(figsize=(12, 4))\nplt.imshow(image[:48, :48], aspect='auto', origin='lower', cmap='viridis')\nplt.colorbar(label='Amplitude')\nplt.title(\"Bird Call Spectrogram ('48x48')\")\nplt.xlabel('Time')\nplt.ylabel('Frequency')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.682872Z","iopub.execute_input":"2025-12-16T17:52:03.683218Z","iopub.status.idle":"2025-12-16T17:52:03.907997Z","shell.execute_reply.started":"2025-12-16T17:52:03.683201Z","shell.execute_reply":"2025-12-16T17:52:03.90742Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"height, width = 48, 48\npatch_size = 16\n\nimage_batch = tf.expand_dims(tf.constant(image[:height, :width], dtype=tf.float32), 0)\nimage_batch = tf.expand_dims(image_batch, axis=-1)\nprint(image_batch.shape)\npatch_layer = Patches(patch_size)\npatches = patch_layer(image_batch)\nprint(f\"Image shape: {image_batch.shape}\")  # Should be (1, 32, 32, 1)\nprint(f\"Patches shape: {patches.shape}\")  # Should be (1, 4, 256)\n\n\nfig, axes = plt.subplots(1, 9, figsize=(16, 4))\nfor patch_idx in range(9):  # You have 5 patches\n    # Get and reshape patch\n    flat_patch = patches[0, patch_idx, :]\n    patch_2d = tf.reshape(flat_patch, [patch_size, patch_size])\n    \n    # Plot\n    axes[patch_idx].imshow(patch_2d, cmap='viridis')\n    axes[patch_idx].set_title(f'Patch {patch_idx}')\n    axes[patch_idx].axis('off')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:03.908628Z","iopub.execute_input":"2025-12-16T17:52:03.908835Z","iopub.status.idle":"2025-12-16T17:52:04.35293Z","shell.execute_reply.started":"2025-12-16T17:52:03.908818Z","shell.execute_reply":"2025-12-16T17:52:04.352311Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create reconstruction layer\nrecon_layer = PatchReconstruction(patch_size=patch_size, \n                                  img_height=height, \n                                  img_width=width, \n                                  channels=1)\n\n# Reconstruct the image\nreconstructed = recon_layer(patches)\n\n# Visualize the spectrogram\nplt.figure(figsize=(12, 4))\nplt.imshow(reconstructed[0, :, :, 0], aspect='auto', origin='lower', cmap='viridis')\nplt.colorbar(label='Amplitude')\nplt.title(\"Bird Call Spectrogram ('32x32')\")\nplt.xlabel('Time')\nplt.ylabel('Frequency')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:04.353682Z","iopub.execute_input":"2025-12-16T17:52:04.353929Z","iopub.status.idle":"2025-12-16T17:52:04.574154Z","shell.execute_reply.started":"2025-12-16T17:52:04.353912Z","shell.execute_reply":"2025-12-16T17:52:04.57342Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Patch Blinding","metadata":{}},{"cell_type":"code","source":"# TRUE MAE MASKING FUNCTION (FIXED - TensorFlow compatible)\ndef create_mae_mask(batch_size, num_patches, mask_ratio):\n    \"\"\"\n    Create MAE-style mask for each item in batch\n    Creates different random masks for each item in the batch.\n    \n    Args:\n        batch_size: Number of items in batch\n        num_patches: Total number of patches\n        mask_ratio: Fraction of patches to mask\n    \n    Returns:\n        keep_indices: [batch, num_keep] - sorted indices of visible patches, by default patches indexes go row by row.\n        mask_indices: [batch, num_mask] - sorted indices of masked patches, by default patches indexes go row by row.\n    \n    Example:\n        batch_size=2, num_patches=10, mask_ratio=0.5\n        Returns:\n          keep_indices: [[0,2,4,6,8], [1,3,5,7,9]]  # 5 kept per image sample\n          mask_indices: [[1,3,5,7,9], [0,2,4,6,8]]  # 5 masked per image sample\n    \"\"\"\n    num_keep = tf.cast(tf.cast(num_patches, tf.float32) * (1.0 - mask_ratio), tf.int32) #Raw number of masks kept\n    \n    # Generate random noise for each batch item and patch: Shape: [batch_size, num_patches]\n    noise = tf.random.uniform([batch_size, num_patches], minval=0, maxval=1) # Round up allways\n    \n    # Sort indices by noise values (creates random permutation per batch item) TensorFlow-compatible\n    # argsort gives us the indices that would sort the noise:\n    # e.g. batch size = 1\n    #noise =  [[0.3, 0.9, 0.1, 0.7, 0.5, 0.2, 0.8, 0.4, 0.6, 0.0],  # Batch item 0\n            #[0.6, 0.2, 0.8, 0.1, 0.9, 0.5, 0.3, 0.7, 0.4, 0.0]]  # Batch item 1\n    #shuffled_indices = [[9, 2, 5, 0, 7, 4, 8, 3, 6, 1],  # Batch item 0\n    #                   [9, 3, 1, 6, 8, 5, 0, 2, 4]]\n\n    shuffled_indices = tf.argsort(noise, axis=1)\n\n    # Split the into keep and mask indices from the ratio% of the shuffled_indices.\n    keep_indices = tf.sort(shuffled_indices[:, :num_keep], axis=1)\n    mask_indices = tf.sort(shuffled_indices[:, num_keep:], axis=1)\n    \n    return keep_indices, mask_indices","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:04.574896Z","iopub.execute_input":"2025-12-16T17:52:04.575164Z","iopub.status.idle":"2025-12-16T17:52:04.580877Z","shell.execute_reply.started":"2025-12-16T17:52:04.575138Z","shell.execute_reply":"2025-12-16T17:52:04.58014Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Patch Blinding - Test","metadata":{}},{"cell_type":"code","source":"batch_size = patches.shape[0]\nnum_patches = patches.shape[1]\nmask_ratio = 0.75\nkeep_idxs, masked_idxs = create_mae_mask(1, 9, 0.75)\n\nprint(f\"Total patches: {num_patches}\")\nprint(f\"Visible: {keep_idxs.shape[1]}\")\nprint(f\"Masked: {masked_idxs.shape[1]}\")\n\npatches_masked = patches.numpy().copy()\n\nfor masked_idx in masked_idxs[0].numpy():\n    patches_masked[0, masked_idx, :] = 0\n\npatches_masked.shape\n\nfig, axes = plt.subplots(1, 9, figsize=(16, 4))\nfor patch_idx in range(9):  # You have 5 patches\n    # Get and reshape patch\n    flat_patch = patches_masked[0, patch_idx, :]\n    patch_2d = tf.reshape(flat_patch, [patch_size, patch_size])\n    \n    # Plot\n    axes[patch_idx].imshow(patch_2d, cmap='viridis')\n    axes[patch_idx].set_title(f'Patch {patch_idx}')\n    axes[patch_idx].axis('off')\n\nplt.tight_layout()\nplt.show()\n    \n# #https://www.geeksforgeeks.org/python/python-tensorflow-gather/\nvisible_patches = tf.gather(patches, keep_idxs, axis=1, batch_dims=1)\nmasked_patches = tf.gather(patches, masked_idxs, axis=1, batch_dims=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:04.581581Z","iopub.execute_input":"2025-12-16T17:52:04.581821Z","iopub.status.idle":"2025-12-16T17:52:05.344365Z","shell.execute_reply.started":"2025-12-16T17:52:04.581797Z","shell.execute_reply":"2025-12-16T17:52:05.343748Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig, axes = plt.subplots(1, visible_patches.shape[1], figsize=(16, 4))\nfor patch_idx in range(visible_patches.shape[1]):  # You have 5 patches\n    # Get and reshape patch\n    flat_patch = visible_patches[0, patch_idx, :]\n    patch_2d = tf.reshape(flat_patch, [patch_size, patch_size])\n    \n    # Plot\n    axes[patch_idx].imshow(patch_2d, cmap='viridis')\n    axes[patch_idx].set_title(f'Patch {patch_idx}')\n    axes[patch_idx].axis('off')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:05.345124Z","iopub.execute_input":"2025-12-16T17:52:05.345397Z","iopub.status.idle":"2025-12-16T17:52:05.50295Z","shell.execute_reply.started":"2025-12-16T17:52:05.345369Z","shell.execute_reply":"2025-12-16T17:52:05.502408Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Patch Encoder","metadata":{}},{"cell_type":"code","source":"#Inspired by: https://towardsdatascience.com/how-to-implement-state-of-the-art-masked-autoencoders-mae-6f454b736087/\nclass MAEPatchEncoder(layers.Layer):\n    \"\"\"Encode only visible patches with positional embeddings\"\"\"\n    def __init__(self, num_patches, projection_dim, **kwargs):\n        super().__init__(**kwargs)\n        self.num_patches = num_patches\n        self.projection_dim = projection_dim\n        \n    def build(self, input_shape):\n        self.projection = layers.Dense(units=self.projection_dim, name='projection')\n        self.position_embedding = layers.Embedding(\n            input_dim=self.num_patches, \n            output_dim=self.projection_dim,\n            name='position_embedding'\n        )\n        super().build(input_shape)\n    \n    def call(self, patches, keep_indices):\n        # Mimicking steps from: https://towardsdatascience.com/how-to-implement-state-of-the-art-masked-autoencoders-mae-6f454b736087/\n        \"\"\"\n        input:\n            patches: [batch, num_patches, patch_dim] - all batches\n            keep_indices: [batch, num_keep] - indices of patches to encode, output from create_mae_mask\n        \n        Returns:\n            [batch, num_keep, projection_dim]\n        \"\"\"\n        # batch_size = tf.shape(patches)[0]\n        # num_keep = tf.shape(keep_indices)[1]\n        \n        # # Gather visible patches for each batch item\n        # batch_indices = tf.repeat(tf.range(batch_size)[:, None], num_keep, axis=1)\n        # gather_indices = tf.stack([batch_indices, keep_indices], axis=-1)\n        # visible_patches = tf.gather_nd(patches, gather_indices)\n        visible_patches = tf.gather(patches, keep_indices, axis=1, batch_dims=1)\n        \n        # Project to embedding space\n        projected = self.projection(visible_patches)\n        \n        # Add positional embeddings across batch\n        position_emb = self.position_embedding(keep_indices[0])  # Use first batch's positions\n        encoded = projected + position_emb\n        \n        return encoded\n    \n    def get_config(self):\n        config = super().get_config()\n        config.update({\n            \"num_patches\": self.num_patches,\n            \"projection_dim\": self.projection_dim\n        })\n        return config","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:05.503656Z","iopub.execute_input":"2025-12-16T17:52:05.503937Z","iopub.status.idle":"2025-12-16T17:52:05.510398Z","shell.execute_reply.started":"2025-12-16T17:52:05.50392Z","shell.execute_reply":"2025-12-16T17:52:05.509692Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Patch Encoder - Test","metadata":{}},{"cell_type":"code","source":"num_batches = 1\nprojection_dim = 512\n\nencoder = MAEPatchEncoder(num_patches=num_patches, projection_dim=projection_dim)\nencoded = encoder(patches, keep_idxs)\nprint(f\"Output:\")\nprint(f\"  encoded.shape: {encoded.shape}\")\nprint(f\"  Expected shape: ({num_batches}, {keep_idxs.shape[1]}, {projection_dim})\")\n\n\nprint(f\"Projection Output:\")\nvisible_patches_manual = tf.gather(patches, keep_idxs, axis=1, batch_dims=1)\nprint(f\"  visible_patches shape: {visible_patches_manual.shape}\")\nprojected_manual = encoder.projection(visible_patches_manual)\nprint(f\"  projected shape: {projected_manual.shape}\")\n\n\nprint(f\"Positional Embedding:\")\nposition_emb = encoder.position_embedding(keep_idxs[0])\nprint(f\"  position_emb.shape: {position_emb.shape}\")\nexpected_encoded = projected_manual + position_emb\ndifference = tf.reduce_max(tf.abs(encoded - expected_encoded))\nprint(difference)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:05.511105Z","iopub.execute_input":"2025-12-16T17:52:05.511382Z","iopub.status.idle":"2025-12-16T17:52:05.54664Z","shell.execute_reply.started":"2025-12-16T17:52:05.511359Z","shell.execute_reply":"2025-12-16T17:52:05.545998Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# MAE Model","metadata":{}},{"cell_type":"markdown","source":"## Encoder","metadata":{}},{"cell_type":"code","source":"class MAEEncoder(layers.Layer):\n    \"\"\"Transformer or MLP encoder for MAE\"\"\"\n    \n    def __init__(self, config, **kwargs):\n        super().__init__(**kwargs)\n        self.config = config\n        \n        # Encoder layers\n        self.encoder_layers = []\n        if config.num_heads == 0:\n            # MLP encoder\n            print(f\"Building MLP encoder with {config.num_layers} layers.\")\n            for layer in range(config.num_layers):\n                # https://arxiv.org/html/2510.12819v1 justifies using gelu and dropout\n                self.encoder_layers.append({\n                    'dense': layers.Dense(config.hidden_dim, activation='gelu', name=f'encoder_dense_{layer}'),\n                    'norm': layers.LayerNormalization(epsilon=1e-6, name=f'encoder_norm_{layer}'),\n                    'drop': layers.Dropout(0.1, name=f'encoder_drop_{layer}')\n                })\n        else:\n            # Transformer encoder\n            print(f\"Building Transformer encoder with {config.num_layers} layers, {config.num_heads} attention heads.\")\n            for layer in range(config.num_layers):\n                # Transformer block/layer taken from:https://uvadlc-notebooks.readthedocs.io/en/latest/tutorial_notebooks/tutorial6/Transformers_and_MHAttention.html\n                self.encoder_layers.append({\n                    'norm1': layers.LayerNormalization(epsilon=1e-6, name=f'encoder_norm1_{layer}'),\n                    'attn': layers.MultiHeadAttention(num_heads=config.num_heads, \n                                                      key_dim=config.hidden_dim // config.num_heads,\n                                                      dropout=0.1, \n                                                      name=f'encoder_attn_{layer}'),\n                    'add1': layers.Add(name=f'encoder_add1_{layer}'),\n                    'norm2': layers.LayerNormalization(epsilon=1e-6, name=f'encoder_norm2_{layer}'),\n                    'mlp1': layers.Dense(config.hidden_dim * 4, activation='gelu', name=f'encoder_mlp1_{layer}'),\n                    'mlp2': layers.Dense(config.hidden_dim, name=f'encoder_mlp2_{layer}'),\n                    'drop': layers.Dropout(0.1, name=f'encoder_drop_{layer}'),\n                    'add2': layers.Add(name=f'encoder_add2_{layer}')\n                })\n        \n        self.encoder_norm = layers.LayerNormalization(epsilon=1e-6, name='encoder_output')\n    \n    def call(self, encoded, training=False):\n        \"\"\"\n        Args:\n            encoded: [batch, num_visible_patches, hidden_dim]\n            training: Boolean for dropout\n        \n        Returns:\n            [batch, num_visible_patches, hidden_dim]\n        \"\"\"\n        if self.config.num_heads == 0:\n            # MLP encoder\n            for layer in self.encoder_layers:\n                encoded = layer['dense'](encoded)\n                encoded = layer['norm'](encoded)\n                encoded = layer['drop'](encoded, training=training)\n        else:\n            # Transformer encoder (Pre-LN)\n            for layer in self.encoder_layers:\n                # Attention block\n                x1 = layer['norm1'](encoded)\n                attn = layer['attn'](x1, x1, training=training)\n                x2 = layer['add1']([attn, encoded])\n                \n                # MLP block\n                x3 = layer['norm2'](x2)\n                x3 = layer['mlp1'](x3)\n                x3 = layer['mlp2'](x3)\n                x3 = layer['drop'](x3, training=training)\n                encoded = layer['add2']([x3, x2])\n        \n        # Final normalization\n        encoded = self.encoder_norm(encoded)\n        \n        return encoded","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:05.549576Z","iopub.execute_input":"2025-12-16T17:52:05.549806Z","iopub.status.idle":"2025-12-16T17:52:05.559341Z","shell.execute_reply.started":"2025-12-16T17:52:05.549792Z","shell.execute_reply":"2025-12-16T17:52:05.558634Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Decoder","metadata":{}},{"cell_type":"markdown","source":"### Idea and Implementation","metadata":{}},{"cell_type":"code","source":"# projection_dim   \npatches2 = tf.random.normal([batch_size, num_patches, 512])\nvisible_patches2 = tf.gather(patches2, keep_idxs, axis=1, batch_dims=1)\n\nfor patch in range(num_patches):\n    patches2 = tf.tensor_scatter_nd_update(patches2, \n                                          [[0, patch, 0]], \n                                          [float(patch)])\nnum_masked_patches = masked_idxs.shape[1]\n#make a temporal embedding of feature vector size for masked embeddings.\nmask_embedding = tf.constant([[[-99.0] + [0.0]*(projection_dim-1)]], dtype=tf.float32)\n#make 1-rate copies of this placeholder\nmask_embeddings = tf.tile(mask_embedding, [batch_size, num_masked_patches, 1])\n#concatenate the patches embeddings of size (feature vector len) with the generated masked patch embedding\nfull_tokens = tf.concat([visible_patches2, mask_embeddings], axis=1)\n#concatenate the patches embeddings indexes\nfull_indices = tf.concat([keep_idxs, masked_idxs], axis=1)\n\n#sort indices and sort the embeddings.\nsorted_indices = tf.argsort(full_indices, axis=1)\nfull_embeddings_sorted = tf.gather(full_tokens, sorted_indices, axis=1, batch_dims=1)\nfull_indices, sorted_indices, full_embeddings_sorted","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:05.560813Z","iopub.execute_input":"2025-12-16T17:52:05.561068Z","iopub.status.idle":"2025-12-16T17:52:05.583995Z","shell.execute_reply.started":"2025-12-16T17:52:05.561047Z","shell.execute_reply":"2025-12-16T17:52:05.583424Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Decoder Design","metadata":{}},{"cell_type":"code","source":"class MAEDecoder(layers.Layer):\n    \"\"\"Transformer or MLP decoder for MAE\"\"\"\n    \n    def __init__(self, config, **kwargs):\n        super().__init__(**kwargs)\n        self.config = config\n        \n        # Learnable mask token - TO keep track of the added ptches that replace the missing/masked ones from the input\n        self.mask_token = None  # Will be created in build()\n        \n        # Decoder positional embeddings - now for all image pathces and not for only input patches\n        self.decoder_pos_emb = layers.Embedding(input_dim=config.num_patches, \n                                                output_dim=config.hidden_dim, \n                                                name='decoder_pos_embedding',)\n        \n        # Decoder layers (shallower than encoder)\n        num_decoder_layers = max(1, config.num_layers // 2)\n        self.decoder_layers = []\n        \n        if config.num_heads == 0:\n            # MLP decoder\n            print(f\"Building MLP decoder with {num_decoder_layers} layers.\")\n            for layer in range(num_decoder_layers):\n                self.decoder_layers.append({\n                    'dense': layers.Dense(config.hidden_dim, activation='gelu', name=f'decoder_dense_{layer}'),\n                    'norm': layers.LayerNormalization(epsilon=1e-6, name=f'decoder_norm_{layer}'),\n                    'drop': layers.Dropout(0.1, name=f'decoder_drop_{layer}')\n                })\n        else:\n            # Transformer decoder\n            print(f\"Building Transformer decoder with {num_decoder_layers} layers.\")\n            for layer in range(num_decoder_layers):\n                self.decoder_layers.append({\n                    'norm1': layers.LayerNormalization(epsilon=1e-6, name=f'decoder_norm1_{layer}'),\n                    'attn': layers.MultiHeadAttention(num_heads=config.num_heads, \n                                                      key_dim=config.hidden_dim // config.num_heads, \n                                                      dropout=0.1, name=f'decoder_attn_{layer}'),\n                    'add1': layers.Add(name=f'decoder_add1_{layer}'),\n                    'norm2': layers.LayerNormalization(epsilon=1e-6, name=f'decoder_norm2_{layer}'),\n                    'mlp1': layers.Dense(config.hidden_dim * 4, activation='gelu', name=f'decoder_mlp1_{layer}'),\n                    'mlp2': layers.Dense(config.hidden_dim, name=f'decoder_mlp2_{layer}'),\n                    'drop': layers.Dropout(0.1, name=f'decoder_drop_{layer}'),\n                    'add2': layers.Add(name=f'decoder_add2_{layer}')\n                })\n        \n        # Patch reconstruction projection, converts abstract decoder features back into actual pixel values.\n        # convert from feature size to flattened patch size.\n        patch_dim = config.patch_size * config.patch_size * 1\n        self.patch_projection = layers.Dense(patch_dim, \n                                             activation='sigmoid',\n                                             name='patch_reconstruction',)\n    \n    def build(self, input_shape):\n        # Create learnable mask embedding to represent \n        self.mask_token = self.add_weight(\n            shape=(1, 1, self.config.hidden_dim),\n            initializer='glorot_uniform',\n            trainable=True,\n            name='mask_token'\n        )\n        super().build(input_shape)\n    \n    def call(self, encoded, keep_indices, mask_indices, training=False):\n        \"\"\"\n        Args:\n            encoded: [batch, num_visible, hidden_dim] - encoded visible patches\n            keep_indices: [batch, num_visible] - indices of visible patches\n            mask_indices: [batch, num_masked] - indices of masked patches\n            training: Boolean for dropout\n        \n        Returns:\n            [batch, num_patches, patch_dim] - reconstructed patches\n        \"\"\"\n        batch_size = tf.shape(encoded)[0]\n        num_mask = tf.shape(mask_indices)[1]\n\n        #make 1-rate copies of initialized embedding\n        \n        # Create mask tokens for masked positions along each batche's blinded patches.\n        mask_embeddings = tf.tile(self.mask_token, [batch_size, num_mask, 1])\n\n        #concatenate the patches embeddings of size (feature vector len) with the generated masked patch embedding\n        full_tokens = tf.concat([encoded, mask_embeddings], axis=1)\n        #concatenate the patches embeddings indexes\n        full_indices = tf.concat([keep_indices, mask_indices], axis=1)\n        \n        #sort indices and sort the embeddings.\n        sorted_indices = tf.argsort(full_indices, axis=1)\n        full_tokens = tf.gather(full_tokens, sorted_indices, axis=1, batch_dims=1)\n        \n        #add decoder positional embeddings\n        positions = tf.range(self.config.num_patches)\n        pos_emb = self.decoder_pos_emb(positions)\n        decoded = full_tokens + pos_emb\n        \n        #Process through decoder layers\n        if self.config.num_heads == 0:\n            # MLP decoder\n            for layer in self.decoder_layers:\n                decoded = layer['dense'](decoded)\n                decoded = layer['norm'](decoded)\n                decoded = layer['drop'](decoded, training=training)\n        else:\n            # Transformer decoder (Pre-LN)\n            for layer in self.decoder_layers:\n                # Attention block\n                x1 = layer['norm1'](decoded)\n                attn = layer['attn'](x1, x1, training=training)\n                x2 = layer['add1']([attn, decoded])\n                \n                # MLP block\n                x3 = layer['norm2'](x2)\n                x3 = layer['mlp1'](x3)\n                x3 = layer['mlp2'](x3)\n                x3 = layer['drop'](x3, training=training)\n                decoded = layer['add2']([x3, x2])\n        \n        #Project to patch dimensions between 0-1\n        reconstructed_patches = self.patch_projection(decoded)\n        \n        return reconstructed_patches","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:05.584658Z","iopub.execute_input":"2025-12-16T17:52:05.584854Z","iopub.status.idle":"2025-12-16T17:52:05.598658Z","shell.execute_reply.started":"2025-12-16T17:52:05.584822Z","shell.execute_reply":"2025-12-16T17:52:05.598143Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class MAEModel(keras.Model):\n    def __init__(self, config, **kwargs):\n        super().__init__(**kwargs)\n        self.config = config\n        \n        # Patchify layer\n        self.patchify = Patches(config.patch_size, name='patchify')\n        self.patch_encoder = MAEPatchEncoder(config.num_patches, \n                                             config.hidden_dim, \n                                             name='patch_encoder')\n        \n        # Encoder\n        self.encoder = MAEEncoder(config, name='mae_encoder')\n        # Decoder\n        self.decoder = MAEDecoder(config, name='mae_decoder')\n        \n        # Image reconstruction\n        self.reconstruct = PatchDecoder(\n            patch_size=config.patch_size,\n            img_height=config.img_size[0],\n            img_width=config.img_size[1],\n            channels=1,\n            name='image_reconstruction'\n        )\n    def call(self, inputs, training=False):\n        \"\"\"\n        Args:\n            inputs: [batch, height, width, 1] - input spectrograms\n            training: Boolean for dropout\n        \n        Returns:\n            [batch, height, width, 1] - reconstructed spectrograms\n        \"\"\"\n        \n        batch_size = tf.shape(inputs)[0]\n        \n        # Patchify\n        patches = self.patchify(inputs)\n        \n        # Create random mask\n        keep_indices, mask_indices = create_mae_mask(batch_size, \n                                                     self.config.num_patches, \n                                                     self.config.mask_ratio)\n        \n        # Encode only visible patches\n        encoded = self.patch_encoder(patches, keep_indices)\n        encoded = self.encoder(encoded, training=training)\n\n        # Decode (reconstructs all patches)\n        reconstructed_patches = self.decoder(encoded, \n                                             keep_indices, \n                                             mask_indices, \n                                             training=training)\n        # Reconstruct image\n        reconstructed_image = self.reconstruct(reconstructed_patches)\n\n        return reconstructed_image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:05.599378Z","iopub.execute_input":"2025-12-16T17:52:05.599639Z","iopub.status.idle":"2025-12-16T17:52:05.614227Z","shell.execute_reply.started":"2025-12-16T17:52:05.599622Z","shell.execute_reply":"2025-12-16T17:52:05.613559Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Loop","metadata":{}},{"cell_type":"code","source":"# TRAINING FUNCTION\ndef train_true_mae_experiment(config, train_ds, val_ds):\n    \"\"\"Train True MAE experiment\"\"\"\n    # Datasets review\n    print(f\"\\nTrain batches: {len(train_ds)}\")\n    print(f\"Val batches: {len(val_ds)}\")\n\n    os.makedirs('models', exist_ok=True)\n    os.makedirs('logs', exist_ok=True)\n    \n    # Build model\n    print(\"\\nBuilding model...\")\n    model = MAEModel(config)\n\n    #Cosine Decay\n    steps_per_epoch = len(train_ds)\n    total_steps = config.epochs * steps_per_epoch\n    \n    lr_schedule = CosineDecay(\n        initial_learning_rate=config.learning_rate,\n        decay_steps=total_steps,\n        alpha=0.0)\n\n    print(f\"  Initial LR: {config.learning_rate}\")\n    print(f\"  Total steps: {total_steps}\")\n    print(f\"  Steps per epoch: {steps_per_epoch}\")\n\n    optimizer = keras.optimizers.Adam(learning_rate=lr_schedule)\n    \n    # Compile\n    model.compile(\n        optimizer=optimizer,\n        loss='mse',\n        metrics=['mae']\n    )\n    \n    # Callbacks\n    callbacks = [\n        # keras.callbacks.ReduceLROnPlateau(\n        #     monitor='val_loss', factor=0.5, patience=2, min_lr=1e-6, verbose=1\n        # ),\n        keras.callbacks.EarlyStopping(monitor='val_loss', \n                                      patience=5, \n                                      restore_best_weights=True, \n                                      verbose=1),\n        keras.callbacks.ModelCheckpoint(f'models/true_mae_{config.name}_best.keras', \n                                        monitor='val_loss', save_best_only=True, verbose=1\n        ),\n        keras.callbacks.CSVLogger(f'logs/true_mae_{config.name}_training.csv'),\n        #keras.callbacks.LambdaCallback(on_epoch_end=lambda epoch, \n         #                              logs: logs.update({'lr': model.optimizer.learning_rate(model.optimizer.iterations).numpy()})\n        #LRLogger()\n    ]\n    \n    # Train\n    print(\"\\nStarting training...\")\n    history = model.fit(train_ds, \n                        validation_data=val_ds, \n                        epochs=config.epochs, \n                        callbacks=callbacks, \n                        verbose=1)\n    \n    # Save final model\n    model.save(f'models/true_mae_{config.name}_final.keras')\n    \n    return model, history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:05.614979Z","iopub.execute_input":"2025-12-16T17:52:05.615282Z","iopub.status.idle":"2025-12-16T17:52:05.628733Z","shell.execute_reply.started":"2025-12-16T17:52:05.615261Z","shell.execute_reply":"2025-12-16T17:52:05.628214Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Experiment Configuration","metadata":{}},{"cell_type":"code","source":"from dataclasses import dataclass\n\n@dataclass\nclass ExperimentConfig:\n    \"\"\"Configuration for MAE experiments\"\"\"\n    name: str\n    hidden_dim: int = 256\n    num_layers: int = 2\n    num_heads: int = 4\n    batch_size: int = 16\n    mask_ratio: float = 0.75\n    learning_rate: float = 0.001\n    epochs: int = 11\n    seed: int = 42\n    img_size: list = None\n    patch_size: int = 16\n    num_patches: int = None  # Will be auto-calculated\n    \n    def __post_init__(self):\n        \"\"\"Auto-calculate num_patches from img_size and patch_size\"\"\"\n        # Use defaults if not provided\n        if self.img_size is None:\n            #self.img_size = [256, 944]\n            self.img_size = [128, 512]\n        \n        # Calculate num_patches\n        h_patches = self.img_size[0] // self.patch_size\n        w_patches = self.img_size[1] // self.patch_size\n        self.num_patches = h_patches * w_patches\n        \n        print(f\"Config '{self.name}':\")\n        print(f\"  Image size: {self.img_size}\")\n        print(f\"  Patch size: {self.patch_size}\")\n        print(f\"  Patches: {h_patches} x {w_patches} = {self.num_patches}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:05.629449Z","iopub.execute_input":"2025-12-16T17:52:05.629717Z","iopub.status.idle":"2025-12-16T17:52:05.642986Z","shell.execute_reply.started":"2025-12-16T17:52:05.629702Z","shell.execute_reply":"2025-12-16T17:52:05.642269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EXPERIMENTS = [\n    # Phase 1: MLP Baseline (no attention heads)\n    ExperimentConfig(\"mae_mlp_mask50_10secs\", hidden_dim=256, num_layers=2, num_heads=0, \n                     batch_size=64, mask_ratio=0.50),\n    ExperimentConfig(\"mae_mlp_mask75_10secs\", hidden_dim=256, num_layers=2, num_heads=0, \n                     batch_size=64, mask_ratio=0.75),\n    \n    # Phase 2: Small Transformer\n    ExperimentConfig(\"mae_small_mask50_4_heads_10secs\", hidden_dim=256, num_layers=2, num_heads=4, \n                     batch_size=64, mask_ratio=0.50),\n    ExperimentConfig(\"mae_small_mask75_4_heads_10secs\", hidden_dim=256, num_layers=2, num_heads=4, \n                     batch_size=64, mask_ratio=0.75),\n    \n    # Phase 3: Medium Transformer\n    ExperimentConfig(\"mae_medium_mask50_6_heads_10secs\", hidden_dim=384, num_layers=4, num_heads=6, \n                     batch_size=32, mask_ratio=0.50),\n    ExperimentConfig(\"mae_medium_mask75_6_heads_10secs\", hidden_dim=384, num_layers=4, num_heads=6, \n                     batch_size=32, mask_ratio=0.75),\n    \n    # Phase 4: Large Transformer\n    ExperimentConfig(\"mae_large_mask50_8_heads_10secs\", hidden_dim=512, num_layers=6, num_heads=8, \n                     batch_size=16, mask_ratio=0.50),\n    ExperimentConfig(\"mae_large_mask75_8_heads_10secs\", hidden_dim=512, num_layers=6, num_heads=8, \n                     batch_size=16, mask_ratio=0.75),\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:05.643693Z","iopub.execute_input":"2025-12-16T17:52:05.643875Z","iopub.status.idle":"2025-12-16T17:52:05.65585Z","shell.execute_reply.started":"2025-12-16T17:52:05.643862Z","shell.execute_reply":"2025-12-16T17:52:05.655167Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"K.clear_session()\n# model, history = train_true_mae_experiment(EXPERIMENTS[0], train_ds, val_ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:05.656711Z","iopub.execute_input":"2025-12-16T17:52:05.656944Z","iopub.status.idle":"2025-12-16T17:52:25.099584Z","shell.execute_reply.started":"2025-12-16T17:52:05.656918Z","shell.execute_reply":"2025-12-16T17:52:25.098065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model, history = train_true_mae_experiment(EXPERIMENTS[1], train_ds, val_ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:25.100271Z","iopub.status.idle":"2025-12-16T17:52:25.100595Z","shell.execute_reply.started":"2025-12-16T17:52:25.100427Z","shell.execute_reply":"2025-12-16T17:52:25.100442Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model, history = train_true_mae_experiment(EXPERIMENTS[2], train_ds, val_ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:25.102013Z","iopub.status.idle":"2025-12-16T17:52:25.102372Z","shell.execute_reply.started":"2025-12-16T17:52:25.102194Z","shell.execute_reply":"2025-12-16T17:52:25.102213Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model, history = train_true_mae_experiment(EXPERIMENTS[3], train_ds, val_ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:25.103501Z","iopub.status.idle":"2025-12-16T17:52:25.103819Z","shell.execute_reply.started":"2025-12-16T17:52:25.10366Z","shell.execute_reply":"2025-12-16T17:52:25.103674Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model, history = train_true_mae_experiment(EXPERIMENTS[4], train_ds, val_ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:25.105127Z","iopub.status.idle":"2025-12-16T17:52:25.105445Z","shell.execute_reply.started":"2025-12-16T17:52:25.105295Z","shell.execute_reply":"2025-12-16T17:52:25.105309Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model, history = train_true_mae_experiment(EXPERIMENTS[5], train_ds, val_ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:52:25.106603Z","iopub.status.idle":"2025-12-16T17:52:25.10688Z","shell.execute_reply.started":"2025-12-16T17:52:25.10673Z","shell.execute_reply":"2025-12-16T17:52:25.106742Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model, history = train_true_mae_experiment(EXPERIMENTS[6], train_ds, val_ds)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model, history = train_true_mae_experiment(EXPERIMENTS[7], train_ds, val_ds)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}