{"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"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\n\"\"\"\nYale/UNC-CH - Geophysical Waveform Inversion Competition - Baseline Script\n\nDescription:\nThis script implements a baseline machine learning model for the Kaggle\ncompetition focused on Full Waveform Inversion (FWI). The goal is to estimate\nsubsurface velocity models from seismic waveform data.\n\nApproach:\n- Model: Multi-Layer Perceptron (MLP) applied to pooled input features.\n- Input Processing: Max Pooling reduces dimensionality before MLP layers.\n- Features: LeakyReLU activations, Dropout, Kaiming Initialization.\n- Training: AdamW optimizer, ReduceLROnPlateau scheduler, L1Loss (MAE),\n            Mixed Precision (AMP) on GPU, Early Stopping.\n- Data Handling: PyTorch Dataset with memory mapping for large files.\n- Augmentation: Training data augmentation (noise, horizontal flip).\n- Evaluation: Mean Absolute Error (MAE), as per competition rules.\n- Inference: Test-Time Augmentation (TTA) using horizontal flips.\n\nStructure:\n1. Imports\n2. Configuration (CFG Class)\n3. Utility Functions (Seeding)\n4. Data Loading & Preparation (Datasets, Splitting, Augmentation)\n5. Model Definition (FWINet Class)\n6. Training & Validation Functions\n7. Prediction & Submission Functions\n8. Main Execution Block\n\n\"\"\"\n\n# ======================================================\n# 1. Library Imports\n# ======================================================\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nimport csv\nimport time\nimport gc  # Garbage Collection\nimport random\nimport os\n\nprint(\"Libraries imported successfully.\")\n\n# Set Matplotlib style for better plots\nplt.style.use('seaborn-v0_8-whitegrid')\n\n# ======================================================\n# 2. Configuration\n# ======================================================\nclass CFG:\n    \"\"\"Configuration class for hyperparameters and settings.\"\"\"\n    # Paths\n    BASE_DIR = '/kaggle/input/waveform-inversion'\n    TRAIN_DIR = os.path.join(BASE_DIR, 'train_samples')\n    TEST_DIR = os.path.join(BASE_DIR, 'test')\n    OUTPUT_DIR = '/kaggle/working/' # Writable directory in Kaggle\n    SUBMISSION_FILE = os.path.join(OUTPUT_DIR, 'submission.csv')\n    BEST_MODEL_PATH = os.path.join(OUTPUT_DIR, 'best_model_final.pth') # Path to save best model\n\n    # Data Parameters\n    N_EXAMPLES_PER_FILE = 500  # Assumed number of samples in each .npy file\n    IMG_HEIGHT = 70 # Target velocity model height\n    IMG_WIDTH = 70  # Target velocity model width\n    INPUT_CHANNELS = 5 # Deduced from typical seismic data shape (Channels, Time, Width)\n    INPUT_TIME_STEPS = 1000 # Deduced from typical seismic data shape\n\n    # Model Parameters\n    POOL_SIZE = (8, 2) # Kernel size for Max Pooling (Time/Depth, Width)\n    HIDDEN_SIZE_FACTOR = 1.0 # Factor to scale hidden layer size relative to output size\n    DROPOUT_RATE = 0.3 # Dropout probability for regularization\n    OUTPUT_SCALE = 1000.0 # Scaling factor applied to model output\n    OUTPUT_OFFSET = 1500.0 # Offset applied to model output (adjusts range)\n\n    # Training Parameters\n    SEED = 42 # Random seed for reproducibility\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # Auto-detect GPU\n    N_EPOCHS = 70 # Maximum number of training epochs (use early stopping)\n    BATCH_SIZE = 32 # Samples per batch (adjust based on GPU memory)\n    NUM_WORKERS = 2 # Number of parallel workers for DataLoader (adjust based on Kaggle limits)\n    LEARNING_RATE = 3e-4 # Initial learning rate for AdamW\n    WEIGHT_DECAY = 1e-5 # Weight decay for AdamW (L2 regularization)\n    LR_SCHEDULER_PATIENCE = 5 # Epochs to wait for improvement before reducing LR\n    LR_SCHEDULER_FACTOR = 0.5 # Factor by which LR is reduced\n    EARLY_STOPPING_PATIENCE = 10 # Epochs to wait for improvement before stopping training\n    GRADIENT_CLIP_NORM = 1.0 # Maximum norm for gradient clipping\n    VALIDATION_SPLIT_RATIO = 0.5 # Use every 1/ratio file for validation (0.5 -> every 2nd file)\n\n    # Augmentation\n    AUGMENT_TRAIN = True # Enable/disable training data augmentation\n    AUG_NOISE_LEVEL_MIN = 0.001 # Min noise level for augmentation\n    AUG_NOISE_LEVEL_MAX = 0.01 # Max noise level for augmentation\n    AUG_FLIP_PROB = 0.5 # Probability of applying horizontal flip augmentation\n\n    # TTA Parameters\n    USE_TTA = True # Enable/disable Test-Time Augmentation\n    TTA_N_AUGMENTATIONS = 1 # Number of *additional* TTA predictions (currently only flip, so 1)\n\n    # Visualization\n    VISUALIZATION_INTERVAL = 5 # Visualize validation predictions every N epochs\n\nprint(f\"Configuration loaded. Device set to: {CFG.DEVICE}\")\n\n# ======================================================\n# 3. Utility Functions\n# ======================================================\ndef set_seed(seed=CFG.SEED):\n    \"\"\"Sets random seeds for reproducibility across libraries.\"\"\"\n    print(f\"Setting random seed to {seed}\")\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed) # For multi-GPU setups\n        # Ensure deterministic behavior for CuDNN operations if reproducibility is critical\n        # Note: This can impact performance. Set False if speed is prioritized over exact reproducibility.\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False # Disable benchmarking for determinism\n    print(\"Seed setting complete.\")\n\n# Set the seed globally at the start of the script\nset_seed(CFG.SEED)\n\n# ======================================================\n# 4. Data Loading and Preparation\n# ======================================================\ndef find_data_files(base_dir):\n    \"\"\"Finds all input seismic data files based on naming conventions.\"\"\"\n    print(f\"Searching for input data files in: {base_dir}\")\n    # Sort files for consistent splitting behavior across runs\n    files = sorted([\n        f for f in Path(base_dir).rglob('*.npy')\n        if ('seis' in f.stem) or ('data' in f.stem)\n    ])\n    print(f\"Found {len(files)} potential input files.\")\n    if not files:\n        print(\"Warning: No input files found in the specified directory.\")\n    return files\n\ndef map_inputs_to_outputs(input_files):\n    \"\"\"Maps input file paths to corresponding output velocity model file paths, checking existence.\"\"\"\n    output_files = []\n    missing_outputs = []\n    filtered_inputs = []\n    print(\"Mapping input files to output files and verifying existence...\")\n    for f in tqdm(input_files, desc=\"Mapping files\", leave=False, dynamic_ncols=True):\n        # Construct the expected output filename based on convention\n        out_f = Path(str(f).replace('seis', 'vel').replace('data', 'model'))\n        if out_f.exists():\n            output_files.append(out_f)\n            filtered_inputs.append(f) # Only keep input if its corresponding output exists\n        else:\n            print(f\"Warning: Output file not found for input {f.name}, skipping this pair.\")\n            missing_outputs.append(f)\n\n    if missing_outputs:\n         print(f\"Warning: {len(missing_outputs)} input files were skipped due to missing outputs.\")\n\n    assert len(filtered_inputs) == len(output_files), \"Input and output file counts mismatch after filtering.\"\n    print(f\"Successfully mapped {len(filtered_inputs)} valid input/output file pairs.\")\n    return filtered_inputs, output_files\n\ndef split_train_validation(all_inputs, all_outputs, split_ratio=CFG.VALIDATION_SPLIT_RATIO):\n    \"\"\"\n    Splits the data into training and validation sets based on file indices using a stride.\n    \"\"\"\n    assert 0 < split_ratio < 1, \"Validation split ratio must be between 0 and 1.\"\n    num_files = len(all_inputs)\n    if num_files == 0:\n        raise ValueError(\"Cannot split empty file lists.\")\n\n    # Calculate stride (e.g., split_ratio=0.5 -> stride=2 -> use every 2nd file for validation)\n    stride = max(1, int(round(1.0 / split_ratio))) # Ensure stride is at least 1\n    print(f\"Splitting data with validation ratio ~{1.0/stride:.2f} (using stride={stride})\")\n\n    # Select indices for the validation set using the calculated stride\n    valid_indices = list(range(0, num_files, stride))\n\n    # Create lists based on selected indices\n    valid_inputs = [all_inputs[i] for i in valid_indices]\n    valid_outputs = [all_outputs[i] for i in valid_indices]\n\n    train_inputs = [f for i, f in enumerate(all_inputs) if i not in valid_indices]\n    train_outputs = [f for i, f in enumerate(all_outputs) if i not in valid_indices]\n\n    print(f\"Total file pairs: {num_files}\")\n    print(f\"Training file pairs: {len(train_inputs)}\")\n    print(f\"Validation file pairs: {len(valid_inputs)}\")\n\n    if not train_inputs or not valid_inputs:\n         print(\"Warning: Training or validation set is empty after splitting. Check ratio and file count.\")\n\n    return train_inputs, train_outputs, valid_inputs, valid_outputs\n\nclass SeismicDataset(Dataset):\n    \"\"\"\n    PyTorch Dataset for loading seismic data (input) and velocity models (target).\n    Utilizes memory mapping for efficient handling of potentially large files.\n    Includes options for data augmentation during training.\n    \"\"\"\n    def __init__(self, inputs_files, output_files, n_examples_per_file=CFG.N_EXAMPLES_PER_FILE, augment=False, cfg=CFG):\n        \"\"\"Initializes the dataset.\"\"\"\n        assert len(inputs_files) == len(output_files), \"Input and output file counts must match.\"\n        if not inputs_files:\n            print(\"Warning: Initializing SeismicDataset with zero files.\")\n        self.inputs_files = inputs_files\n        self.output_files = output_files\n        self.n_examples_per_file = n_examples_per_file\n        self.augment = augment\n        self.cfg = cfg\n        # Basic check for file existence (checks only the first file pair for speed)\n        self._check_first_file_exists()\n\n    def _check_first_file_exists(self):\n        \"\"\"Quick check if the first files in the lists exist.\"\"\"\n        if self.inputs_files and not self.inputs_files[0].exists():\n             raise FileNotFoundError(f\"First input file specified does not exist: {self.inputs_files[0]}\")\n        if self.output_files and not self.output_files[0].exists():\n             raise FileNotFoundError(f\"First output file specified does not exist: {self.output_files[0]}\")\n        print(\"First input/output file pair checked successfully.\")\n\n    def __len__(self):\n        \"\"\"Returns the total number of samples across all files.\"\"\"\n        return len(self.inputs_files) * self.n_examples_per_file\n\n    def _apply_augmentation(self, x, y):\n        \"\"\"Applies configured data augmentation techniques (noise, flip).\"\"\"\n        # Ensure operating on copies\n        x = x.copy()\n        y = y.copy()\n\n        # 1. Random Noise Injection\n        if np.random.random() < 0.5: # 50% chance to add noise\n            noise_level = np.random.uniform(self.cfg.AUG_NOISE_LEVEL_MIN, self.cfg.AUG_NOISE_LEVEL_MAX)\n            noise = noise_level * np.random.randn(*x.shape).astype(np.float32)\n            x = x + noise\n\n        # 2. Random Horizontal Flip\n        if np.random.random() < self.cfg.AUG_FLIP_PROB:\n            # Flip input (Channels, Time, Width) along Width axis (axis=2)\n            # Flip target (1, Height, Width) along Width axis (axis=2 if dim is 3, axis=1 if dim is 2)\n            x = np.flip(x, axis=2).copy() # .copy() ensures positive strides\n            # Adjust target flip axis based on actual dimensions before applying this!\n            # Assuming y is [1, H, W] -> flip axis 2. If y is [H, W] -> flip axis 1.\n            y_flip_axis = 2 if y.ndim == 3 else 1\n            y = np.flip(y, axis=y_flip_axis).copy()\n\n        return x, y\n\n    def __getitem__(self, idx):\n        \"\"\"Loads a single sample (input, target) using memory mapping.\"\"\"\n        if len(self.inputs_files) == 0:\n            raise IndexError(\"Dataset is empty, cannot retrieve item.\")\n\n        # Determine which file and sample within that file corresponds to the global index 'idx'\n        file_idx = idx // self.n_examples_per_file\n        sample_idx = idx % self.n_examples_per_file\n\n        if file_idx >= len(self.inputs_files):\n             raise IndexError(f\"Calculated file index {file_idx} is out of bounds for {len(self.inputs_files)} files.\")\n\n        input_file_path = self.inputs_files[file_idx]\n        output_file_path = self.output_files[file_idx]\n\n        try:\n            # Use memory mapping ('r' mode) for efficient read-only access.\n            X_mmap = np.load(input_file_path, mmap_mode='r')\n            y_mmap = np.load(output_file_path, mmap_mode='r')\n\n            # Check if sample_idx is valid for the loaded memory maps\n            if sample_idx >= X_mmap.shape[0] or sample_idx >= y_mmap.shape[0]:\n                raise IndexError(f\"Sample index {sample_idx} out of bounds for file {input_file_path.name} (Shapes: X={X_mmap.shape}, y={y_mmap.shape})\")\n\n            # Access the specific sample, triggering data read.\n            # Convert to float32 and copy data from mmap to ensure it's in memory.\n            X_sample = X_mmap[sample_idx].copy().astype(np.float32)\n            y_sample = y_mmap[sample_idx].copy().astype(np.float32)\n\n            # Explicitly close mmap objects (optional but good practice)\n            del X_mmap, y_mmap\n\n            # Ensure target has a channel dimension: [H, W] -> [1, H, W]\n            if y_sample.ndim == 2:\n                y_sample = np.expand_dims(y_sample, axis=0)\n\n            # Apply augmentation only if specified for this dataset instance (typically training set)\n            if self.augment:\n                X_sample, y_sample = self._apply_augmentation(X_sample, y_sample)\n\n            # Convert numpy arrays to PyTorch tensors\n            X_tensor = torch.from_numpy(X_sample)\n            y_tensor = torch.from_numpy(y_sample)\n\n            return X_tensor, y_tensor\n\n        except FileNotFoundError as e:\n            print(f\"FATAL ERROR: File not found during loading! {e}\")\n            raise e\n        except IndexError as e:\n            print(f\"FATAL ERROR: Index out of bounds during loading! {e}\")\n            raise e\n        except Exception as e:\n            print(f\"FATAL ERROR loading data: File='{input_file_path.name}', Sample Index={sample_idx}. Exception: {e}\")\n            raise e\n\nprint(\"Dataset class defined.\")\n\n# ======================================================\n# 5. Model Definition\n# ======================================================\nclass FWINet(nn.Module):\n    \"\"\"\n    Feedforward Network (MLP) with initial Pooling layer for Full Waveform Inversion.\n    Takes seismic data (Batch, Channels, Time, Width) and predicts velocity maps (Batch, 1, Height, Width).\n    \"\"\"\n    def __init__(self, cfg=CFG):\n        super().__init__()\n        self.cfg = cfg\n\n        # Input dimensions\n        input_channels = cfg.INPUT_CHANNELS\n        input_time_steps = cfg.INPUT_TIME_STEPS\n        input_width = cfg.IMG_WIDTH # Assuming input width relevant for pooling\n\n        # 1. Pooling layer: Reduces Time/Depth and Width dimensions\n        self.pool = nn.MaxPool2d(kernel_size=cfg.POOL_SIZE)\n\n        # Calculate the flattened feature size after pooling\n        pooled_time = input_time_steps // cfg.POOL_SIZE[0]\n        pooled_width = input_width // cfg.POOL_SIZE[1]\n        flattened_size = input_channels * pooled_time * pooled_width\n        if flattened_size <= 0:\n             raise ValueError(f\"Flattened size after pooling is non-positive ({flattened_size}). Check input dims ({input_channels}x{input_time_steps}x{input_width}) and POOL_SIZE {cfg.POOL_SIZE}.\")\n        print(f\"Flattened feature size after pooling: {flattened_size}\")\n\n        # 2. MLP Layers\n        output_size = cfg.IMG_HEIGHT * cfg.IMG_WIDTH\n        hidden_size_base = int(output_size * cfg.HIDDEN_SIZE_FACTOR)\n        # Define hidden layer sizes, ensuring a minimum size relative to output\n        hidden_sizes = [\n            max(output_size // 2, hidden_size_base),\n            max(output_size // 4, hidden_size_base // 2),\n            max(output_size // 8, hidden_size_base // 4)\n        ]\n        print(f\"MLP hidden layer sizes: {hidden_sizes}\")\n\n        self.network = nn.Sequential(\n            nn.Linear(flattened_size, hidden_sizes[0]),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Dropout(cfg.DROPOUT_RATE),\n\n            nn.Linear(hidden_sizes[0], hidden_sizes[1]),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Dropout(cfg.DROPOUT_RATE),\n\n            nn.Linear(hidden_sizes[1], hidden_sizes[2]),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Dropout(cfg.DROPOUT_RATE),\n\n            nn.Linear(hidden_sizes[2], output_size) # Final layer projects to flattened output map size\n        )\n\n        # Apply Kaiming initialization, suitable for LeakyReLU\n        self._initialize_weights()\n        print(\"Model initialized with Kaiming Normal weights.\")\n\n    def _initialize_weights(self):\n        \"\"\"Initializes Linear layer weights using Kaiming Normal.\"\"\"\n        for m in self.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.kaiming_normal_(m.weight, a=0.2, mode='fan_in', nonlinearity='leaky_relu')\n                if m.bias is not None:\n                    nn.init.zeros_(m.bias)\n\n    def forward(self, x):\n        \"\"\"Defines the forward pass of the model.\"\"\"\n        batch_size = x.shape[0]\n\n        # Ensure input tensor is float32\n        x = x.float()\n\n        # --- 1. Input Feature Scaling (Instance-like normalization per sample) ---\n        mean = torch.mean(x, dim=(2, 3), keepdim=True)\n        std = torch.std(x, dim=(2, 3), keepdim=True)\n        x_norm = (x - mean) / (std + 1e-8) # Epsilon for numerical stability\n\n        # --- 2. Max Pooling ---\n        x_pooled = self.pool(x_norm)\n\n        # --- 3. Flatten Features ---\n        x_flattened = x_pooled.view(batch_size, -1) # Flatten all dims except batch\n\n        # --- 4. Pass through MLP ---\n        output_flat = self.network(x_flattened)\n\n        # --- 5. Reshape to Output Image Dimensions ---\n        # Reshape from (batch, height * width) -> (batch, 1, height, width)\n        output_image = output_flat.view(batch_size, 1, self.cfg.IMG_HEIGHT, self.cfg.IMG_WIDTH)\n\n        # --- 6. Apply Output Scaling and Offset ---\n        # Map network output to the expected physical velocity range\n        output_scaled = output_image * self.cfg.OUTPUT_SCALE + self.cfg.OUTPUT_OFFSET\n\n        return output_scaled\n\nprint(\"Model class 'FWINet' defined.\")\n\n# ======================================================\n# 6. Training and Validation Functions\n# ======================================================\ndef train_one_epoch(model, dataloader, criterion, optimizer, device, scaler, cfg):\n    \"\"\"Performs one training epoch.\"\"\"\n    model.train()\n    total_loss = 0.0\n    progress_bar = tqdm(dataloader, desc=f'Training Epoch', leave=False, dynamic_ncols=True)\n\n    for inputs, targets in progress_bar:\n        inputs = inputs.to(device, non_blocking=True).float()\n        targets = targets.to(device, non_blocking=True).float()\n\n        optimizer.zero_grad(set_to_none=True)\n\n        with torch.cuda.amp.autocast(enabled=(scaler is not None)):\n            outputs = model(inputs)\n            loss = criterion(outputs, targets)\n\n        if scaler: # Mixed precision backward pass\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer) # Unscale before clipping\n            torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.GRADIENT_CLIP_NORM)\n            scaler.step(optimizer)\n            scaler.update()\n        else: # Standard precision backward pass\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.GRADIENT_CLIP_NORM)\n            optimizer.step()\n\n        loss_item = loss.item()\n        total_loss += loss_item\n        progress_bar.set_postfix(loss=f'{loss_item:.4f}')\n\n    avg_loss = total_loss / len(dataloader)\n    # print(f\"Epoch Training Avg Loss: {avg_loss:.5f}\")\n    return avg_loss\n\n\ndef validate_one_epoch(model, dataloader, criterion, device, cfg):\n    \"\"\"Performs one validation epoch.\"\"\"\n    model.eval()\n    total_loss = 0.0\n    first_batch_outputs = None\n    first_batch_targets = None\n    progress_bar = tqdm(dataloader, desc=f'Validation Epoch', leave=False, dynamic_ncols=True)\n\n    with torch.inference_mode():\n        for i, (inputs, targets) in enumerate(progress_bar):\n            inputs = inputs.to(device, non_blocking=True).float()\n            targets = targets.to(device, non_blocking=True).float()\n\n            with torch.cuda.amp.autocast(enabled=(device == torch.device('cuda'))):\n                 outputs = model(inputs)\n\n            loss = criterion(outputs, targets)\n            loss_item = loss.item()\n            total_loss += loss_item\n            progress_bar.set_postfix(loss=f'{loss_item:.4f}')\n\n            # Store the first sample of the first batch for visualization\n            if i == 0:\n                first_batch_outputs = outputs[0:1].detach().cpu()\n                first_batch_targets = targets[0:1].detach().cpu()\n\n    avg_loss = total_loss / len(dataloader)\n    # print(f\"Epoch Validation Avg Loss: {avg_loss:.5f}\")\n    return avg_loss, first_batch_outputs, first_batch_targets\n\n\ndef visualize_prediction(target, prediction, epoch, loss, cfg):\n    \"\"\"Visualizes a comparison between ground truth and model prediction.\"\"\"\n    if target is None or prediction is None:\n        print(f\"Epoch {epoch}: Skipping visualization due to missing sample data.\")\n        return\n\n    target_np = target.squeeze().numpy()\n    prediction_np = prediction.squeeze().numpy()\n\n    fig, axes = plt.subplots(1, 2, figsize=(12, 5.5))\n    fig.suptitle(f'Epoch {epoch} | Validation MAE Loss: {loss:.5f}', fontsize=14, y=0.98)\n\n    # Determine shared color range using percentiles for robustness\n    vmin = np.percentile(target_np, 1)\n    vmax = np.percentile(target_np, 99)\n\n    # Plot Ground Truth\n    im1 = axes[0].imshow(target_np, cmap='viridis', vmin=vmin, vmax=vmax, aspect='auto')\n    axes[0].set_title('Ground Truth Velocity')\n    axes[0].set_xlabel('Width Index')\n    axes[0].set_ylabel('Depth/Time Index')\n    fig.colorbar(im1, ax=axes[0], label='Velocity (units)', fraction=0.046, pad=0.04) # Adjust unit label if known\n\n    # Plot Prediction\n    im2 = axes[1].imshow(prediction_np, cmap='viridis', vmin=vmin, vmax=vmax, aspect='auto')\n    axes[1].set_title('Predicted Velocity')\n    axes[1].set_xlabel('Width Index')\n    axes[1].set_ylabel('Depth/Time Index')\n    fig.colorbar(im2, ax=axes[1], label='Velocity (units)', fraction=0.046, pad=0.04)\n\n    plt.tight_layout(rect=[0, 0.03, 1, 0.95])\n\n    fig_path = os.path.join(cfg.OUTPUT_DIR, f'validation_epoch_{epoch}.png')\n    try:\n        plt.savefig(fig_path, dpi=150)\n        # print(f\"Saved validation visualization: {fig_path}\") # Reduce verbose printing\n    except Exception as e:\n        print(f\"Warning: Error saving visualization - {e}\")\n    plt.show() # Display plot in notebook context\n\ndef plot_history(history, cfg):\n    \"\"\"Plots training/validation loss and learning rate curves.\"\"\"\n    epochs = range(1, len(history['train_loss']) + 1)\n    if not epochs:\n        print(\"No history data to plot.\")\n        return\n\n    fig, ax1 = plt.subplots(figsize=(12, 5))\n\n    # Plot Losses\n    color = 'tab:blue'\n    ax1.set_xlabel('Epoch', fontsize=12)\n    ax1.set_ylabel('MAE Loss', color=color, fontsize=12)\n    ax1.plot(epochs, history['train_loss'], color=color, linestyle='-', marker='o', markersize=4, label='Train Loss')\n    ax1.plot(epochs, history['valid_loss'], color='tab:orange', linestyle='--', marker='x', markersize=4, label='Validation Loss')\n    ax1.tick_params(axis='y', labelcolor=color, labelsize=10)\n    ax1.tick_params(axis='x', labelsize=10)\n    ax1.legend(loc='upper left', fontsize=10)\n    ax1.grid(True, linestyle='--', alpha=0.6)\n\n    # Plot Learning Rate on secondary y-axis\n    ax2 = ax1.twinx()\n    color = 'tab:green'\n    ax2.set_ylabel('Learning Rate', color=color, fontsize=12)\n    ax2.plot(epochs, history['lr'], color=color, linestyle=':', marker='s', markersize=4, label='Learning Rate')\n    ax2.tick_params(axis='y', labelcolor=color, labelsize=10)\n    ax2.ticklabel_format(style='sci', axis='y', scilimits=(-5, 4))\n    ax2.legend(loc='upper right', fontsize=10)\n\n    plt.title('Training History: Loss & Learning Rate', fontsize=14)\n    fig.tight_layout()\n\n    history_plot_path = os.path.join(cfg.OUTPUT_DIR, 'training_history.png')\n    try:\n        plt.savefig(history_plot_path, dpi=150)\n        print(f\"📊 Saved training history plot: {history_plot_path}\")\n    except Exception as e:\n        print(f\"Warning: Error saving history plot - {e}\")\n    plt.show()\n\n\ndef run_training(model, train_loader, valid_loader, cfg):\n    \"\"\"Coordinates the model training process.\"\"\"\n    print(f\"🚀 Starting training run...\")\n    model.to(cfg.DEVICE)\n    criterion = nn.L1Loss() # MAE Loss\n    optimizer = optim.AdamW(model.parameters(), lr=cfg.LEARNING_RATE, weight_decay=cfg.WEIGHT_DECAY)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=cfg.LR_SCHEDULER_FACTOR, patience=cfg.LR_SCHEDULER_PATIENCE, verbose=True)\n    scaler = torch.cuda.amp.GradScaler() if cfg.DEVICE == torch.device('cuda') else None\n    if scaler: print(\"AMP enabled.\")\n\n    history = {'train_loss': [], 'valid_loss': [], 'lr': []}\n    best_valid_loss = float('inf')\n    best_model_state = None\n    epochs_no_improve = 0\n    start_time = time.time()\n\n    for epoch in range(1, cfg.N_EPOCHS + 1):\n        epoch_start_time = time.time()\n        print(f\"\\n===== Epoch {epoch}/{cfg.N_EPOCHS} =====\")\n\n        train_loss = train_one_epoch(model, train_loader, criterion, optimizer, cfg.DEVICE, scaler, cfg)\n        valid_loss, sample_output, sample_target = validate_one_epoch(model, valid_loader, criterion, cfg.DEVICE, cfg)\n\n        history['train_loss'].append(train_loss)\n        history['valid_loss'].append(valid_loss)\n        current_lr = optimizer.param_groups[0]['lr']\n        history['lr'].append(current_lr)\n        scheduler.step(valid_loss)\n\n        epoch_duration = time.time() - epoch_start_time\n        print(f\"Epoch {epoch} Summary | Train Loss: {train_loss:.5f} | Valid Loss: {valid_loss:.5f} (Best: {best_valid_loss:.5f}) | LR: {current_lr:.1e} | Time: {epoch_duration:.2f}s\")\n\n        if valid_loss < best_valid_loss:\n            best_valid_loss = valid_loss\n            best_model_state = model.state_dict().copy()\n            epochs_no_improve = 0\n            print(f\"✨ Validation Loss Improved! Saving best model state.\")\n            # Save immediately in case of interruption\n            try:\n                 torch.save(best_model_state, cfg.BEST_MODEL_PATH)\n                 # print(f\"Best model state saved to {cfg.BEST_MODEL_PATH}\") # Less verbose\n            except Exception as e:\n                 print(f\"Warning: Error saving best model state - {e}\")\n        else:\n            epochs_no_improve += 1\n            print(f\"Patience: {epochs_no_improve}/{cfg.EARLY_STOPPING_PATIENCE}\")\n\n        if epoch % cfg.VISUALIZATION_INTERVAL == 0 or epoch == 1:\n            visualize_prediction(sample_target, sample_output, epoch, valid_loss, cfg)\n\n        if epochs_no_improve >= cfg.EARLY_STOPPING_PATIENCE:\n            print(f\"\\n🛑 Early stopping triggered after {epoch} epochs.\")\n            break\n\n        gc.collect()\n        if cfg.DEVICE == torch.device('cuda'):\n            torch.cuda.empty_cache()\n\n    total_training_time = time.time() - start_time\n    print(f\"\\n===== Training Finished =====\")\n    print(f\"Total Training Time: {total_training_time / 60:.2f} minutes\")\n    print(f\"Best Validation MAE: {best_valid_loss:.5f}\")\n\n    if best_model_state:\n        print(\"Loading best model weights achieved during training...\")\n        model.load_state_dict(best_model_state)\n        # Final save already happened when loss improved\n    else:\n        print(\"Warning: No improvement observed or training too short. Using final model state.\")\n\n    plot_history(history, cfg)\n    return history, model # Return history and model with best weights\n\nprint(\"Training and validation helper functions defined.\")\n\n# ======================================================\n# 7. Prediction and Submission Functions\n# ======================================================\nclass TestDataset(Dataset):\n    \"\"\"Dataset for loading test data files for inference.\"\"\"\n    def __init__(self, test_files):\n        self.test_files = sorted(test_files)\n        if not self.test_files: print(\"Warning: TestDataset initialized with zero files.\")\n\n    def __len__(self):\n        return len(self.test_files)\n\n    def __getitem__(self, i):\n        if i >= len(self.test_files): raise IndexError(\"Test dataset index out of range.\")\n        test_file_path = self.test_files[i]\n        try:\n            data = np.load(test_file_path).astype(np.float32)\n            oid = test_file_path.stem # Filename without extension is the ID\n            return torch.from_numpy(data), oid\n        except Exception as e:\n            print(f\"ERROR loading test file: {test_file_path}. Exception: {e}\")\n            raise e\n\ndef generate_submission(model, test_files, cfg):\n    \"\"\"Generates predictions on test data and creates the submission file.\"\"\"\n    print(\"\\n===== Generating Submission File =====\")\n    if not test_files:\n        print(\"No test files found. Skipping submission.\")\n        return\n\n    test_dataset = TestDataset(test_files)\n    inference_batch_size = max(1, cfg.BATCH_SIZE // 2) # Use potentially smaller batch for inference\n    test_loader = DataLoader(test_dataset, batch_size=inference_batch_size, shuffle=False, num_workers=cfg.NUM_WORKERS, pin_memory=True)\n    print(f\"Created Test DataLoader: Batches={len(test_loader)}, Batch Size={inference_batch_size}\")\n\n    # Required columns for submission (oid_ypos + odd x indices)\n    x_cols = [f'x_{i}' for i in range(1, cfg.IMG_WIDTH, 2)]\n    fieldnames = ['oid_ypos'] + x_cols\n\n    model.eval()\n    model.to(cfg.DEVICE)\n    results = []\n    progress_bar = tqdm(test_loader, desc='Predicting', dynamic_ncols=True)\n    inference_start_time = time.time()\n\n    with torch.inference_mode():\n        for inputs_batch, oids_batch in progress_bar:\n            inputs_batch = inputs_batch.to(cfg.DEVICE, non_blocking=True).float()\n            current_batch_size = inputs_batch.shape[0]\n\n            # --- Test-Time Augmentation (TTA) ---\n            if cfg.USE_TTA:\n                tta_predictions = []\n                # Original\n                with torch.cuda.amp.autocast(enabled=(cfg.DEVICE == torch.device('cuda'))):\n                    outputs_original = model(inputs_batch)\n                tta_predictions.append(outputs_original)\n\n                # Flipped\n                inputs_flipped = torch.flip(inputs_batch, dims=[3]).clone()\n                with torch.cuda.amp.autocast(enabled=(cfg.DEVICE == torch.device('cuda'))):\n                    outputs_flipped = model(inputs_flipped)\n                outputs_flipped_restored = torch.flip(outputs_flipped, dims=[3]).clone()\n                tta_predictions.append(outputs_flipped_restored)\n\n                # Average TTA predictions\n                final_outputs = torch.mean(torch.stack(tta_predictions), dim=0)\n                del outputs_original, inputs_flipped, outputs_flipped, outputs_flipped_restored, tta_predictions # Cleanup\n            else: # No TTA\n                 with torch.cuda.amp.autocast(enabled=(cfg.DEVICE == torch.device('cuda'))):\n                      final_outputs = model(inputs_batch)\n\n            # --- Process Outputs ---\n            y_preds_batch_np = final_outputs.squeeze(1).cpu().numpy() # (Batch, H, W)\n\n            for i in range(current_batch_size):\n                oid = oids_batch[i]\n                y_pred_single_np = y_preds_batch_np[i] # (H, W)\n                for y_pos in range(cfg.IMG_HEIGHT):\n                    oid_ypos = f\"{oid}_y_{y_pos}\"\n                    odd_x_values = y_pred_single_np[y_pos, 1::2] # Slice to get odd columns (1, 3, 5...)\n                    row_data = {'oid_ypos': oid_ypos}\n                    row_data.update(dict(zip(x_cols, odd_x_values)))\n                    results.append(row_data)\n\n            del inputs_batch, final_outputs, y_preds_batch_np\n            if cfg.DEVICE == torch.device('cuda'): torch.cuda.empty_cache()\n\n    inference_duration = time.time() - inference_start_time\n    print(f\"Inference finished in {inference_duration:.2f} seconds ({len(results)} rows generated).\")\n\n    # --- Write CSV ---\n    if not results:\n        print(\"Warning: No results were generated.\")\n        return\n    print(f\"Writing {len(results)} rows to submission file: {cfg.SUBMISSION_FILE}\")\n    submission_df = pd.DataFrame(results)\n    submission_df = submission_df[['oid_ypos'] + x_cols] # Ensure correct column order\n    try:\n        submission_df.to_csv(cfg.SUBMISSION_FILE, index=False, float_format='%.4f') # Format floats\n        print(f\"✅ Submission file created successfully!\")\n    except Exception as e:\n        print(f\"Error writing submission CSV: {e}\")\n\nprint(\"Prediction and submission functions defined.\")\n\n\n# ======================================================\n# 8. Main Execution Block\n# ======================================================\ndef main():\n    \"\"\"Main function to orchestrate the entire pipeline.\"\"\"\n    print(\"\\n\" + \"=\"*40)\n    print(\"===== FWI MLP Baseline Pipeline Start =====\")\n    print(\"=\"*40 + \"\\n\")\n\n    # --- Print Key Config ---\n    print(\"--- Configuration Summary ---\")\n    print(f\"Device: {CFG.DEVICE}, Seed: {CFG.SEED}\")\n    print(f\"Epochs: {CFG.N_EPOCHS}, Batch Size: {CFG.BATCH_SIZE}, LR: {CFG.LEARNING_RATE}\")\n    print(f\"Output Dir: {CFG.OUTPUT_DIR}\")\n    print(f\"Augmentation: {CFG.AUGMENT_TRAIN}, TTA: {CFG.USE_TTA}\\n\")\n\n    # Create output directory if it doesn't exist\n    Path(CFG.OUTPUT_DIR).mkdir(parents=True, exist_ok=True)\n\n    # --- Data Preparation ---\n    print(\"--- 1. Preparing Data ---\")\n    try:\n        all_input_files = find_data_files(CFG.TRAIN_DIR)\n        if not all_input_files: raise FileNotFoundError(\"No training input files found.\")\n        all_input_files, all_output_files = map_inputs_to_outputs(all_input_files)\n        if not all_input_files: raise ValueError(\"No valid input/output pairs found.\")\n        train_inputs, train_outputs, valid_inputs, valid_outputs = split_train_validation(\n            all_input_files, all_output_files, CFG.VALIDATION_SPLIT_RATIO\n        )\n        # Create Datasets\n        train_dataset = SeismicDataset(train_inputs, train_outputs, augment=CFG.AUGMENT_TRAIN, cfg=CFG)\n        valid_dataset = SeismicDataset(valid_inputs, valid_outputs, augment=False, cfg=CFG)\n        # Create DataLoaders\n        train_loader = DataLoader(\n            train_dataset, batch_size=CFG.BATCH_SIZE, shuffle=True,\n            num_workers=CFG.NUM_WORKERS, pin_memory=True, drop_last=True,\n            persistent_workers=(CFG.NUM_WORKERS > 0)\n        )\n        valid_loader = DataLoader(\n            valid_dataset, batch_size=max(1, CFG.BATCH_SIZE * 2), shuffle=False, # Often larger valid BS possible\n            num_workers=CFG.NUM_WORKERS, pin_memory=True, drop_last=False,\n            persistent_workers=(CFG.NUM_WORKERS > 0)\n        )\n        print(f\"DataLoaders ready: Train batches={len(train_loader)}, Valid batches={len(valid_loader)}\")\n    except Exception as e:\n        print(f\"FATAL ERROR during Data Preparation: {e}\")\n        return # Stop execution if data fails\n\n    # --- Model Initialization ---\n    print(\"\\n--- 2. Initializing Model ---\")\n    try:\n        model = FWINet(cfg=CFG)\n        num_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n        print(f\"Model '{model.__class__.__name__}' initialized with {num_params:,} trainable parameters.\")\n    except Exception as e:\n        print(f\"FATAL ERROR during Model Initialization: {e}\")\n        return # Stop execution if model fails\n\n    # --- Training ---\n    print(\"\\n--- 3. Starting Training ---\")\n    try:\n        history, trained_model = run_training(model, train_loader, valid_loader, cfg=CFG)\n        # 'trained_model' holds the model with the best weights loaded\n        print(\"Training complete.\")\n    except Exception as e:\n        print(f\"FATAL ERROR during Training: {e}\")\n        # Optionally try to generate submission with model state before error? For now, stop.\n        return\n\n    # --- Prediction & Submission ---\n    print(\"\\n--- 4. Generating Submission ---\")\n    try:\n        test_files = list(Path(CFG.TEST_DIR).glob('*.npy'))\n        if not test_files:\n            print(\"No test files found. Submission cannot be generated.\")\n        else:\n            print(f\"Found {len(test_files)} test files in {CFG.TEST_DIR}\")\n            # Make sure the best model state is actually in the model instance\n            # It should be loaded by run_training if successful.\n            generate_submission(trained_model, test_files, cfg=CFG)\n    except Exception as e:\n        print(f\"ERROR during Submission Generation: {e}\")\n        # Training might have finished, but submission failed.\n\n    print(\"\\n\" + \"=\"*40)\n    print(\"===== FWI MLP Baseline Pipeline End =====\")\n    print(\"=\"*40 + \"\\n\")\n\n\n# --- Entry Point ---\nif __name__ == \"__main__\":\n    # Record overall script execution time\n    script_start_time = time.time()\n    main() # Run the main pipeline\n    script_end_time = time.time()\n    total_duration = script_end_time - script_start_time\n    print(f\"\\nTotal script execution time: {total_duration / 60:.2f} minutes ({total_duration:.2f} seconds).\")","metadata":{"_uuid":"83959588-5bce-4664-93b6-e2c9486041a6","_cell_guid":"fb713b5c-b5e9-4f49-9716-da2b00c1f553","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-04-10T20:35:24.061621Z","iopub.execute_input":"2025-04-10T20:35:24.061810Z"}},"outputs":[],"execution_count":null}]}