{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"},{"sourceId":11569755,"sourceType":"datasetVersion","datasetId":7253661},{"sourceId":12038896,"sourceType":"datasetVersion","datasetId":7377931},{"sourceId":244283798,"sourceType":"kernelVersion"}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Train dataset","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport os\n\n# Define the path to the specific .npy file\n# Ensure this path is correct for your Kaggle environment where the dataset is mounted.\nfile_path = '/kaggle/input/open-wfi-test/test/000039dca2.npy'\n\ntry:\n    # Load the .npy file\n    data = np.load(file_path)\n    \n    # Print the shape of the loaded data\n    print(f\"Shape of the sample data from {os.path.basename(file_path)}: {data.shape}\")\n    \n    # Determine the image slice to plot\n    img_to_plot = None\n    if data.ndim == 3:\n        # If the data is (channels, height, width) or (time_steps, depth, width),\n        # take the first channel/slice for plotting.\n        # This assumes the first dimension is features/channels/time-steps.\n        img_to_plot = data[0, :, :]\n        print(f\"Plotting the first slice (index 0) from the 3D data. Slice shape: {img_to_plot.shape}\")\n    elif data.ndim == 2:\n        # If the data is already 2-dimensional (height, width), plot it directly.\n        img_to_plot = data\n        print(f\"Data is 2-dimensional. Plotting directly. Shape: {img_to_plot.shape}\")\n    else:\n        print(f\"Cannot plot data with {data.ndim} dimensions. Expected 2 or 3 dimensions for image visualization.\")\n\n    # Plot the image if a valid slice was extracted\n    if img_to_plot is not None:\n        plt.figure(figsize=(10, 6))\n        # Use 'gray' colormap for intensity data, 'aspect='auto' to prevent stretching if dimensions are very different\n        plt.imshow(img_to_plot, cmap='gray', aspect='auto') \n        plt.title(f\"Sample from {os.path.basename(file_path)}\\n(Slice 0, Shape: {img_to_plot.shape})\")\n        plt.colorbar(label='Amplitude')\n        plt.xlabel('X (Width)')\n        plt.ylabel('Y (Time/Depth)')\n        plt.tight_layout()\n        plt.show()\n    else:\n        print(\"No image could be generated for plotting.\")\n\nexcept FileNotFoundError:\n    print(f\"Error: The file '{file_path}' was not found. Please ensure the path is correct and the dataset is mounted.\")\nexcept Exception as e:\n    print(f\"An unexpected error occurred: {e}\")","metadata":{"trusted":true,"jupyter":{"source_hidden":true,"outputs_hidden":true},"execution":{"iopub.status.busy":"2025-06-08T06:59:40.239116Z","iopub.execute_input":"2025-06-08T06:59:40.239416Z","iopub.status.idle":"2025-06-08T06:59:40.765528Z","shell.execute_reply.started":"2025-06-08T06:59:40.239387Z","shell.execute_reply":"2025-06-08T06:59:40.764755Z"},"collapsed":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the path to the specific .npy file\n# Ensure this path is correct for your Kaggle environment where the dataset is mounted.\nfile_path = '/kaggle/input/waveform-inversion/test/0000fd8ec8.npy'\n\ntry:\n    # Load the .npy file\n    data = np.load(file_path)\n    \n    # Print the shape of the loaded data\n    print(f\"Shape of the sample data from {os.path.basename(file_path)}: {data.shape}\")\n    \n    # Determine the image slice to plot\n    img_to_plot = None\n    if data.ndim == 3:\n        # If the data is (channels, height, width) or (time_steps, depth, width),\n        # take the first channel/slice for plotting.\n        # This assumes the first dimension is features/channels/time-steps.\n        img_to_plot = data[0, :, :]\n        print(f\"Plotting the first slice (index 0) from the 3D data. Slice shape: {img_to_plot.shape}\")\n    elif data.ndim == 2:\n        # If the data is already 2-dimensional (height, width), plot it directly.\n        img_to_plot = data\n        print(f\"Data is 2-dimensional. Plotting directly. Shape: {img_to_plot.shape}\")\n    else:\n        print(f\"Cannot plot data with {data.ndim} dimensions. Expected 2 or 3 dimensions for image visualization.\")\n\n    # Plot the image if a valid slice was extracted\n    if img_to_plot is not None:\n        plt.figure(figsize=(10, 6))\n        # Use 'gray' colormap for intensity data, 'aspect='auto' to prevent stretching if dimensions are very different\n        plt.imshow(img_to_plot, cmap='gray', aspect='auto') \n        plt.title(f\"Sample from {os.path.basename(file_path)}\\n(Slice 0, Shape: {img_to_plot.shape})\")\n        plt.colorbar(label='Amplitude')\n        plt.xlabel('X (Width)')\n        plt.ylabel('Y (Time/Depth)')\n        plt.tight_layout()\n        plt.show()\n    else:\n        print(\"No image could be generated for plotting.\")\n\nexcept FileNotFoundError:\n    print(f\"Error: The file '{file_path}' was not found. Please ensure the path is correct and the dataset is mounted.\")\nexcept Exception as e:\n    print(f\"An unexpected error occurred: {e}\")","metadata":{"trusted":true,"jupyter":{"source_hidden":true,"outputs_hidden":true},"execution":{"iopub.status.busy":"2025-06-08T06:59:40.766423Z","iopub.execute_input":"2025-06-08T06:59:40.766665Z","iopub.status.idle":"2025-06-08T06:59:41.177194Z","shell.execute_reply.started":"2025-06-08T06:59:40.766645Z","shell.execute_reply":"2025-06-08T06:59:41.176395Z"},"collapsed":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport tensorflow as tf\nfrom tqdm import tqdm # Used for progress bar during initial metadata loading\nimport glob # For finding test files\nimport csv # For writing submission.csv\nimport time # For timing inference\nimport matplotlib.pyplot as plt # For plotting predictions\n\n# --- Weights & Biases (WandB) Imports and Setup ---\nimport wandb\nfrom kaggle_secrets import UserSecretsClient # Required for Kaggle environments to access secrets\nfrom wandb.integration.keras import WandbCallback # Recommended for Keras integration with WandB\n\n# --- Global Flags to Control Script Execution ---\n# Set these to True or False to control which parts of the script run.\nRUN_TRAIN = True  # Set to True to initialize and train the model on the training dataset.\n# RUN_VALID flag removed as validation is now integrated into training.\nRUN_TEST  = True  # Set to True to initialize the test dataset, run inference, and plot predictions.\n\n# --- Experimental Mode ---\n# If True, a limited number of samples will be used for training, validation, and prediction\n# to quickly check if the script is running correctly without processing the entire dataset.\nEXPERIMENTAL_MODE = True\nEXPERIMENTAL_SUBSAMPLE_LIMIT = 100 # Number of records/files to use in experimental mode\n\n# --- Load Pre-trained Weights Option ---\n# Set this path if you want to load a previously saved model's weights.\n# Example: LOAD_WEIGHTS_PATH = \"/kaggle/input/seismic_velocity_model.keras\"\n# If set to None, the model will initialize with random weights.\n# LOAD_WEIGHTS_PATH = None\nLOAD_WEIGHTS_PATH = \"/kaggle/input/seismic-velocity-model-checkpoint/seismic_velocity_model.keras\" # Example: uncomment to load weights\n\n\n# --- Configuration Class ---\n# A simple configuration class to hold dataset parameters.\n# This mimics the 'cfg' object used in your original PyTorch dataset.\nclass Cfg:\n    def __init__(self, data_dir, subsample=None, local_rank=0, samples_per_record=500, batch_size_val=4):\n        \"\"\"\n        Initializes the configuration for the dataset.\n\n        Args:\n            data_dir (str): Path to the directory containing the actual .npy data and label files.\n                            This should be the base directory where data_fpath/label_fpath in folds.csv are relative to.\n            subsample (int, optional): If specified, limits the number of records (files) in the dataset.\n                                       Defaults to None (no subsampling).\n            local_rank (int): Used to control tqdm's progress bar visibility (only show for rank 0).\n            samples_per_record (int): The number of individual samples contained within each .npy file.\n                                      Derived from your original __len__ and __getitem__ logic (500).\n            batch_size_val (int): Batch size for validation/test datasets.\n        \"\"\"\n        self.data_dir = data_dir\n        self.subsample = subsample\n        self.local_rank = local_rank\n        self.samples_per_record = samples_per_record\n        self.batch_size_val = batch_size_val\n        # Define a dummy device for compatibility, TF handles device placement\n        self.device = tf.config.list_physical_devices('GPU')[0] if tf.config.list_physical_devices('GPU') else tf.config.list_physical_devices('CPU')[0]\n\n\n# --- TensorFlow Training/Evaluation Dataset Adapter Class ---\nclass CustomTFDataset:\n    def __init__(self, cfg, mode=\"train\"):\n        \"\"\"\n        Initializes the custom TensorFlow dataset adapter.\n\n        Args:\n            cfg (Cfg): Configuration object containing data_dir, subsample, local_rank, samples_per_record.\n            mode (str): \"train\" or \"eval\". Determines data split and augmentation logic.\n        \"\"\"\n        self.cfg = cfg\n        self.mode = mode\n        # Load file paths and other metadata (but not the actual numpy arrays)\n        self.metadata_list = self._load_metadata()\n        # Calculate the total number of samples across all selected files\n        self.total_samples = len(self.metadata_list) * self.cfg.samples_per_record\n\n        # Crucially, load one dummy sample to infer the expected shapes and data types\n        # of the output tensors. This information is required by tf.py_function.\n        print(f\"Loading a dummy sample to determine data types and shapes for {self.mode} dataset...\")\n        # tf.constant(0, dtype=tf.int64) ensures the input type matches what tf.data.Dataset.range yields\n        # Ensure that total_samples is at least 1 before trying to load a dummy sample\n        if self.total_samples == 0:\n            raise ValueError(f\"No samples found in the {self.mode} dataset based on the provided configuration and folds.csv. \"\n                             \"Please check your data_dir, folds.csv, and mode settings.\")\n        \n        dummy_x, dummy_y = self._load_and_process_single_sample(tf.constant(0, dtype=tf.int64))\n        self.x_dtype = dummy_x.dtype\n        self.y_dtype = dummy_y.dtype\n        self.x_shape = dummy_x.shape\n        self.y_shape = dummy_y.shape\n        print(f\"Determined X shape: {self.x_shape}, X dtype: {self.x_dtype}\")\n        print(f\"Determined Y shape: {self.y_shape}, Y dtype: {self.y_dtype}\")\n\n    def _load_metadata(self):\n        \"\"\"\n        Loads the paths to data and label files from the folds.csv.\n        This function performs the initial file system scan and filtering based on mode and subsample,\n        but it does NOT load the actual numpy array data into memory.\n        \"\"\"\n        # Explicitly set the full path to folds.csv as per your data organization\n        folds_csv_path = \"/kaggle/input/openfwi-preprocessed-72x72/folds.csv\"\n        if not os.path.exists(folds_csv_path):\n            raise FileNotFoundError(f\"Error: folds.csv not found at {folds_csv_path}. \"\n                                    f\"Please ensure the path is correct and the file exists.\")\n\n        df = pd.read_csv(folds_csv_path)\n\n        # Filter dataframe rows based on the dataset mode (\"train\" or \"eval\")\n        if self.mode == \"train\":\n            df = df[df[\"fold\"] != 0] # Use folds other than 0 for training\n        else:\n            df = df[df[\"fold\"] == 0] # Use fold 0 for evaluation/testing\n\n        # Apply subsampling if specified in the config (limits total files after mode filter)\n        if self.cfg.subsample is not None:\n            df = df.head(self.cfg.subsample) # Take the first N files directly\n\n        metadata_list = []\n        # Use tqdm for a progress bar during initial metadata loading.\n        # Disable tqdm if local_rank is not 0 (e.g., in a distributed training setup).\n        for _, row in tqdm(df.iterrows(), total=len(df), disable=(self.cfg.local_rank != 0),\n                           desc=f\"Loading {self.mode} metadata\"):\n            row_dict = row.to_dict()\n            metadata_list.append({\n                # Construct full paths to data and label files using the base data_dir\n                \"data_fpath\": os.path.join(self.cfg.data_dir, row_dict[\"data_fpath\"]),\n                \"label_fpath\": os.path.join(self.cfg.data_dir, row_dict[\"label_fpath\"]),\n                \"dataset_name\": row_dict[\"dataset\"]\n            })\n        return metadata_list\n\n    def _load_and_process_single_sample(self, global_idx_tensor):\n        \"\"\"\n        Loads and processes a single sample given its global index.\n        This function is intended to be wrapped by tf.py_function.\n        It performs actual numpy file loading, slicing, and augmentation.\n\n        Args:\n            global_idx_tensor (tf.Tensor): The global index of the sample, as a TensorFlow tensor.\n\n        Returns:\n            tuple: A tuple containing the processed data (x) and label (y) as tf.Tensor.\n        \"\"\"\n        # Convert the TensorFlow tensor to a numpy integer for Python indexing\n        global_idx = global_idx_tensor.numpy().item() # .item() extracts scalar from 0-dim array\n\n        # Calculate the row (file) and column (sample within file) index\n        row_idx = global_idx // self.cfg.samples_per_record\n        col_idx = global_idx % self.cfg.samples_per_record\n\n        # Retrieve file paths from the pre-loaded metadata list\n        record_metadata = self.metadata_list[row_idx]\n        data_fpath = record_metadata[\"data_fpath\"]\n        label_fpath = record_metadata[\"label_fpath\"]\n\n        # Determine the memory map mode based on the dataset mode.\n        # 'r' is read-only memory map; None means load fully into memory (often for smaller files).\n        mmap_mode = \"r\" if self.mode == \"train\" else None\n        \n        try:\n            # Load the full numpy arrays. Using mmap_mode efficiently handles large files\n            # by mapping them directly into memory without reading the entire content.\n            arr = np.load(data_fpath, mmap_mode=mmap_mode)\n            lbl = np.load(label_fpath, mmap_mode=mmap_mode)\n        except Exception as e:\n            # Handle potential errors during file loading (e.g., file not found, corruption)\n            print(f\"Error loading numpy file: {e}. Data file: {data_fpath}, Label file: {label_fpath}\")\n            # In a real scenario, you might log this, skip the sample, or return dummy data.\n            # Raising the exception here will stop the dataset pipeline if a file is missing/corrupt.\n            raise e\n\n        # Extract the specific sample using the column index\n        # '...' ensures that all other dimensions are included\n        x = arr[col_idx, ...]\n        y = lbl[col_idx, ...]\n\n        # Apply augmentations only when in training mode\n        if self.mode == \"train\":\n            # Temporal flip augmentation:\n            # Flips the first dimension (time) and the last spatial dimension (width) for x.\n            # Flips the last spatial dimension (width) for y.\n            if np.random.random() < 0.5:\n                x = x[::-1, :, ::-1] # Example: if x is (Channels, Time, Width) or (Time, Height, Width)\n                y = y[..., ::-1]     # Example: if y is (Height, Width) or (Height, Width, Channels)\n\n        # It's crucial to make a copy of the numpy arrays.\n        # Memory-mapped arrays are views, and if not copied, subsequent TensorFlow operations\n        # or Python garbage collection might cause issues.\n        x = x.copy()\n        y = y.copy()\n        \n        # Convert the numpy arrays to TensorFlow Tensors and ensure they have the desired dtype.\n        # Assuming data and labels are floating point. If they are integers, adjust dtype accordingly.\n        return tf.convert_to_tensor(x, dtype=tf.float32), tf.convert_to_tensor(y, dtype=tf.float32)\n\n    def create_tf_dataset(self):\n        \"\"\"\n        Creates and returns a tf.data.Dataset object for the custom dataset.\n\n        Returns:\n            tf.data.Dataset: A TensorFlow dataset ready for training or evaluation.\n        \"\"\"\n        # Create a dataset that yields sequential global indices from 0 to total_samples - 1\n        indices_ds = tf.data.Dataset.range(self.total_samples)\n\n        # Map the Python processing function (_load_and_process_single_sample) to each index.\n        # tf.py_function allows arbitrary Python code to run within the TF graph.\n        # Tout specifies the output types of the Python function (inferred during initialization).\n        dataset = indices_ds.map(\n            lambda idx: tf.py_function(\n                self._load_and_process_single_sample, # The Python function to execute\n                inp=[idx],                            # Input tensors to the Python function\n                Tout=[self.x_dtype, self.y_dtype],     # Expected output types (for compatibility)\n                name=\"load_and_process_sample_train\"\n            ),\n            num_parallel_calls=tf.data.AUTOTUNE # Automatically determines optimal number of parallel calls\n        )\n\n        # Re-add tf.ensure_shape after tf.py_function as experimental_output_signature is not supported\n        dataset = dataset.map(\n            lambda x, y: (tf.ensure_shape(x, self.x_shape), tf.ensure_shape(y, self.y_shape)),\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n        \n        # Shuffle the dataset only for training mode to ensure randomness in batches\n        if self.mode == \"train\":\n            # The buffer size determines how many elements from the dataset are buffered\n            # for shuffling. A larger buffer leads to better shuffling but uses more memory.\n            # Common heuristic: a few thousand or min(self.total_samples, 10000).\n            shuffle_buffer_size = min(self.total_samples, 10000) \n            dataset = dataset.shuffle(buffer_size=shuffle_buffer_size, reshuffle_each_iteration=True)\n        \n        # As explicitly requested: \"the dataset is huge so no cache is needed\".\n        # The .cache() operation is intentionally omitted.\n\n        # Prefetch data to overlap data preprocessing with model execution (GPU/TPU training).\n        # This significantly improves pipeline throughput.\n        dataset = dataset.prefetch(buffer_size=tf.data.AUTOTUNE)\n\n        return dataset\n\n# --- TensorFlow Test (Prediction) Dataset Adapter Class ---\nclass CustomTFTestDataset:\n    def __init__(self, cfg, test_files):\n        \"\"\"\n        Initializes the custom TensorFlow test dataset adapter for prediction.\n\n        Args:\n            cfg (Cfg): Configuration object.\n            test_files (list): List of paths to individual test .npy files.\n        \"\"\"\n        self.cfg = cfg\n        self.test_files = test_files\n        self.total_samples = len(self.test_files)\n\n        print(f\"Loading a dummy sample to determine data types and shapes for test dataset...\")\n        if self.total_samples == 0:\n            raise ValueError(\"No test files found. Please check `test_files` path.\")\n        \n        # Load one dummy sample to infer the expected shapes and data types\n        # Input to _load_single_test_sample is a tf.string tensor (file path)\n        dummy_x, dummy_oid = self._load_single_test_sample(tf.constant(self.test_files[0], dtype=tf.string))\n        self.x_dtype = dummy_x.dtype\n        # IMPORTANT: The output shape of _preprocess_test_sample is (channels, 72, 72)\n        # So, self.x_shape must reflect this.\n        self.x_shape = dummy_x.shape \n        # OID is a string (stem), its shape will be scalar\n        self.oid_dtype = dummy_oid.dtype \n        self.oid_shape = tf.TensorShape([]) \n        print(f\"Determined X shape (after preprocessing in dataset): {self.x_shape}, X dtype: {self.x_dtype}\")\n        print(f\"Determined OID shape: {self.oid_shape}, OID dtype: {self.oid_dtype}\")\n\n\n    def _preprocess_test_sample(self, x):\n        \"\"\"\n        Applies preprocessing similar to your PyTorch _preprocess function,\n        but using TensorFlow operations.\n        Input x is expected to be (channels, height, width).\n        \"\"\"\n        # Permute to (height, width, channels) for tf.image.resize\n        x_permuted = tf.transpose(x, perm=[1, 2, 0]) # (H, W, C)\n\n        # Interpolate to (70, 70) using AREA method\n        # tf.image.resize expects (height, width)\n        x_interpolated = tf.image.resize(\n            x_permuted, size=(70, 70), method=tf.image.ResizeMethod.AREA\n        )\n\n        # Pad with 1 pixel on all sides, replicate mode\n        # tf.pad expects [[top, bottom], [left, right], [channel_pad_top, channel_pad_bottom]]\n        # For replicate mode, we use 'REFLECT' which is often a good substitute for 'replicate' at edges in TF.\n        # This padding will make the 70x70 image into a 72x72 image.\n        x_padded = tf.pad(x_interpolated, [[1, 1], [1, 1], [0, 0]], mode='REFLECT') # (70+1+1, 70+1+1, C) = (72, 72, C)\n\n        # Permute back to (channels, height, width) as the model input expects channels-first\n        x_final = tf.transpose(x_padded, perm=[2, 0, 1]) # (C, 72, 72)\n        \n        return x_final\n\n    def _load_single_test_sample(self, file_path_tensor):\n        \"\"\"\n        Loads a single test sample and its OID (stem) given its file path.\n        This function is intended to be wrapped by tf.py_function.\n        It now includes the preprocessing step.\n\n        Args:\n            file_path_tensor (tf.Tensor): The path to the test .npy file, as a TensorFlow string tensor.\n\n        Returns:\n            tuple: A tuple containing the processed data (x) as tf.Tensor and the OID (stem) as tf.Tensor (string).\n        \"\"\"\n        # Convert the TensorFlow string tensor to a Python string\n        file_path = file_path_tensor.numpy().decode('utf-8')\n\n        try:\n            # Load the numpy array from the file\n            x_np = np.load(file_path)\n            # Convert to TensorFlow tensor for preprocessing\n            x_tf = tf.convert_to_tensor(x_np, dtype=tf.float32)\n        except Exception as e:\n            print(f\"Error loading test numpy file: {e}. File: {file_path}\")\n            raise e\n\n        # Apply preprocessing\n        x_processed = self._preprocess_test_sample(x_tf)\n\n        # Extract the stem (OID) from the file path\n        test_stem = os.path.basename(file_path).split(\".\")[0]\n\n        # Convert stem to TensorFlow Tensor\n        return x_processed, tf.convert_to_tensor(test_stem, dtype=tf.string)\n\n    def create_tf_dataset(self, include_oids=True):\n        \"\"\"\n        Creates and returns a tf.data.Dataset object for the test dataset.\n\n        Args:\n            include_oids (bool): If True, the dataset yields (x, oid). If False, it yields only x.\n\n        Returns:\n            tf.data.Dataset: A TensorFlow dataset ready for prediction.\n        \"\"\"\n        # Create a dataset from the list of test file paths\n        file_paths_ds = tf.data.Dataset.from_tensor_slices(self.test_files)\n\n        # Determine Tout based on include_oids\n        tout_types = [self.x_dtype, self.oid_dtype] if include_oids else [self.x_dtype]\n        \n        # Map the Python processing function (_load_single_test_sample) to each file path.\n        dataset = file_paths_ds.map(\n            lambda fp: tf.py_function(\n                self._load_single_test_sample, # The Python function to execute\n                inp=[fp],                     # Input tensor (file path)\n                Tout=tout_types,              # Expected output types (for compatibility)\n                name=\"load_and_process_sample_test\"\n            ),\n            num_parallel_calls=tf.data.AUTOTUNE # Automatically determines optimal number of parallel calls\n        )\n        \n        # Re-add tf.ensure_shape after tf.py_function for the output of tf.py_function\n        if include_oids:\n            dataset = dataset.map(\n                lambda x, oid: (tf.ensure_shape(x, self.x_shape), tf.ensure_shape(oid, self.oid_shape)),\n                num_parallel_calls=tf.data.AUTOTUNE\n            )\n        else:\n            dataset = dataset.map(\n                lambda x: tf.ensure_shape(x, self.x_shape),\n                num_parallel_calls=tf.data.AUTOTUNE\n            )\n\n        # No shuffling for test/prediction dataset\n        # Prefetch data to overlap data preprocessing with model execution\n        dataset = dataset.prefetch(buffer_size=tf.data.AUTOTUNE)\n\n        return dataset\n\n# --- Helper function for formatting time ---\ndef format_time(seconds):\n    \"\"\"Formats a duration in seconds into H:MM:SS string.\"\"\"\n    minutes, seconds = divmod(seconds, 60)\n    hours, minutes = divmod(minutes, 60)\n    return f\"{int(hours):02}:{int(minutes):02}:{int(seconds):02}\"\n\n# --- Custom Keras Layer for Noise Band Augmentation ---\nclass AddNoiseBands(tf.keras.layers.Layer):\n    \"\"\"\n    A Keras layer to add horizontal and/or vertical noise bands to input images.\n    Intended for data augmentation during training.\n    \"\"\"\n    def __init__(self, horizontal_bands_prob=0.3, vertical_bands_prob=0.3,\n                 max_band_width_h_ratio=0.05, max_band_width_v_ratio=0.05,\n                 max_noise_intensity=0.1, name=None, **kwargs):\n        super(AddNoiseBands, self).__init__(name=name, **kwargs)\n        self.horizontal_bands_prob = tf.constant(horizontal_bands_prob, dtype=tf.float32)\n        self.vertical_bands_prob = tf.constant(vertical_bands_prob, dtype=tf.float32)\n        self.max_band_width_h_ratio = tf.constant(max_band_width_h_ratio, dtype=tf.float32)\n        self.max_band_width_v_ratio = tf.constant(max_band_width_v_ratio, dtype=tf.float32)\n        self.max_noise_intensity = tf.constant(max_noise_intensity, dtype=tf.float32)\n\n    @tf.function\n    def call(self, inputs, training=None):\n        # Only apply augmentation during training\n        if training is False:\n            return inputs\n        if training is None:\n            # If `training` is None, assume training mode.\n            # This is common when the layer is used directly in a `model.call` or Sequential model\n            # where the `training` argument is propagated from `model.fit`.\n            pass \n\n        batch_size = tf.shape(inputs)[0]\n        height = tf.shape(inputs)[1]\n        width = tf.shape(inputs)[2]\n        channels = tf.shape(inputs)[3] # Assuming channels-last format here\n\n        current_output = inputs\n\n        def add_horizontal_band_fn():\n            # Calculate band width, cast height to float32 first for ratio calculation\n            band_width_h = tf.cast(tf.cast(height, tf.float32) * self.max_band_width_h_ratio, dtype=tf.int32)\n            band_width_h = tf.maximum(1, tf.minimum(band_width_h, height - 1)) # Ensure min 1 pixel and not exceed height\n\n            start_row_h = tf.random.uniform([], minval=0, maxval=height - band_width_h + 1, dtype=tf.int32)\n\n            noise_value_h_per_channel = tf.random.uniform((1, 1, 1, channels), # Noise applies per channel\n                                                           minval=-self.max_noise_intensity,\n                                                           maxval=self.max_noise_intensity)\n\n            noise_tensor = tf.random.normal(tf.shape(inputs), mean=0.0, stddev=1.0, dtype=inputs.dtype) * noise_value_h_per_channel\n            noise_tensor = tf.clip_by_value(noise_tensor, -self.max_noise_intensity, self.max_noise_intensity)\n\n            # XLA-COMPATIBLE MASKING FOR HORIZONTAL BAND\n            row_indices = tf.range(height, dtype=tf.int32)\n            # Create a boolean mask for rows within the band\n            is_in_band_h = tf.logical_and(row_indices >= start_row_h,\n                                        row_indices < start_row_h + band_width_h)\n            # Reshape and tile to match input tensor dimensions (batch, height, width, channels)\n            is_in_band_h = tf.reshape(is_in_band_h, (1, height, 1, 1))\n            is_in_band_h = tf.tile(is_in_band_h, [batch_size, 1, width, channels])\n\n            # Use tf.where to apply ones where the condition is true, zeros otherwise\n            band_mask_applied = tf.where(is_in_band_h, tf.ones_like(inputs, dtype=inputs.dtype), tf.zeros_like(inputs, dtype=inputs.dtype))\n\n            return current_output + noise_tensor * band_mask_applied\n\n        def add_vertical_band_fn():\n            # Calculate band width, cast width to float32 first for ratio calculation\n            band_width_v = tf.cast(tf.cast(width, tf.float32) * self.max_band_width_v_ratio, dtype=tf.int32)\n            band_width_v = tf.maximum(1, tf.minimum(band_width_v, width - 1)) # Ensure min 1 pixel and not exceed width\n\n            start_col_v = tf.random.uniform([], minval=0, maxval=width - band_width_v + 1, dtype=tf.int32)\n\n            noise_value_v_per_channel = tf.random.uniform((1, 1, 1, channels), # Noise applies per channel\n                                                           minval=-self.max_noise_intensity,\n                                                           maxval=self.max_noise_intensity)\n\n            noise_tensor = tf.random.normal(tf.shape(inputs), mean=0.0, stddev=1.0, dtype=inputs.dtype) * noise_value_v_per_channel\n            noise_tensor = tf.clip_by_value(noise_tensor, -self.max_noise_intensity, self.max_noise_intensity)\n\n            # XLA-COMPATIBLE MASKING FOR VERTICAL BAND\n            col_indices = tf.range(width, dtype=tf.int32)\n            # Create a boolean mask for columns within the band\n            is_in_band_v = tf.logical_and(col_indices >= start_col_v,\n                                        col_indices < start_col_v + band_width_v)\n            # Reshape and tile to match input tensor dimensions (batch, height, width, channels)\n            is_in_band_v = tf.reshape(is_in_band_v, (1, 1, width, 1))\n            is_in_band_v = tf.tile(is_in_band_v, [batch_size, height, 1, channels])\n\n            # Use tf.where to apply ones where the condition is true, zeros otherwise\n            band_mask_applied = tf.where(is_in_band_v, tf.ones_like(inputs, dtype=inputs.dtype), tf.zeros_like(inputs, dtype=inputs.dtype))\n\n            return current_output + noise_tensor * band_mask_applied\n\n        # Randomly apply horizontal and vertical bands with specified probabilities\n        current_output = tf.cond(tf.random.uniform(()) < self.horizontal_bands_prob,\n                                 add_horizontal_band_fn,\n                                 lambda: current_output)\n\n        current_output = tf.cond(tf.random.uniform(()) < self.vertical_bands_prob,\n                                 add_vertical_band_fn,\n                                 lambda: current_output)\n\n        return current_output\n\n    def get_config(self):\n        config = super(AddNoiseBands, self).get_config()\n        config.update({\n            'horizontal_bands_prob': self.horizontal_bands_prob.numpy(),\n            'vertical_bands_prob': self.vertical_bands_prob.numpy(),\n            'max_band_width_h_ratio': self.max_band_width_h_ratio.numpy(),\n            'max_band_width_v_ratio': self.max_band_width_v_ratio.numpy(),\n            'max_noise_intensity': self.max_noise_intensity.numpy(),\n        })\n        return config\n\n# --- Model Definition Function ---\ndef build_model(input_shape, output_shape, model_name=\"seismic_model\"):\n    \"\"\"\n    Builds the TensorFlow Keras model based on the provided architecture,\n    adapting to the specific input_shape.\n\n    Args:\n        input_shape (tuple): Expected input shape for the model (e.g., (5, 1000, 70) or (5, 72, 72)).\n        output_shape (tuple): Expected output shape for the model (e.g., (1, 70, 70)).\n        model_name (str): Name for the Keras model.\n\n    Returns:\n        tf.keras.Model: Compiled TensorFlow Keras model.\n    \"\"\"\n    from tensorflow.keras.applications import EfficientNetB2, MobileNetV2\n    from tensorflow.keras import layers, models\n\n    input_channels = input_shape[0]\n    input_height = input_shape[1]\n    input_width = input_shape[2]\n    \n    print(f\"\\n--- Building Model '{model_name}' with Input Shape: {input_shape} ---\")\n\n    # Input layer: raw seismic data in channels-first format.\n    inputs = layers.Input(shape=input_shape, name='seismic_input') # (None, C, H, W)\n\n    # --- BRANCH 1: First 3 channels with EfficientNetB2 ---\n    print(f\"--- Building Branch 1 (Channels 0, 1, 2) with EfficientNetB2 for {model_name} ---\")\n    x1 = layers.Lambda(lambda tensor: tensor[:, :3, :, :], name='select_3_channels_b1')(inputs) # (None, 3, H, W)\n    x1 = layers.Permute((2, 3, 1), name='permute_input_b1')(x1) # (None, H, W, 3)\n    \n    # Adaptive Pooling and Cropping based on input_height\n    if input_height == 1000 and input_width == 70:\n        # Original logic for (5, 1000, 70) data\n        print(f\"  Branch 1: Applying pooling (14,1) and cropping ((1,0), (0,0)) for 1000-height input.\")\n        x1 = layers.AveragePooling2D(pool_size=(14, 1), padding='valid', name='time_pooling_b1')(x1) # (None, 71, 70, 3)\n        x1 = layers.Cropping2D(cropping=((1, 0), (0, 0)), name='crop_time_b1')(x1) # (None, 70, 70, 3)\n    elif input_height == 72 and input_width == 72:\n        # Logic for (5, 72, 72) data - directly crop to 70x70\n        print(f\"  Branch 1: Applying cropping ((1,1), (1,1)) for 72x72 input.\")\n        x1 = layers.Cropping2D(cropping=((1, 1), (1, 1)), name='crop_72x72_b1')(x1) # (None, 70, 70, 3)\n    else:\n        raise ValueError(f\"Unsupported input spatial dimensions for Branch 1 in {model_name}: ({input_height}, {input_width}). \"\n                         \"Model expects (5, 1000, 70) or (5, 72, 72).\")\n\n    x1 = layers.Resizing(224, 224, name='resize_for_cnn_encoder_b1')(x1) # (None, 224, 224, 3)\n    \n    # Add Noise Bands for augmentation during training\n    x1 = AddNoiseBands(\n        horizontal_bands_prob=0.5, vertical_bands_prob=0.2,\n        max_band_width_h_ratio=0.01, max_band_width_v_ratio=0.01,\n        max_noise_intensity=0.2, name='noise_augmentation_b1')(x1)\n    \n    efficientnet_b2_model = EfficientNetB2(\n        include_top=False,\n        weights='imagenet',\n        input_shape=(224, 224, 3), \n        pooling=None # We'll add our own pooling layer later\n    )\n    efficientnet_b2_model.trainable = True # Fine-tune the EfficientNet model\n    x1 = efficientnet_b2_model(x1)\n    x1 = layers.GlobalAveragePooling2D(name='global_avg_pooling_b1')(x1)\n\n    # --- BRANCH 2: Remaining 2 channels with MobileNetV2 ---\n    print(f\"\\n--- Building Branch 2 (Channels 3, 4) with MobileNetV2 for {model_name} ---\")\n    x2 = layers.Lambda(lambda tensor: tensor[:, 3:, :, :], name='select_2_channels_b2')(inputs) # (None, 2, H, W)\n    \n    # MobileNetV2 expects 3 channels. Duplicate the last channel to make it 3-channel.\n    x2 = layers.Lambda(lambda t: tf.concat([t, t[:, -1:, :, :]], axis=1), name='duplicate_channel_to_3d_b2')(x2) # (None, 3, H, W)\n    \n    x2 = layers.Permute((2, 3, 1), name='permute_input_b2')(x2) # (None, H, W, 3)\n    \n    # Adaptive Pooling and Cropping based on input_height\n    if input_height == 1000 and input_width == 70:\n        # Logic for (5, 1000, 70) data\n        print(f\"  Branch 2: Applying pooling (14,1) and cropping ((1,0), (0,0)) for 1000-height input.\")\n        x2 = layers.AveragePooling2D(pool_size=(14, 1), padding='valid', name='time_pooling_b2')(x2) # (None, 71, 70, 3)\n        x2 = layers.Cropping2D(cropping=((1, 0), (0, 0)), name='crop_time_b2')(x2) # (None, 70, 70, 3)\n    elif input_height == 72 and input_width == 72:\n        # Logic for (5, 72, 72) data - directly crop to 70x70\n        print(f\"  Branch 2: Applying cropping ((1,1), (1,1)) for 72x72 input.\")\n        x2 = layers.Cropping2D(cropping=((1, 1), (1, 1)), name='crop_72x72_b2')(x2) # (None, 70, 70, 3)\n    else:\n        # This error should have been caught for Branch 1 already, but included for robustness\n        raise ValueError(f\"Unsupported input spatial dimensions for Branch 2 in {model_name}: ({input_height}, {input_width}). \"\n                         \"Model expects (5, 1000, 70) or (5, 72, 72).\")\n\n    x2 = layers.Resizing(224, 224, name='resize_for_cnn_encoder_b2')(x2) # (None, 224, 224, 3)\n    \n    # Add Noise Bands for augmentation during training\n    x2 = AddNoiseBands(\n        horizontal_bands_prob=0.5, vertical_bands_prob=0.2,\n        max_band_width_h_ratio=0.01, max_band_width_v_ratio=0.01,\n        max_noise_intensity=0.2, name='noise_augmentation_b2')(x2)\n    \n    mobilenet_v2_model = MobileNetV2(\n        include_top=False,\n        weights='imagenet',\n        input_shape=(224, 224, 3),\n        pooling=None # We'll add our own pooling layer later\n    )\n    mobilenet_v2_model.trainable = True # Fine-tune the MobileNetV2 model\n    x2 = mobilenet_v2_model(x2)\n    x2 = layers.GlobalAveragePooling2D(name='global_avg_pooling_b2')(x2)\n\n    # --- Feature Assembly / Combination ---\n    print(\"\\n--- Assembling Features ---\")\n    # Concatenate the features from both branches\n    combined_features = layers.Concatenate(axis=-1, name='combine_features')([x1, x2])\n\n    # --- Custom Head (now taking combined features) ---\n    print(\"\\n--- Building Custom Head ---\")\n    x = layers.Dense(512, activation='gelu', name='dense_head_1')(combined_features)\n    x = layers.Dropout(0.5, name='dropout_head_1')(x)\n    # x = layers.Dense(256, activation='gelu', name='dense_head_2')(x) # Optional layer, commented out as in your code\n    # x = layers.Dropout(0.5, name='dropout_head_2')(x) # Optional layer, commented out as in your code\n    \n    # The Dense layer's output units must match the total number of elements\n    # in the target output_shape (C * H * W).\n    x = layers.Dense(np.prod(output_shape), activation='linear', name='dense_output')(x)\n    outputs = layers.Reshape(output_shape, name='reshape_output')(x)\n\n    # Build the model.\n    model = models.Model(inputs=inputs, outputs=outputs, name=model_name)\n\n    # --- Compilation ---\n    # For seismic velocity prediction (regression), Mean Absolute Error (MAE) or\n    # Mean Squared Error (MSE) are common choices. MAE is more robust to outliers.\n    # MSE penalizes larger errors more heavily.\n    model.compile(optimizer='adam', loss='mae', metrics=['mae'])\n    \n    model.summary()\n    return model\n\n\nif __name__ == '__main__':\n    # Disable XLA JIT compilation globally for debugging.\n    # This might resolve 'layout failed' and 'CollectiveReduceV2' errors\n    # if they are due to XLA's aggressive graph optimizations.\n    tf.config.optimizer.set_jit(False)\n    print(\"XLA JIT compilation disabled for debugging.\")\n\n    # --- WandB Initialization ---\n    # This block retrieves your WandB API key from Kaggle secrets\n    # and initializes a new WandB run.\n    try:\n        user_secrets = UserSecretsClient()\n        secret_value_0 = user_secrets.get_secret(\"wandb_api\")\n        wandb.login(key=secret_value_0)\n        \n        # Define default config for WandB, can be overridden by sweep or command line\n        wandb_config_defaults = {\n            \"learning_rate\": 0.001,\n            \"epochs\": 5,\n            \"batch_size_per_replica\": 8, # Changed to per_replica batch size\n            \"batch_size_val_per_replica\": 4, # Changed to per_replica batch size\n            \"subsample\": None # Default to no subsampling\n        }\n        if EXPERIMENTAL_MODE:\n            wandb_config_defaults[\"subsample\"] = EXPERIMENTAL_SUBSAMPLE_LIMIT\n\n        wandb.init(project=\"seismic-velocity-adapted-dataset\", entity=\"crischir\", \n                   config=wandb_config_defaults)\n        print(\"WandB initialized successfully!\")\n    except Exception as e:\n        print(f\"Error initializing WandB: {e}. Please ensure 'wandb_api' secret is set on Kaggle.\")\n        wandb.init(mode=\"disabled\") # Disable wandb if login fails to allow script to continue\n        print(\"WandB is disabled. Script will continue without WandB logging.\")\n\n    # --- Setup for Multi-GPU Distribution Strategy ---\n    gpus = tf.config.list_physical_devices('GPU')\n    if len(gpus) > 1:\n        print(f\"Detected {len(gpus)} GPUs. Using MirroredStrategy for distributed training.\")\n        strategy = tf.distribute.MirroredStrategy()\n    else:\n        print(\"Detected 1 or no GPU. Using default strategy.\")\n        # If no GPUs, it will use CPU. If 1 GPU, it will use OneDeviceStrategy on that GPU.\n        strategy = tf.distribute.OneDeviceStrategy(device=\"/gpu:0\" if gpus else \"/cpu:0\")\n    \n    num_replicas = strategy.num_replicas_in_sync\n    print(f\"Number of replicas (devices in sync): {num_replicas}\")\n\n    # --- Configuration for your real data ---\n    # Set data_dir to the base directory where your .npy files are located.\n    # The paths in folds.csv (e.g., \"data/data_A_f0.npy\") will be joined with this data_dir.\n    REAL_DATA_DIR = \"/kaggle/input/openfwi-preprocessed-72x72/openfwi_72x72/\"\n    # Corrected path to the directory containing test .npy files\n    REAL_TEST_DIR = \"/kaggle/input/waveform-inversion/test/\" \n\n    # Initialize configuration object using WandB's config if available\n    cfg = Cfg(\n        data_dir=REAL_DATA_DIR,\n        subsample=wandb.config.subsample if wandb.run and \"subsample\" in wandb.config else None, \n        local_rank=0,        # Show tqdm progress bar\n        samples_per_record=500, # This must match the number of samples stored in each .npy file\n        # Batch sizes are per-replica from WandB, but cfg stores the global batch size.\n        # This will be overridden later with the global batch size.\n        batch_size_val=wandb.config.batch_size_val_per_replica if wandb.run and \"batch_size_val_per_replica\" in wandb.config else 4,\n    )\n    # Override subsample if in experimental mode (explicitly setting to EXPERIMENTAL_SUBSAMPLE_LIMIT)\n    if EXPERIMENTAL_MODE:\n        cfg.subsample = EXPERIMENTAL_SUBSAMPLE_LIMIT\n\n    # Calculate global batch sizes\n    GLOBAL_BATCH_SIZE_TRAIN = (wandb.config.batch_size_per_replica if wandb.run and \"batch_size_per_replica\" in wandb.config else 8) * num_replicas\n    GLOBAL_BATCH_SIZE_VAL = (wandb.config.batch_size_val_per_replica if wandb.run and \"batch_size_val_per_replica\" in wandb.config else 4) * num_replicas\n\n    # --- Infer Dataset Shapes for Model Building ---\n    # Infer shapes from the training dataset. This shape will be the input to the 'model'.\n    print(\"\\n--- Inferring Dataset Shapes for Model Building (using training data) ---\")\n    inferred_x_shape = None\n    inferred_y_shape = None\n\n    try:\n        # Use a minimal subsample to infer shapes quickly without loading too much data\n        temp_config_infer = Cfg(data_dir=REAL_DATA_DIR, subsample=1, local_rank=0, samples_per_record=500)\n        temp_train_dataset_adapter = CustomTFDataset(temp_config_infer, mode=\"train\")\n        inferred_x_shape = temp_train_dataset_adapter.x_shape\n        inferred_y_shape = temp_train_dataset_adapter.y_shape\n        print(f\"Inferred input (X) shape from dataset: {inferred_x_shape}\")\n        print(f\"Inferred output (Y) shape from dataset: {inferred_y_shape}\")\n        \n    except Exception as e:\n        print(f\"\\nCRITICAL ERROR: Could not infer dataset shapes. {e}\")\n        if wandb.run:\n            wandb.finish()\n        exit()\n\n    model = None\n\n    # Build and compile model within the distribution strategy scope\n    with strategy.scope():\n        try:\n            if inferred_x_shape and inferred_y_shape:\n                print(\"\\n--- Building Main Model ---\")\n                model = build_model(input_shape=inferred_x_shape, output_shape=inferred_y_shape, model_name=\"seismic_velocity_model\")\n            else:\n                print(\"Skipping model build due to missing shape inference.\")\n\n            # Load weights if path is provided and model was built\n            if model and LOAD_WEIGHTS_PATH:\n                if os.path.exists(LOAD_WEIGHTS_PATH):\n                    print(f\"Loading pre-trained weights from: {LOAD_WEIGHTS_PATH}\")\n                    try:\n                        # Load by name for custom layers like AddNoiseBands\n                        model.load_weights(LOAD_WEIGHTS_PATH)\n                        print(\"Weights loaded successfully.\")\n                    except Exception as e:\n                        print(f\"Error loading weights: {e}. Model will start with random weights.\")\n                else:\n                    print(f\"WARNING: No pre-trained weights found at {LOAD_WEIGHTS_PATH}. Model will initialize with random weights.\")\n\n        except Exception as e:\n            print(f\"\\nCRITICAL ERROR: Could not build or load weights for model. {e}\")\n            if wandb.run:\n                wandb.finish()\n            exit()\n\n\n    # 1. Create the training dataset and run training\n    if RUN_TRAIN and model:\n        print(\"\\n--- Initializing Training Dataset ---\")\n        try:\n            train_dataset_adapter = CustomTFDataset(cfg, mode=\"train\")\n            tf_train_dataset = train_dataset_adapter.create_tf_dataset()\n\n            # Distribute the training dataset across replicas\n            dist_train_dataset = strategy.experimental_distribute_dataset(tf_train_dataset.batch(GLOBAL_BATCH_SIZE_TRAIN))\n\n            # Initialize and prepare the validation dataset for validation during training\n            print(\"\\n--- Initializing Evaluation Dataset for Validation During Training ---\")\n            eval_dataset_adapter = CustomTFDataset(cfg, mode=\"eval\")\n            tf_eval_dataset = eval_dataset_adapter.create_tf_dataset()\n            dist_eval_dataset = strategy.experimental_distribute_dataset(tf_eval_dataset.batch(GLOBAL_BATCH_SIZE_VAL))\n\n\n            print(f\"\\n--- Running Training Loop with WandB Callback (Global Batch Size: {GLOBAL_BATCH_SIZE_TRAIN}) ---\")\n            history = model.fit( # Use the single model instance for training\n                dist_train_dataset,\n                epochs=wandb.config.epochs if wandb.run and \"epochs\" in wandb.config else 5, # Use epochs from wandb.config\n                callbacks=[WandbCallback(save_graph=False, save_model=False)], # Integrate WandB callback\n                validation_data=dist_eval_dataset, # Pass validation dataset for per-epoch evaluation\n                verbose=1 # Show training progress\n            )\n            print(\"Training complete. Check WandB dashboard for logs, including validation metrics.\")\n\n            # Save the trained model for later use\n            model_save_path = \"seismic_velocity_model.keras\" # Recommended Keras format for TF 2.x\n            print(f\"Saving the model to: {model_save_path}\")\n            model.save(model_save_path) # Save the trained model\n            print(\"Model saved successfully.\")\n\n        except Exception as e:\n            print(f\"\\nError initializing or training model: {e}\")\n            print(\"Please ensure the training data paths, model, and setup are correct.\")\n\n\n    # 2. Evaluation is now integrated into training, so the RUN_VALID block is removed.\n    #    If you need a separate evaluation run, re-enable the RUN_VALID flag and\n    #    copy the evaluation logic here, but remember to load the model first if RUN_TRAIN is False.\n    # if RUN_VALID and model:\n    #     print(\"\\n--- Initializing Evaluation Dataset ---\")\n    #     try:\n    #         eval_dataset_adapter = CustomTFDataset(cfg, mode=\"eval\")\n    #         tf_eval_dataset = eval_dataset_adapter.create_tf_dataset()\n\n    #         # Distribute the dataset across replicas\n    #         dist_eval_dataset = strategy.experimental_distribute_dataset(tf_eval_dataset.batch(GLOBAL_BATCH_SIZE_VAL))\n\n    #         print(f\"\\n--- Running Evaluation on Validation Dataset (Global Batch Size: {GLOBAL_BATCH_SIZE_VAL}) ---\")\n    #         # Evaluate the trained model\n    #         evaluation_results = model.evaluate(dist_eval_dataset, verbose=1) # Use main model for eval\n    #         print(f\"Validation Loss: {evaluation_results[0]:.4f}, Validation MAE: {evaluation_results[1]:.4f}\")\n            \n    #         # Log evaluation metrics to WandB\n    #         if wandb.run:\n    #             wandb.log({\"val_loss\": evaluation_results[0], \"val_mae\": evaluation_results[1]})\n    #             print(\"Validation metrics logged to WandB.\")\n\n    #     except Exception as e:\n    #         print(f\"\\nError initializing or evaluating model: {e}\")\n    #         print(\"Please ensure the evaluation data paths and file contents are correct.\")\n\n\n    # 3. Create the Test (Prediction) Dataset and run inference\n    if RUN_TEST and model: # Ensure model is built\n        print(\"\\n--- Initializing Test Dataset and Running Inference ---\")\n        row_count = 0\n        t0 = time.time()\n        \n        # Find all .npy files in the test directory\n        test_files_full = sorted(glob.glob(os.path.join(REAL_TEST_DIR, \"*.npy\")))\n        \n        # Apply experimental subsample limit if enabled\n        if EXPERIMENTAL_MODE and EXPERIMENTAL_SUBSAMPLE_LIMIT is not None:\n            test_files = test_files_full[:EXPERIMENTAL_SUBSAMPLE_LIMIT]\n            print(f\"EXPERIMENTAL_MODE: Limiting test files to {len(test_files)} out of {len(test_files_full)}\")\n        else:\n            test_files = test_files_full\n\n        if not test_files:\n            print(f\"Warning: No test files found in {REAL_TEST_DIR} after applying filters. Skipping inference.\")\n        else:\n            # Column names for the submission CSV, adapted from your original\n            x_cols = [f\"x_{i}\" for i in range(1, 70, 2)]\n            # Store integer indices corresponding to x_cols for easier numpy slicing\n            x_col_indices = [int(col.split('_')[1]) for col in x_cols] \n            fieldnames = [\"oid_ypos\"] + x_cols\n            \n            try:\n                # The CustomTFTestDataset will now preprocess the input 'x' to (C, 72, 72)\n                # This aligns with the model's expected input shape (C, H, W) where H=72, W=72 after its internal permute.\n                test_dataset_adapter = CustomTFTestDataset(cfg, test_files)\n                \n                # Create a dataset that yields only 'x' for prediction\n                tf_test_dataset_for_predict = test_dataset_adapter.create_tf_dataset(include_oids=False)\n                # Distribute this dataset\n                dist_test_dataset_for_predict = strategy.experimental_distribute_dataset(\n                    tf_test_dataset_for_predict.batch(GLOBAL_BATCH_SIZE_VAL)\n                )\n\n                # Create a separate dataset that yields (x, oid) for collecting OIDs\n                # This dataset is NOT batched or distributed as it's used for sequential OID collection.\n                tf_test_dataset_with_oids = test_dataset_adapter.create_tf_dataset(include_oids=True)\n\n                # Open the submission CSV file for writing\n                with open(\"submission.csv\", \"wt\", newline=\"\") as csvfile:\n                    writer = csv.DictWriter(csvfile, fieldnames=fieldnames)\n                    writer.writeheader()\n\n                    print(f\"Starting inference on {len(test_files)} test files...\")\n                    \n                    # Collect all OIDs upfront to ensure proper pairing with predictions\n                    all_oids = []\n                    # tqdm over the non-batched dataset to get OIDs one by one\n                    for _, oid_tensor in tqdm(tf_test_dataset_with_oids, total=test_dataset_adapter.total_samples, desc=\"Collecting OIDs\"):\n                        all_oids.append(oid_tensor.numpy().decode('utf-8'))\n                    \n                    # Perform prediction with the TensorFlow model directly on the distributed dataset\n                    all_outputs = model.predict(dist_test_dataset_for_predict, verbose=1) # verbose=1 to see progress\n\n                    # Ensure all_outputs and all_oids have the same length\n                    if len(all_outputs) != len(all_oids):\n                        raise ValueError(f\"Mismatch between number of predictions ({len(all_outputs)}) and OIDs collected ({len(all_oids)}).\")\n\n                    # Store a few outputs and oids for later plotting\n                    plot_y_preds = []\n                    plot_oids_test = []\n                    max_plots = 15 # Max number of plots (3x5 grid)\n\n                    # Iterate through the predictions and OIDs\n                    for i in tqdm(range(len(all_outputs)), desc=\"Writing Submission and Collecting Plots\"):\n                        y_pred_single_batched = all_outputs[i] # This is (1, 70, 70)\n                        oid_test = all_oids[i]\n\n                        # Collect samples for plotting from the first few\n                        if len(plot_y_preds) < max_plots:\n                            plot_y_preds.append(y_pred_single_batched[0, :, :]) # Squeeze channel dim for plotting\n                            plot_oids_test.append(oid_test)\n\n                        # Iterate through y_pos (0 to 69) and select specific x_pos for CSV\n                        y_pred_single = np.squeeze(y_pred_single_batched, axis=0) # (70, 70)\n                        for y_pos in range(70):\n                            row_values = y_pred_single[y_pos, x_col_indices] \n                            row = dict(zip(x_cols, row_values))\n                            row[\"oid_ypos\"] = f\"{oid_test}_y_{y_pos}\"\n                    \n                            writer.writerow(row)\n                            row_count += 1\n\n                            # Clear buffer periodically\n                            if row_count % 100_000 == 0:\n                                csvfile.flush()\n                \n                t1 = format_time(time.time() - t0)\n                print(f\"Inference complete. Total rows written: {row_count}\")\n                print(f\"Inference Time: {t1}\")\n\n                # --- Plotting predictions ---\n                if plot_y_preds: # Only plot if we collected any samples\n                    print(\"\\n--- Plotting a few predicted samples ---\")\n                    fig, axes = plt.subplots(3, 5, figsize=(15, 9)) # Adjust figsize for better view\n                    axes= axes.flatten()\n\n                    # Use the collected samples for plotting\n                    n = min(len(plot_y_preds), len(axes)) \n                    \n                    for i in range(n):\n                        img = plot_y_preds[i] # y_preds is (70, 70)\n                        idx = plot_oids_test[i] # Get the OID for the title\n                    \n                        # Plot\n                        axes[i].imshow(img, cmap='gray')\n                        axes[i].set_title(idx, fontsize=8) # Reduce font size if titles are long\n                        axes[i].axis('off')\n\n                    # Turn off any unused subplots\n                    for i in range(n, len(axes)):\n                        axes[i].axis('off')\n                    \n                    plt.tight_layout()\n                    plt.show()\n                else:\n                    print(\"No samples collected for plotting. Ensure there are test files and RUN_TEST is True.\")\n\n            except Exception as e:\n                print(f\"\\nError during test dataset initialization or inference: {e}\")\n                print(\"Please ensure the test data paths, model, and output logic are correct.\")\n    \n    # Finish the WandB run if it was initialized\n    if wandb.run:\n        wandb.finish()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T06:59:41.18232Z","iopub.execute_input":"2025-06-08T06:59:41.182522Z","execution_failed":"2025-06-08T07:05:39.443Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}