{"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":"gpu","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport glob\nimport re\nimport math\nfrom typing import List, Dict, Tuple\nfrom collections import defaultdict\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as F\nfrom torchmetrics.functional import structural_similarity_index_measure\nfrom sklearn.metrics import mean_squared_error, mean_absolute_error\nfrom torch.cuda.amp import autocast, GradScaler\n# !pip install torch-optimizer\n# import torch_optimizer as optim\n# !pip install torchinfo\nfrom torchinfo import summary\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport random\nfrom tqdm.auto import tqdm\nfrom pathlib import Path\nimport time\nimport warnings\nimport copy\nfrom copy import deepcopy","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:34.412339Z","iopub.execute_input":"2025-06-20T17:16:34.412639Z","iopub.status.idle":"2025-06-20T17:16:39.585041Z","shell.execute_reply.started":"2025-06-20T17:16:34.412618Z","shell.execute_reply":"2025-06-20T17:16:39.584445Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Suppress warnings\nwarnings.filterwarnings('ignore')\n\n# Set plotting style\nplt.style.use('seaborn-v0_8')\nsns.set_palette(\"husl\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.586277Z","iopub.execute_input":"2025-06-20T17:16:39.586734Z","iopub.status.idle":"2025-06-20T17:16:39.591229Z","shell.execute_reply.started":"2025-06-20T17:16:39.586714Z","shell.execute_reply":"2025-06-20T17:16:39.590554Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Config:\n    BASE_PATH = Path(\"/kaggle/input/waveform-inversion\")\n    TRAIN_PATH = BASE_PATH / \"train_samples\"\n    TEST_PATH = BASE_PATH / \"test\"\n\n    # Data dimensions\n    SEISMIC_NUM_SOURCES = 5\n    SEISMIC_TIME_STEPS = 1000\n    SEISMIC_NUM_RECEIVERS = 70\n    VELOCITY_MAP_HEIGHT = 70\n    VELOCITY_MAP_WIDTH = 70\n    POS_ENC_CHANNELS = 4\n\n    # Model params\n    UNET_INPUT_CHANNELS = SEISMIC_NUM_SOURCES * 2 + POS_ENC_CHANNELS\n    UNET_OUTPUT_CHANNELS = 1\n    BASE_CHANNELS = 96\n    BILINEAR = True\n    USE_TTA = True\n    USE_EMA = True\n    USE_SCSE = True\n\n    # Training params\n    BATCH_SIZE = 32\n    LEARNING_RATE = 1e-4\n    NUM_EPOCHS = 100 # Increased epochs since EarlyStopping is used\n    VALIDATION_SPLIT = 0.15\n    RANDOM_SEED = 42\n    PATIENCE = 10  # Increased patience for early stopping\n    WEIGHT_DECAY = 1e-3\n\n    # Loss params\n    # ALPHA = 0.85\n    # BETA  = 0.15\n    W_MAE = 0.7\n    W_SSIM = 0.15\n    W_GRAD = 0.15\n\n    \n    # Normalization (to be computed)\n    SEISMIC_MEAN = None\n    SEISMIC_STD = None\n    VELOCITY_MEAN = None\n    VELOCITY_STD = None\n    GROUP_STATS = None\n\n    # Augmentation/Transformation\n    MAX_SEISMIC_SHIFT = 50\n    LOG_TRANSFORM_VELOCITY = True\n\n    # Augmentation parameters\n    NOISE_STD = 0.015  # Relative noise std for Gaussian noise\n    RECEIVER_DROP_PROB = 0.1\n    MAX_RECEIVER_DROPS = 2\n    SCALE_MIN = 0.9\n    SCALE_MAX = 1.1\n    TRANSLATE_PIXELS = 2\n\n    # Dropout for residual blocks\n    RES_BLOCK_DROPOUT = 0.1\n\n    # Sampling for stats\n    NORMALIZATION_SAMPLE_FRACTION = 1\n    CACHE_SIZE = 50\n    NUM_WORKERS = 4\n    PREFETCH_FACTOR = 4\n    PIN_MEMORY = True\n    PERSISTENT_WORKERS = True\n\n    DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\ncfg = Config()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.591988Z","iopub.execute_input":"2025-06-20T17:16:39.592228Z","iopub.status.idle":"2025-06-20T17:16:39.611664Z","shell.execute_reply.started":"2025-06-20T17:16:39.592212Z","shell.execute_reply":"2025-06-20T17:16:39.611109Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Reproducibility\ntorch.manual_seed(cfg.RANDOM_SEED)\nnp.random.seed(cfg.RANDOM_SEED)\nrandom.seed(cfg.RANDOM_SEED)\nif torch.cuda.is_available():\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\ndef seed_worker(worker_id):\n    worker_seed = cfg.RANDOM_SEED + worker_id\n    np.random.seed(worker_seed)\n    random.seed(worker_seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.613079Z","iopub.execute_input":"2025-06-20T17:16:39.613305Z","iopub.status.idle":"2025-06-20T17:16:39.629473Z","shell.execute_reply.started":"2025-06-20T17:16:39.613290Z","shell.execute_reply":"2025-06-20T17:16:39.628956Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def scan_files(root_dir: Path) -> list[dict[str, Path | str]]:\n    if not root_dir.exists():\n        raise FileNotFoundError(f\"Directory {root_dir} does not exist.\")\n\n    file_pairs: list[dict[str, Path | str]] = []\n    for item in sorted(root_dir.iterdir()):\n        if not item.is_dir():\n            continue\n\n        name = item.name\n        if (\"Vel\" in name) or (\"Style\" in name):\n            group = \"Vel\" if \"Vel\" in name else \"Style\"\n            data_dir = item / \"data\"\n            model_dir = item / \"model\"\n            if data_dir.exists() and model_dir.exists():\n                for data_file in sorted(data_dir.glob(\"data*.npy\")):\n                    idx_match = re.search(r\"data(\\d+)\\.npy\", data_file.name)\n                    if idx_match:\n                        idx = idx_match.group(1)\n                        model_file = model_dir / f\"model{idx}.npy\"\n                        if model_file.exists():\n                            file_pairs.append({\n                                \"input\": data_file,\n                                \"target\": model_file,\n                                \"group\": group\n                            })\n        elif \"Fault\" in name:\n            for seis_file in sorted(item.glob(\"seis*.npy\")):\n                base_name = seis_file.name.replace(\"seis\", \"vel\")\n                vel_file = item / base_name\n                if vel_file.exists():\n                    file_pairs.append({\n                        \"input\": seis_file,\n                        \"target\": vel_file,\n                        \"group\": \"Fault\"\n                    })\n    return file_pairs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.630156Z","iopub.execute_input":"2025-06-20T17:16:39.630341Z","iopub.status.idle":"2025-06-20T17:16:39.645720Z","shell.execute_reply.started":"2025-06-20T17:16:39.630317Z","shell.execute_reply":"2025-06-20T17:16:39.645108Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_stratified_split(\n    pairs: List[Dict[str, Path]],\n    val_frac: float = 0.15\n) -> Tuple[List[Dict[str, Path]], List[Dict[str, Path]]]:\n    \"\"\"\n    Stratify by directory type so each family (Vel/Style/Fault)\n    appears in both train and val.\n    Returns full dictionaries with 'input', 'target', 'group'.\n    \"\"\"\n    groups = {\"Vel\": [], \"Style\": [], \"Fault\": []}\n    for p in pairs:\n        group = p.get(\"group\")\n        if group in groups:\n            groups[group].append(p)\n        else:\n            print(f\"Warning: Unknown group type in {p}\")\n\n    train, val = [], []\n    for group_name, items in groups.items():\n        random.shuffle(items)\n        n_val = max(1, int(len(items) * val_frac))\n        val.extend(items[:n_val])\n        train.extend(items[n_val:])\n\n    print(f\"Stratified split → Train: {len(train)}, Val: {len(val)}\")\n    return train, val","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.646387Z","iopub.execute_input":"2025-06-20T17:16:39.646636Z","iopub.status.idle":"2025-06-20T17:16:39.660021Z","shell.execute_reply.started":"2025-06-20T17:16:39.646619Z","shell.execute_reply":"2025-06-20T17:16:39.659320Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compute_stratified_stats(file_pairs):\n    group_stats = {}\n\n    for group in ['Vel', 'Style', 'Fault']:\n        group_files = [(p['input'], p['target']) for p in file_pairs if p['group'] == group]\n        if not group_files:\n            continue\n\n        ds = StatsDataset(group_files, log_transform_velocity=cfg.LOG_TRANSFORM_VELOCITY)\n        loader = DataLoader(ds, batch_size=cfg.BATCH_SIZE, num_workers=cfg.NUM_WORKERS)\n\n        s_mean, s_std, v_mean, v_std = compute_stats_gpu(loader, sample_fraction=1.0)\n        group_stats[group] = {\n            \"seismic_mean\": s_mean,\n            \"seismic_std\": s_std,\n            \"vel_mean\": v_mean,\n            \"vel_std\": v_std,\n        }\n        print(f\"[{group}] SEISMIC μ={s_mean:.2f}, σ={s_std:.2f} | VELOCITY μ={v_mean:.2f}, σ={v_std:.2f}\")\n\n    return group_stats","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.660824Z","iopub.execute_input":"2025-06-20T17:16:39.661067Z","iopub.status.idle":"2025-06-20T17:16:39.676341Z","shell.execute_reply.started":"2025-06-20T17:16:39.661046Z","shell.execute_reply":"2025-06-20T17:16:39.675721Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class StatsDataset(Dataset):\n    \"\"\"\n    Dataset to compute normalization statistics on the final 14-channel processed data\n    (10 FFT channels + 4 Positional Encoding channels).\n    \"\"\"\n    def __init__(self, file_paths, log_transform_velocity=False, cache_in_memory=False):\n        self.file_metadata = []\n        self.log_transform_velocity = log_transform_velocity\n        self.cache_in_memory = cache_in_memory\n        self.seismic_cache = {}\n        self.velocity_cache = {}\n\n        # Pre-generate the positional encoding once, as it's the same for all samples\n        self.pos_encoding = generate_positional_encoding(\n            cfg.VELOCITY_MAP_HEIGHT,\n            cfg.VELOCITY_MAP_WIDTH,\n            cfg.POS_ENC_CHANNELS\n        )\n\n        for seismic_path, vel_path in file_paths:\n            if self.cache_in_memory:\n                self.seismic_cache[seismic_path] = np.load(seismic_path)\n                self.velocity_cache[vel_path] = np.load(vel_path)\n\n            # Correctly load the array to get its shape without a 'with' statement\n            data = np.load(seismic_path)\n            num_samples = data.shape[0]\n\n            self.file_metadata.append((seismic_path, vel_path, num_samples))\n\n        self.indices = []\n        for file_idx, (_, _, num_samples) in enumerate(self.file_metadata):\n            self.indices.extend([(file_idx, i) for i in range(num_samples)])\n\n    def __len__(self):\n        return len(self.indices)\n\n    def __getitem__(self, idx):\n        file_idx, sample_idx = self.indices[idx]\n        seismic_path, vel_path, _ = self.file_metadata[file_idx]\n\n        if self.cache_in_memory:\n            seismic_data = self.seismic_cache[seismic_path]\n            vel_data = self.velocity_cache[vel_path]\n        else:\n            seismic_data = np.load(seismic_path, mmap_mode='r')\n            vel_data = np.load(vel_path, mmap_mode='r')\n\n        seismic = torch.from_numpy(seismic_data[sample_idx].copy()).float()\n        vel = torch.from_numpy(vel_data[sample_idx].copy()).float()\n\n        # 1. Perform FFT and separate into Real/Imaginary parts\n        ffted = torch.fft.fft(seismic, dim=1)\n        real_part = ffted.real[:, :cfg.SEISMIC_NUM_RECEIVERS, :]\n        imag_part = ffted.imag[:, :cfg.SEISMIC_NUM_RECEIVERS, :]\n        seismic_fft = torch.cat([real_part, imag_part], dim=0)\n\n        # 2. Concatenate with positional encoding to create the final 14-channel tensor\n        seismic = torch.cat([seismic_fft, self.pos_encoding], dim=0)\n\n        # Process velocity map\n        if self.log_transform_velocity:\n            vel = torch.log(vel + 1e-6)\n\n        if vel.dim() == 2:\n            vel = vel.unsqueeze(0)\n\n        return seismic, vel","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.677094Z","iopub.execute_input":"2025-06-20T17:16:39.677843Z","iopub.status.idle":"2025-06-20T17:16:39.692100Z","shell.execute_reply.started":"2025-06-20T17:16:39.677818Z","shell.execute_reply":"2025-06-20T17:16:39.691482Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RunningStats:\n    \"\"\"Online mean & std via Welford's algorithm.\"\"\"\n    def __init__(self, device='cuda'):\n        self.n = torch.tensor(0, dtype=torch.long, device=device)\n        self.mean = torch.tensor(0.0, device=device)\n        self.M2 = torch.tensor(0.0, device=device)\n\n    def update(self, x):\n        with torch.no_grad():\n            flat = x.detach().flatten()\n            batch_n = flat.numel()\n            batch_mean = flat.mean()\n            batch_M2 = flat.var(unbiased=False) * batch_n\n\n            delta = batch_mean - self.mean\n            total_n = self.n + batch_n\n\n            self.mean = (self.n * self.mean + batch_n * batch_mean) / total_n\n            self.M2 = self.M2 + batch_M2 + delta**2 * self.n * batch_n / total_n\n            self.n = total_n\n\n    @property\n    def std(self):\n        return torch.sqrt(self.M2 / self.n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.692794Z","iopub.execute_input":"2025-06-20T17:16:39.692955Z","iopub.status.idle":"2025-06-20T17:16:39.711064Z","shell.execute_reply.started":"2025-06-20T17:16:39.692942Z","shell.execute_reply":"2025-06-20T17:16:39.710393Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compute_stats_gpu(dataloader, sample_fraction=0.1):\n    \"\"\"Compute mean/std on GPU for seismic and velocity.\"\"\"\n    seismic_stats = RunningStats(device=cfg.DEVICE)\n    velocity_stats = RunningStats(device=cfg.DEVICE)\n\n    n_samples = int(len(dataloader.dataset) * sample_fraction)\n    count = 0\n\n    with torch.no_grad():\n        for seismic, velocity in dataloader:\n            seismic = seismic.to(cfg.DEVICE, non_blocking=True)\n            velocity = velocity.to(cfg.DEVICE, non_blocking=True)\n\n            seismic_stats.update(seismic)\n            velocity_stats.update(velocity)\n\n            count += seismic.shape[0]\n            if count >= n_samples:\n                break\n\n    return (\n        seismic_stats.mean.item(), seismic_stats.std.item(),\n        velocity_stats.mean.item(), velocity_stats.std.item()\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.713303Z","iopub.execute_input":"2025-06-20T17:16:39.713501Z","iopub.status.idle":"2025-06-20T17:16:39.727969Z","shell.execute_reply.started":"2025-06-20T17:16:39.713487Z","shell.execute_reply":"2025-06-20T17:16:39.727279Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class WaveformInversionDataset(Dataset):\n    \"\"\"\n    Main dataset for training and validation. It creates the 14-channel input,\n    applies group-aware normalization, and performs data augmentation.\n    \"\"\"\n    def __init__(self, file_pairs, group_stats,\n                 augment=False, light_augment=False,\n                 log_transform_velocity=False, cache_in_memory=False):\n        self.pairs = file_pairs\n        self.group_stats = group_stats\n        self.augment = augment\n        self.light_augment = light_augment\n        self.log_transform_velocity = log_transform_velocity\n        self.cache_in_memory = cache_in_memory\n        self.seismic_cache = {}\n        self.velocity_cache = {}\n\n        # Pre-generate the positional encoding once\n        self.pos_encoding = generate_positional_encoding(\n            cfg.VELOCITY_MAP_HEIGHT,\n            cfg.VELOCITY_MAP_WIDTH,\n            cfg.POS_ENC_CHANNELS\n        )\n\n        self.metadata = []\n        for pair in self.pairs:\n            seismic_path = pair[\"input\"]\n            vel_path = pair[\"target\"]\n            group = pair[\"group\"]\n\n            if cache_in_memory:\n                self.seismic_cache[seismic_path] = np.load(seismic_path)\n                self.velocity_cache[vel_path] = np.load(vel_path)\n            \n            ### --- CORRECTED CODE --- ###\n            # The 'with' statement is removed to fix the TypeError.\n            data = np.load(seismic_path)\n            num_samples = data.shape[0]\n            ### --- END CORRECTION --- ###\n\n            self.metadata.append((seismic_path, vel_path, num_samples, group))\n\n        self.indices = []\n        for i, (_, _, num_samples, _) in enumerate(self.metadata):\n            self.indices.extend([(i, j) for j in range(num_samples)])\n\n    def __len__(self):\n        return len(self.indices)\n\n    def __getitem__(self, idx):\n        meta_idx, sample_idx = self.indices[idx]\n        seismic_path, vel_path, _, group = self.metadata[meta_idx]\n        stats = self.group_stats[group]\n\n        if self.cache_in_memory:\n            seismic_data = self.seismic_cache[seismic_path]\n            vel_data = self.velocity_cache[vel_path]\n        else:\n            seismic_data = np.load(seismic_path, mmap_mode='r')\n            vel_data = np.load(vel_path, mmap_mode='r')\n\n        seismic = torch.from_numpy(seismic_data[sample_idx].copy()).float()\n        vel = torch.from_numpy(vel_data[sample_idx].copy()).float()\n\n        if self.log_transform_velocity:\n            vel = torch.log(vel + 1e-6)\n\n        # 1. Perform FFT and separate into Real/Imaginary parts\n        ffted = torch.fft.fft(seismic, dim=1)\n        real_part = ffted.real[:, :cfg.SEISMIC_NUM_RECEIVERS, :]\n        imag_part = ffted.imag[:, :cfg.SEISMIC_NUM_RECEIVERS, :]\n        seismic_fft = torch.cat([real_part, imag_part], dim=0)\n\n        # 2. Concatenate with positional encoding\n        seismic_processed = torch.cat([seismic_fft, self.pos_encoding], dim=0)\n\n        # 3. Apply normalization to the final 14-channel tensor\n        seismic = (seismic_processed - stats[\"seismic_mean\"]) / stats[\"seismic_std\"]\n\n        # 4. Apply data augmentation (if specified)\n        if self.augment:\n            seismic += torch.randn_like(seismic) * cfg.NOISE_STD\n            scale = random.uniform(cfg.SCALE_MIN, cfg.SCALE_MAX)\n            seismic *= scale\n            if random.random() < cfg.RECEIVER_DROP_PROB:\n                for _ in range(random.randint(1, cfg.MAX_RECEIVER_DROPS)):\n                    # Only drop FFT channels, not positional encoding channels\n                    ch = random.randint(0, 10 - 1) \n                    seismic[ch, :, :] = 0.0\n\n        elif self.light_augment:\n            seismic += torch.randn_like(seismic) * (cfg.NOISE_STD * 0.5)\n            seismic *= random.uniform(0.98, 1.02)\n\n        if vel.dim() == 2:\n            vel = vel.unsqueeze(0)\n\n        vel = (vel - stats[\"vel_mean\"]) / stats[\"vel_std\"]\n        \n        return seismic, vel, group","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.728532Z","iopub.execute_input":"2025-06-20T17:16:39.728788Z","iopub.status.idle":"2025-06-20T17:16:39.748077Z","shell.execute_reply.started":"2025-06-20T17:16:39.728768Z","shell.execute_reply":"2025-06-20T17:16:39.747466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def generate_positional_encoding(height, width, channels):\n    \"\"\"\n    Generates a 2D positional encoding map of shape (channels, height, width).\n    \"\"\"\n    if channels % 4 != 0:\n        raise ValueError(\"Cannot use sin/cos positional encoding with \"\n                         \"odd number of channels or channels not divisible by 4.\")\n\n    pos_encoding = torch.zeros(channels, height, width)\n    \n    # Create a grid of coordinates\n    y_pos, x_pos = torch.meshgrid(torch.arange(height), torch.arange(width), indexing=\"ij\")\n\n    # Define the division term for the sine/cosine frequencies\n    div_term = torch.exp(torch.arange(0, channels // 2, 2) * -(math.log(10000.0) / (channels // 2)))\n\n    # Calculate positional encoding for x coordinate\n    pos_encoding[0::2, :, :] = torch.sin(x_pos.unsqueeze(0) * div_term.view(-1, 1, 1))\n    pos_encoding[1::2, :, :] = torch.cos(x_pos.unsqueeze(0) * div_term.view(-1, 1, 1))\n    \n    # Calculate positional encoding for y coordinate\n    # Note: We apply this to the second half of the channels\n    y_channel_offset = channels // 2\n    pos_encoding[y_channel_offset::2, :, :] = torch.sin(y_pos.unsqueeze(0) * div_term.view(-1, 1, 1))\n    pos_encoding[y_channel_offset+1::2, :, :] = torch.cos(y_pos.unsqueeze(0) * div_term.view(-1, 1, 1))\n\n    return pos_encoding","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.748685Z","iopub.execute_input":"2025-06-20T17:16:39.748921Z","iopub.status.idle":"2025-06-20T17:16:39.767238Z","shell.execute_reply.started":"2025-06-20T17:16:39.748899Z","shell.execute_reply":"2025-06-20T17:16:39.766620Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ResidualDoubleConv(nn.Module):\n    \"\"\"Residual double‐conv block with dropout.\"\"\"\n    def __init__(self, in_ch, out_ch, dropout=cfg.RES_BLOCK_DROPOUT):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1, bias=False)\n        self.norm1 = nn.InstanceNorm2d(out_ch)\n        self.relu1 = nn.ReLU()\n        self.conv2 = nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1, bias=False)\n        self.norm2 = nn.InstanceNorm2d(out_ch)\n        self.relu2 = nn.ReLU()\n        self.skip = nn.Conv2d(in_ch, out_ch, kernel_size=1) if in_ch != out_ch else nn.Identity()\n        self.dropout = nn.Dropout2d(p=dropout) if dropout > 0 else nn.Identity()\n\n    def forward(self, x):\n        identity = self.skip(x)\n        out = self.relu1(self.norm1(self.conv1(x)))\n        out = self.norm2(self.conv2(out))\n        out = self.dropout(out)\n        return self.relu2(out + identity)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.767814Z","iopub.execute_input":"2025-06-20T17:16:39.767994Z","iopub.status.idle":"2025-06-20T17:16:39.784611Z","shell.execute_reply.started":"2025-06-20T17:16:39.767979Z","shell.execute_reply":"2025-06-20T17:16:39.784018Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SCSEBlock(nn.Module):\n    \"\"\"Concurrent Spatial and Channel Squeeze & Excitation (SCSE) block.\"\"\"\n    def __init__(self, channels, reduction=16):\n        super().__init__()\n        self.cSE = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(channels, channels // reduction, 1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(channels // reduction, channels, 1),\n            nn.Sigmoid()\n        )\n        self.sSE = nn.Sequential(\n            nn.Conv2d(channels, 1, 1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        return x * self.cSE(x) + x * self.sSE(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.785332Z","iopub.execute_input":"2025-06-20T17:16:39.785828Z","iopub.status.idle":"2025-06-20T17:16:39.801321Z","shell.execute_reply.started":"2025-06-20T17:16:39.785812Z","shell.execute_reply":"2025-06-20T17:16:39.800641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Down(nn.Module):\n    \"\"\"Downscaling with maxpool then residual double‐conv.\"\"\"\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.layer = nn.Sequential(\n            nn.MaxPool2d(2),\n            ResidualDoubleConv(in_ch, out_ch, dropout=cfg.RES_BLOCK_DROPOUT)\n        )\n\n    def forward(self, x):\n        return self.layer(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.802092Z","iopub.execute_input":"2025-06-20T17:16:39.802294Z","iopub.status.idle":"2025-06-20T17:16:39.817364Z","shell.execute_reply.started":"2025-06-20T17:16:39.802273Z","shell.execute_reply":"2025-06-20T17:16:39.816831Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AttentionBlock(nn.Module):\n    \"\"\"Attention block used in Attention U-Net.\"\"\"\n    def __init__(self, F_g, F_l, F_int):\n        super().__init__()\n        self.W_g = nn.Sequential(\n            nn.Conv2d(F_g, F_int, kernel_size=1, bias=True),\n            nn.BatchNorm2d(F_int)\n        )\n        self.W_x = nn.Sequential(\n            nn.Conv2d(F_l, F_int, kernel_size=1, bias=True),\n            nn.BatchNorm2d(F_int)\n        )\n        self.psi = nn.Sequential(\n            nn.ReLU(),\n            nn.Conv2d(F_int, 1, kernel_size=1, bias=True),\n            nn.BatchNorm2d(1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, g, x):\n        g1 = self.W_g(g)\n        x1 = self.W_x(x)\n        psi = self.psi(torch.relu(g1 + x1))\n        return x * psi","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.817984Z","iopub.execute_input":"2025-06-20T17:16:39.818179Z","iopub.status.idle":"2025-06-20T17:16:39.832481Z","shell.execute_reply.started":"2025-06-20T17:16:39.818165Z","shell.execute_reply":"2025-06-20T17:16:39.831897Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Up(nn.Module):\n    \"\"\"Upscaling then attention then residual double‐conv.\"\"\"\n    def __init__(self, in_ch, out_ch, bilinear=True):\n        super().__init__()\n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        else:\n            self.up = nn.ConvTranspose2d(in_ch // 2, in_ch // 2, kernel_size=2, stride=2)\n\n        self.attn = AttentionBlock(F_g=in_ch // 2, F_l=in_ch // 2, F_int=in_ch // 4)\n        self.conv = ResidualDoubleConv(in_ch, out_ch, dropout=cfg.RES_BLOCK_DROPOUT)\n        self.scse = SCSEBlock(out_ch)\n\n    def forward(self, x, skip):\n        x = self.up(x)\n\n        # If spatial sizes don't match, apply padding\n        diffY = skip.size(2) - x.size(2)\n        diffX = skip.size(3) - x.size(3)\n\n        if diffY != 0 or diffX != 0:\n            assert abs(diffY) <= 2 and abs(diffX) <= 2, (\n                f\"Padding too large: x={x.shape}, skip={skip.shape}, \"\n                f\"diffY={diffY}, diffX={diffX}\"\n            )\n\n            x = F.pad(x, [diffX // 2, diffX - diffX // 2,\n                          diffY // 2, diffY - diffY // 2])\n\n        # Attention and concat\n        attn_out = self.attn(g=x, x=skip)\n        x = torch.cat([attn_out, x], dim=1)\n        return self.scse(self.conv(x))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.833122Z","iopub.execute_input":"2025-06-20T17:16:39.833319Z","iopub.status.idle":"2025-06-20T17:16:39.845558Z","shell.execute_reply.started":"2025-06-20T17:16:39.833305Z","shell.execute_reply":"2025-06-20T17:16:39.845043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class OutConv(nn.Module):\n    \"\"\"Final 1×1 convolution to produce output channels.\"\"\"\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.conv = nn.Conv2d(in_ch, out_ch, kernel_size=1)\n\n    def forward(self, x):\n        return self.conv(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.846278Z","iopub.execute_input":"2025-06-20T17:16:39.846596Z","iopub.status.idle":"2025-06-20T17:16:39.863973Z","shell.execute_reply.started":"2025-06-20T17:16:39.846553Z","shell.execute_reply":"2025-06-20T17:16:39.863192Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNet(nn.Module):\n    \"\"\"\n    Attention U-Net with residual double-conv blocks, modified for multi-scale outputs.\n    \"\"\"\n    def __init__(self, n_channels, n_classes, bilinear=cfg.BILINEAR, base_channels=cfg.BASE_CHANNELS):\n        super().__init__()\n        B = base_channels\n        self.inc = ResidualDoubleConv(n_channels, B, dropout=cfg.RES_BLOCK_DROPOUT)\n        self.down1 = Down(B, B*2)\n        self.down2 = Down(B*2, B*4)\n        self.down3 = Down(B*4, B*8)\n        factor = 2 if bilinear else 1\n        self.down4 = Down(B*8, B*16 // factor)\n\n        self.up1 = Up(B*16, B*8 // factor, bilinear)\n        self.up2 = Up(B*8, B*4 // factor, bilinear)\n        self.up3 = Up(B*4, B*2 // factor, bilinear)\n        self.up4 = Up(B*2, B, bilinear)\n        \n        # The final, full-resolution output convolution\n        self.outc_final = OutConv(B, n_classes) ### <-- MODIFIED (renamed for clarity)\n        \n        ### --- NEW: Prediction heads for intermediate scales --- ###\n        # Prediction head after up2 (1/4 resolution)\n        self.outc_scale2 = OutConv(B*4 // factor, n_classes)\n        # Prediction head after up3 (1/2 resolution)\n        self.outc_scale1 = OutConv(B*2 // factor, n_classes)\n        ### --- END NEW --- ###\n\n\n    def forward(self, x):\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n\n        # Decoder path\n        u1 = self.up1(x5, x4)\n        u2 = self.up2(u1, x3)\n        u3 = self.up3(u2, x2)\n        u4 = self.up4(u3, x1)\n\n        # Final full-resolution prediction\n        logits_final = self.outc_final(u4)\n        \n        ### --- NEW: Generate predictions from intermediate layers --- ###\n        logits_scale1 = self.outc_scale1(u3) # Prediction at 1/2 resolution\n        logits_scale2 = self.outc_scale2(u2) # Prediction at 1/4 resolution\n        ### --- END NEW --- ###\n\n        # Return a list of predictions, from lowest resolution to highest\n        return [logits_scale2, logits_scale1, logits_final]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.864748Z","iopub.execute_input":"2025-06-20T17:16:39.864963Z","iopub.status.idle":"2025-06-20T17:16:39.879335Z","shell.execute_reply.started":"2025-06-20T17:16:39.864947Z","shell.execute_reply":"2025-06-20T17:16:39.878737Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_sobel_filters(device):\n    \"\"\"Creates Sobel filters for X and Y gradients, moved to the specified device.\"\"\"\n    # Sobel filter for the x-gradient\n    sobel_x = torch.tensor([[-1, 0, 1], [-2, 0, 2], [-1, 0, 1]], dtype=torch.float32).view(1, 1, 3, 3)\n    \n    # Sobel filter for the y-gradient\n    sobel_y = torch.tensor([[-1, -2, -1], [0, 0, 0], [1, 2, 1]], dtype=torch.float32).view(1, 1, 3, 3)\n    \n    # Move filters to the same device as the model\n    return sobel_x.to(device), sobel_y.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.880153Z","iopub.execute_input":"2025-06-20T17:16:39.880313Z","iopub.status.idle":"2025-06-20T17:16:39.897731Z","shell.execute_reply.started":"2025-06-20T17:16:39.880301Z","shell.execute_reply":"2025-06-20T17:16:39.897072Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def gradient_loss(pred, target, sobel_x, sobel_y):\n    \"\"\"\n    Calculates the L1 loss between the gradients of the prediction and the target.\n    \"\"\"\n    # Calculate gradients for the prediction\n    pred_grad_x = F.conv2d(pred, sobel_x, padding=1)\n    pred_grad_y = F.conv2d(pred, sobel_y, padding=1)\n    \n    # Calculate gradients for the target\n    target_grad_x = F.conv2d(target, sobel_x, padding=1)\n    target_grad_y = F.conv2d(target, sobel_y, padding=1)\n    \n    # Calculate L1 loss on the gradients\n    loss = F.l1_loss(pred_grad_x, target_grad_x) + F.l1_loss(pred_grad_y, target_grad_y)\n    \n    return loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.898413Z","iopub.execute_input":"2025-06-20T17:16:39.898643Z","iopub.status.idle":"2025-06-20T17:16:39.918541Z","shell.execute_reply.started":"2025-06-20T17:16:39.898625Z","shell.execute_reply":"2025-06-20T17:16:39.918001Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rescale_to_unit_range(x):\n    \"\"\"\n    Safely rescales a tensor to the [0, 1] range for stable SSIM calculation.\n    \"\"\"\n    min_val = x.amin(dim=(1, 2, 3), keepdim=True)\n    max_val = x.amax(dim=(1, 2, 3), keepdim=True)\n    # The epsilon in the denominator prevents division by zero for flat images\n    return (x - min_val) / (max_val - min_val + 1e-8)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.919348Z","iopub.execute_input":"2025-06-20T17:16:39.919610Z","iopub.status.idle":"2025-06-20T17:16:39.934092Z","shell.execute_reply.started":"2025-06-20T17:16:39.919588Z","shell.execute_reply":"2025-06-20T17:16:39.933497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def safe_ssim_loss(pred, target):\n    \"\"\"\n    Calculates a numerically stable SSIM loss.\n    Returns the loss value (1.0 - SSIM).\n    \"\"\"\n    # Rescale both tensors independently to the [0, 1] range\n    pred_norm = rescale_to_unit_range(pred)\n    target_norm = rescale_to_unit_range(target)\n    \n    # Calculate SSIM on the rescaled tensors with a fixed data_range of 1.0\n    ssim_val = structural_similarity_index_measure(pred_norm, target_norm, data_range=1.0)\n    \n    # The SSIM loss is 1 - ssim\n    ssim_term = 1.0 - ssim_val\n    \n    return ssim_term","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.934712Z","iopub.execute_input":"2025-06-20T17:16:39.935053Z","iopub.status.idle":"2025-06-20T17:16:39.948804Z","shell.execute_reply.started":"2025-06-20T17:16:39.935032Z","shell.execute_reply":"2025-06-20T17:16:39.948074Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def safe_ssim_metric(pred, target):\n    \"\"\"\n    Calculates a numerically stable SSIM score (not the loss).\n    \"\"\"\n    # Use the same rescaling helper as our loss function\n    pred_norm = rescale_to_unit_range(pred)\n    target_norm = rescale_to_unit_range(target)\n    \n    # Calculate and return the SSIM score on the rescaled tensors\n    return structural_similarity_index_measure(pred_norm, target_norm, data_range=1.0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.949618Z","iopub.execute_input":"2025-06-20T17:16:39.949849Z","iopub.status.idle":"2025-06-20T17:16:39.965918Z","shell.execute_reply.started":"2025-06-20T17:16:39.949825Z","shell.execute_reply":"2025-06-20T17:16:39.965218Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class CombinedLoss(nn.Module):\n#     \"\"\"\n#     A combined loss function that includes MAE, a stable SSIM, and a gradient (Sobel) loss.\n#     Loss = w_mae * MAE + w_ssim * SSIM_Loss + w_grad * Gradient_Loss\n#     \"\"\"\n#     def __init__(self, w_mae=cfg.W_MAE, w_ssim=cfg.W_SSIM, w_grad=cfg.W_GRAD, device=cfg.DEVICE):\n#         super().__init__()\n#         self.w_mae = w_mae\n#         self.w_ssim = w_ssim\n#         self.w_grad = w_grad\n        \n#         # Create and store the sobel filters on the correct device\n#         self.sobel_x, self.sobel_y = create_sobel_filters(device)\n\n#     def forward(self, pred, target):\n#         # 1. Mean Absolute Error (L1 Loss)\n#         mae_term = F.l1_loss(pred, target)\n        \n#         # 2. Stable SSIM Loss (using our safe helper function)\n#         ssim_term = safe_ssim_loss(pred, target)\n        \n#         # 3. Gradient Loss (using our sobel helper function)\n#         grad_term = gradient_loss(pred, target, self.sobel_x, self.sobel_y)\n        \n#         # Combine the losses with their respective weights\n#         total_loss = (self.w_mae * mae_term) + \\\n#                      (self.w_ssim * ssim_term) + \\\n#                      (self.w_grad * grad_term)\n        \n#         return total_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.966761Z","iopub.execute_input":"2025-06-20T17:16:39.967008Z","iopub.status.idle":"2025-06-20T17:16:39.980144Z","shell.execute_reply.started":"2025-06-20T17:16:39.966982Z","shell.execute_reply":"2025-06-20T17:16:39.979621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CombinedLoss(nn.Module):\n    \"\"\"\n    A combined loss function adapted for multi-scale supervision.\n    \"\"\"\n    def __init__(self, w_mae=cfg.W_MAE, w_ssim=cfg.W_SSIM, w_grad=cfg.W_GRAD, device=cfg.DEVICE):\n        super().__init__()\n        self.w_mae = w_mae\n        self.w_ssim = w_ssim\n        self.w_grad = w_grad\n        self.sobel_x, self.sobel_y = create_sobel_filters(device)\n\n    def forward(self, preds, target): ### <-- MODIFIED: 'pred' is now 'preds' (a list)\n        \n        total_loss = 0\n        \n        # Iterate through the predictions from each scale\n        for p in preds:\n            # Downsample the ground truth target to match the prediction's size\n            # mode='area' is generally good for downsampling.\n            resized_target = F.interpolate(target, size=p.shape[2:], mode='area')\n            \n            # --- Calculate the combined loss for the current scale ---\n            mae_term = F.l1_loss(p, resized_target)\n            ssim_term = safe_ssim_loss(p, resized_target)\n            grad_term = gradient_loss(p, resized_target, self.sobel_x, self.sobel_y)\n            \n            scale_loss = (self.w_mae * mae_term) + \\\n                         (self.w_ssim * ssim_term) + \\\n                         (self.w_grad * grad_term)\n            \n            # Add the loss for the current scale to the total\n            total_loss += scale_loss\n            \n        return total_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:39.980728Z","iopub.execute_input":"2025-06-20T17:16:39.980922Z","iopub.status.idle":"2025-06-20T17:16:40.000448Z","shell.execute_reply.started":"2025-06-20T17:16:39.980908Z","shell.execute_reply":"2025-06-20T17:16:39.999945Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class MAE_SSIM_Loss(nn.Module):\n#     \"\"\"\n#     Combines L1 (MAE) with SSIM as a composite loss.\n#     Total loss = alpha * MAE + beta * (1 - SSIM)\n#     \"\"\"\n#     def __init__(self, alpha=cfg.ALPHA, beta=cfg.BETA):\n#         super().__init__()\n#         self.alpha = alpha\n#         self.beta = beta\n\n#     def forward(self, pred, target):\n#         mae_loss = F.l1_loss(pred, target)\n#         # For SSIM, we expect the input to be in a predictable range, e.g. [0, 1] or [-1, 1]\n#         # Since our targets are normalized around 0, we can use a large data_range.\n#         # Alternatively, clamp to a range if we know it. Here, we assume a reasonable range after normalization.\n#         data_range = target.max() - target.min()\n#         if data_range < 1e-6: # Handle case of blank targets\n#             data_range = 1.0\n\n#         ssim_val = structural_similarity_index_measure(pred, target, data_range=data_range)\n#         ssim_loss = 1.0 - ssim_val\n#         return self.alpha * mae_loss + self.beta * ssim_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.001071Z","iopub.execute_input":"2025-06-20T17:16:40.001239Z","iopub.status.idle":"2025-06-20T17:16:40.020154Z","shell.execute_reply.started":"2025-06-20T17:16:40.001226Z","shell.execute_reply":"2025-06-20T17:16:40.019626Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ModelEMA(nn.Module):\n    def __init__(self, model, decay=0.99, device=None):\n        super().__init__()\n        self.module = deepcopy(model)\n        self.module.eval()\n        self.decay = decay\n        self.device = device\n        if device:\n            self.module.to(device=device)\n\n    def _update(self, model, update_fn):\n        with torch.no_grad():\n            for ema_v, model_v in zip(self.module.state_dict().values(), model.state_dict().values()):\n                model_v = model_v.to(ema_v.device)\n                ema_v.copy_(update_fn(ema_v, model_v))\n\n    def update(self, model):\n        self._update(model, lambda e, m: self.decay * e + (1. - self.decay) * m)\n\n    def set(self, model):\n        self._update(model, lambda e, m: m)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.023389Z","iopub.execute_input":"2025-06-20T17:16:40.023633Z","iopub.status.idle":"2025-06-20T17:16:40.038443Z","shell.execute_reply.started":"2025-06-20T17:16:40.023614Z","shell.execute_reply":"2025-06-20T17:16:40.037839Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def denormalize_velocity(norm_vel, groups):\n#     \"\"\"\n#     Denormalize batched velocity using per-sample group stats.\n#     groups: list of strings (length = batch_size)\n#     \"\"\"\n#     out = []\n#     for i, group in enumerate(groups):\n#         stats = cfg.GROUP_STATS[group]\n#         v = norm_vel[i] * stats[\"vel_std\"] + stats[\"vel_mean\"]\n#         if cfg.LOG_TRANSFORM_VELOCITY:\n#             v = torch.exp(v)\n#         out.append(v.unsqueeze(0))  # keep batch dimension\n\n#     return torch.cat(out, dim=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.039216Z","iopub.execute_input":"2025-06-20T17:16:40.039480Z","iopub.status.idle":"2025-06-20T17:16:40.053975Z","shell.execute_reply.started":"2025-06-20T17:16:40.039459Z","shell.execute_reply":"2025-06-20T17:16:40.053424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def denormalize_velocity(norm_vel, groups):\n    \"\"\"\n    Denormalize batched velocity using per-sample group stats.\n    Includes a clamp for numerical stability.\n    \"\"\"\n    out = []\n    for i, group in enumerate(groups):\n        stats = cfg.GROUP_STATS[group]\n        v = norm_vel[i] * stats[\"vel_std\"] + stats[\"vel_mean\"]\n        if cfg.LOG_TRANSFORM_VELOCITY:\n            # Clamp before the exponential to prevent overflow to infinity\n            v = torch.clamp(v, max=20.0)\n            v = torch.exp(v)\n        out.append(v.unsqueeze(0))  # keep batch dimension\n\n    return torch.cat(out, dim=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.054562Z","iopub.execute_input":"2025-06-20T17:16:40.054764Z","iopub.status.idle":"2025-06-20T17:16:40.070881Z","shell.execute_reply.started":"2025-06-20T17:16:40.054750Z","shell.execute_reply":"2025-06-20T17:16:40.070275Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def denormalize_velocity_global(norm_vel):\n    \"\"\"Denormalize batched velocity using global stats for inference.\"\"\"\n    v = norm_vel * cfg.VELOCITY_STD + cfg.VELOCITY_MEAN\n    if cfg.LOG_TRANSFORM_VELOCITY:\n        v = torch.exp(v)\n    return v","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.071536Z","iopub.execute_input":"2025-06-20T17:16:40.071807Z","iopub.status.idle":"2025-06-20T17:16:40.086594Z","shell.execute_reply.started":"2025-06-20T17:16:40.071784Z","shell.execute_reply":"2025-06-20T17:16:40.085910Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_model_metrics(history, save_path='model_metrics.png'):\n    \"\"\"Plot MAE & SSIM curves, overfitting gap, and smoothed MAE.\"\"\"\n    epochs = np.arange(1, len(history['train_denorm_mae']) + 1)\n    fig, axes = plt.subplots(2, 2, figsize=(14, 10))\n\n    # 1) MAE curves\n    axes[0, 0].plot(epochs, history['train_denorm_mae'], 'b-', label='Train MAE', linewidth=2)\n    axes[0, 0].plot(epochs, history['val_denorm_mae'], 'r-', label='Val MAE', linewidth=2)\n    axes[0, 0].set_title('MAE Curves (Denormalized)', fontsize=14, fontweight='bold')\n    axes[0, 0].set_xlabel('Epoch')\n    axes[0, 0].set_ylabel('MAE')\n    axes[0, 0].legend()\n    axes[0, 0].grid(alpha=0.3)\n\n    # 2) SSIM curves\n    axes[0, 1].plot(epochs, history['train_ssim'], 'b-', label='Train SSIM', linewidth=2)\n    axes[0, 1].plot(epochs, history['val_ssim'], 'r-', label='Val SSIM', linewidth=2)\n    axes[0, 1].set_title('SSIM Curves', fontsize=14, fontweight='bold')\n    axes[0, 1].set_xlabel('Epoch')\n    axes[0, 1].set_ylabel('SSIM')\n    axes[0, 1].legend()\n    axes[0, 1].grid(alpha=0.3)\n\n    # 3) Overfitting gap (MAE difference)\n    mae_gap = np.array(history['val_denorm_mae']) - np.array(history['train_denorm_mae'])\n    axes[1, 0].plot(epochs, mae_gap, 'purple', linewidth=2)\n    axes[1, 0].axhline(0, color='black', linestyle='--', alpha=0.5)\n    axes[1, 0].set_title('Overfitting Gap (Val MAE – Train MAE)', fontsize=14, fontweight='bold')\n    axes[1, 0].set_xlabel('Epoch')\n    axes[1, 0].set_ylabel('MAE Gap')\n    axes[1, 0].grid(alpha=0.3)\n\n    # 4) Smoothed MAE\n    window = max(3, len(epochs) // 10)\n    if len(epochs) >= window:\n        smooth_train = np.convolve(history['train_denorm_mae'], np.ones(window) / window, mode='valid')\n        smooth_val = np.convolve(history['val_denorm_mae'], np.ones(window) / window, mode='valid')\n        smooth_epochs = epochs[window - 1:]\n        axes[1, 1].plot(epochs, history['train_denorm_mae'], 'b-', alpha=0.3, label='Train MAE', linewidth=1)\n        axes[1, 1].plot(epochs, history['val_denorm_mae'], 'r-', alpha=0.3, label='Val MAE', linewidth=1)\n        axes[1, 1].plot(smooth_epochs, smooth_train, 'b-', linewidth=2, label=f'Train MA({window})')\n        axes[1, 1].plot(smooth_epochs, smooth_val, 'r-', linewidth=2, label=f'Val MA({window})')\n        axes[1, 1].set_title('Smoothed MAE Learning Curves', fontsize=14, fontweight='bold')\n        axes[1, 1].set_xlabel('Epoch')\n        axes[1, 1].set_ylabel('MAE')\n        axes[1, 1].legend()\n        axes[1, 1].grid(alpha=0.3)\n    else:\n        axes[1, 1].text(0.5, 0.5, \"Not enough epochs to smooth\", ha='center', va='center', fontsize=12)\n        axes[1, 1].axis('off')\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.087322Z","iopub.execute_input":"2025-06-20T17:16:40.087657Z","iopub.status.idle":"2025-06-20T17:16:40.103088Z","shell.execute_reply.started":"2025-06-20T17:16:40.087631Z","shell.execute_reply":"2025-06-20T17:16:40.102345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_scatter_and_residuals(model, val_loader, save_path='scatter_residuals.png'):\n    \"\"\"\n    Plots predicted vs actual values, residuals, and error distribution.\n    Updated to handle multi-scale model outputs.\n    \"\"\"\n    model.eval()\n    all_preds = []\n    all_targets = []\n\n    with torch.no_grad():\n        for inputs, targets, groups in tqdm(val_loader, desc=\"Gathering eval data\"):\n            inputs = inputs.to(cfg.DEVICE)\n            targets = targets.to(cfg.DEVICE)\n            \n            # model(inputs) returns a list of predictions for different scales\n            outputs = model(inputs)\n\n            ### --- CORRECTED CODE --- ###\n            # 1. Select the final, full-resolution prediction from the list.\n            final_pred = outputs[-1]\n\n            # 2. Pass this single tensor to the denormalization function.\n            pred_denorm = denormalize_velocity(final_pred, groups).cpu().numpy().flatten()\n            ### --- END CORRECTION --- ###\n            \n            tgt_denorm = denormalize_velocity(targets, groups).cpu().numpy().flatten()\n\n            all_preds.append(pred_denorm)\n            all_targets.append(tgt_denorm)\n\n    all_preds = np.concatenate(all_preds)\n    all_targets = np.concatenate(all_targets)\n    residuals = all_preds - all_targets\n\n    mse = mean_squared_error(all_targets, all_preds)\n    mae = mean_absolute_error(all_targets, all_preds)\n    rmse = np.sqrt(mse)\n    ss_res = np.sum((all_targets - all_preds) ** 2)\n    ss_tot = np.sum((all_targets - all_targets.mean()) ** 2)\n    r2 = 1 - (ss_res / ss_tot) if ss_tot > 0 else 0\n\n    fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\n    # Scatter: Pred vs Actual\n    axes[0].scatter(all_targets, all_preds, alpha=0.3, s=1)\n    axes[0].plot([all_targets.min(), all_targets.max()],\n                 [all_targets.min(), all_targets.max()],\n                 'r--', linewidth=2)\n    axes[0].set_title(f'Pred vs Actual\\nR²={r2:.4f}', fontsize=14, fontweight='bold')\n    axes[0].set_xlabel('Actual Velocity')\n    axes[0].set_ylabel('Predicted Velocity')\n    axes[0].grid(alpha=0.3)\n\n    # Residuals\n    axes[1].scatter(all_targets, residuals, alpha=0.3, s=1)\n    axes[1].axhline(0, color='red', linestyle='--', linewidth=1.5)\n    axes[1].set_title(f'Residuals\\nMAE={mae:.4f}, RMSE={rmse:.4f}', fontsize=14, fontweight='bold')\n    axes[1].set_xlabel('Actual Velocity')\n    axes[1].set_ylabel('Residual (Pred – Actual)')\n    axes[1].grid(alpha=0.3)\n\n    # Error distribution\n    axes[2].hist(residuals, bins=50, alpha=0.7, edgecolor='black')\n    axes[2].axvline(0, color='red', linestyle='--', linewidth=1.5)\n    axes[2].set_title('Error Distribution', fontsize=14, fontweight='bold')\n    axes[2].set_xlabel('Residual')\n    axes[2].set_ylabel('Frequency')\n    axes[2].grid(alpha=0.3)\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')\n    plt.show()\n\n    return {'mse': mse, 'mae': mae, 'rmse': rmse, 'r2': r2}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.103843Z","iopub.execute_input":"2025-06-20T17:16:40.104159Z","iopub.status.idle":"2025-06-20T17:16:40.121747Z","shell.execute_reply.started":"2025-06-20T17:16:40.104136Z","shell.execute_reply":"2025-06-20T17:16:40.121115Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_predictions(model, loader, group_stats, num_samples=4):\n    \"\"\"\n    Visualizes model predictions with group-aware denormalization.\n    Updated to handle multi-scale model outputs.\n    \"\"\"\n    model.eval()\n    count = 0\n\n    with torch.no_grad():\n        for inputs, targets, groups in loader:\n            inputs = inputs.to(cfg.DEVICE)\n            targets = targets.to(cfg.DEVICE)\n            \n            # model(inputs) returns a list of predictions\n            outputs = model(inputs)\n\n            ### --- CORRECTED CODE --- ###\n            # 1. Select the final, full-resolution prediction from the list.\n            final_pred = outputs[-1]\n\n            # 2. Pass this single tensor to the denormalization function.\n            denorm_pred = denormalize_velocity(final_pred, groups)\n            ### --- END CORRECTION --- ###\n            \n            denorm_target = denormalize_velocity(targets, groups)\n\n            for i in range(inputs.size(0)):\n                # Take first source as representative seismic input\n                # This shows the Real part of the FFT for the first source\n                input_img = inputs[i, 0].cpu().numpy() \n                pred_img = denorm_pred[i].squeeze().cpu().numpy()\n                target_img = denorm_target[i].squeeze().cpu().numpy()\n                group = groups[i]\n\n                # Compute MAE and SSIM for this single sample\n                pred_tensor = denorm_pred[i].unsqueeze(0)\n                target_tensor = denorm_target[i].unsqueeze(0)\n                mae = F.l1_loss(pred_tensor, target_tensor).item()\n                \n                # Use our stable SSIM metric function\n                ssim_score = safe_ssim_metric(pred_tensor, target_tensor).item()\n\n                # Plot\n                fig, axs = plt.subplots(1, 3, figsize=(12, 4))\n                axs[0].imshow(input_img, cmap='gray')\n                axs[0].set_title(\"Seismic Input (Source 1, Real Part)\")\n\n                axs[1].imshow(pred_img, cmap='viridis')\n                axs[1].set_title(f\"Prediction\\nMAE: {mae:.1f} | SSIM: {ssim_score:.3f}\")\n\n                axs[2].imshow(target_img, cmap='viridis')\n                axs[2].set_title(f\"Ground Truth\\nGroup: {group}\")\n\n                for ax in axs:\n                    ax.axis(\"off\")\n\n                plt.tight_layout()\n                plt.show()\n\n                count += 1\n                if count >= num_samples:\n                    return","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.122324Z","iopub.execute_input":"2025-06-20T17:16:40.122635Z","iopub.status.idle":"2025-06-20T17:16:40.142410Z","shell.execute_reply.started":"2025-06-20T17:16:40.122611Z","shell.execute_reply":"2025-06-20T17:16:40.141951Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EarlyStopping:\n    \"\"\"\n    Early stops if validation denormalized MAE doesn't improve after `patience` epochs.\n    Saves the best model.\n    \"\"\"\n    def __init__(self, patience: int = cfg.PATIENCE, min_delta: float = 0.0, path: str = 'best_model.pth', verbose: bool = False):\n        self.patience = patience\n        self.min_delta = min_delta\n        self.path = path\n        self.verbose = verbose\n        self.counter = 0\n        self.best_score = None\n        self.early_stop = False\n\n    def __call__(self, val_denorm_mae: float, model_to_save: nn.Module):\n        if self.best_score is None:\n            self.best_score = val_denorm_mae\n            self._save_checkpoint(model_to_save)\n        elif val_denorm_mae < self.best_score - self.min_delta:\n            self.best_score = val_denorm_mae\n            self._save_checkpoint(model_to_save)\n            self.counter = 0\n        else:\n            self.counter += 1\n            if self.verbose:\n                print(f\"EarlyStopping counter: {self.counter} out of {self.patience}\")\n            if self.counter >= self.patience:\n                self.early_stop = True\n\n    def _save_checkpoint(self, model_to_save: nn.Module):\n        torch.save(model_to_save.state_dict(), self.path)\n        if self.verbose:\n            print(f\"Validation MAE improved to {self.best_score:.4f}. Saving model to {self.path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.143497Z","iopub.execute_input":"2025-06-20T17:16:40.143810Z","iopub.status.idle":"2025-06-20T17:16:40.161671Z","shell.execute_reply.started":"2025-06-20T17:16:40.143789Z","shell.execute_reply":"2025-06-20T17:16:40.161057Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(model, train_loader, val_loader, optimizer, scheduler, num_epochs):\n    \"\"\"\n    The main training and validation loop, updated with all enhancements.\n    \"\"\"\n    history = {\n        \"train_loss\": [], \"val_loss\": [],\n        \"train_denorm_mae\": [], \"val_denorm_mae\": [],\n        \"train_ssim\": [], \"val_ssim\": [],\n        \"val_mae_group\": {}, \"val_ssim_group\": {},\n        \"ema_val_denorm_mae\": [], \"ema_val_ssim\": [],\n        \"ema_val_mae_group\": {}, \"ema_val_ssim_group\": {}\n    }\n\n    # Use our new sophisticated and stable loss function\n    criterion = CombinedLoss()\n    \n    early_stopper = EarlyStopping(path='best_model.pth', verbose=True)\n    ema = ModelEMA(model, decay=0.99, device=cfg.DEVICE)\n    scaler = GradScaler()\n\n    def validate(model_to_eval, loader, label):\n        \"\"\"Inner function to perform validation.\"\"\"\n        model_to_eval.eval()\n        loss_total, denorm_mae_total, ssim_total = 0, 0, 0\n        group_mae, group_ssim = defaultdict(list), defaultdict(list)\n\n        with torch.no_grad():\n            for inputs, targets, groups in loader:\n                inputs = inputs.to(cfg.DEVICE)\n                targets = targets.to(cfg.DEVICE)\n\n                with autocast():\n                    outputs = model_to_eval(inputs)\n                    loss = criterion(outputs, targets)\n\n                loss_total += loss.item()\n\n                # Use the final, full-resolution prediction for metrics\n                final_pred = outputs[-1]\n\n                denorm_out = denormalize_velocity(final_pred, groups)\n                denorm_tar = denormalize_velocity(targets, groups)\n\n                for i, group in enumerate(groups):\n                    group_mae[group].append(F.l1_loss(denorm_out[i], denorm_tar[i]).item())\n                    \n                    # Use our numerically stable SSIM metric function to prevent NaNs\n                    ssim_score = safe_ssim_metric(denorm_out[i].unsqueeze(0), denorm_tar[i].unsqueeze(0))\n                    group_ssim[group].append(ssim_score.item())\n                    \n        avg_group_mae = {g: np.mean(v) for g, v in group_mae.items()}\n        avg_group_ssim = {g: np.mean(v) for g, v in group_ssim.items()}\n        denorm_mae = np.mean([mae for v in group_mae.values() for mae in v])\n        ssim = np.mean([s for v in group_ssim.values() for s in v])\n        avg_loss = loss_total / len(loader)\n\n        print(f\"[{label}] Loss: {avg_loss:.4f} | MAE: {denorm_mae:.2f} | SSIM: {ssim:.3f}\")\n        return avg_loss, denorm_mae, ssim, avg_group_mae, avg_group_ssim\n\n    # --- Main Training Loop ---\n    for epoch in range(num_epochs):\n        model.train()\n        train_loss, train_mae, train_ssim = 0, 0, 0\n\n        pbar = tqdm(train_loader, desc=f\"[Epoch {epoch+1}/{num_epochs}] Training\")\n        for inputs, targets, groups in pbar:\n            inputs = inputs.to(cfg.DEVICE)\n            targets = targets.to(cfg.DEVICE)\n\n            optimizer.zero_grad()\n\n            # Forward pass with mixed-precision\n            with autocast():\n                outputs = model(inputs)\n                loss = criterion(outputs, targets)\n\n            # Backward pass with gradient scaling\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n            ema.update(model)\n            train_loss += loss.item()\n\n            # Calculate training metrics for logging\n            with torch.no_grad():\n                # Use the final, full-resolution prediction for metrics\n                final_pred = outputs[-1]\n                denorm_out = denormalize_velocity(final_pred, groups)\n                denorm_tar = denormalize_velocity(targets, groups)\n                train_mae += F.l1_loss(denorm_out, denorm_tar).item()\n                \n                # Use our numerically stable SSIM metric function to prevent NaNs\n                train_ssim += safe_ssim_metric(denorm_out, denorm_tar).item()\n\n        # --- End of Epoch: Calculate and Print Metrics ---\n        train_loss /= len(train_loader)\n        train_mae /= len(train_loader)\n        train_ssim /= len(train_loader)\n\n        # Print training metrics\n        print(f\"[TRAIN] Loss: {train_loss:.4f} | MAE: {train_mae:.2f} | SSIM: {train_ssim:.3f}\")\n\n        # Perform validation\n        val_loss, val_mae, val_ssim, val_grp_mae, val_grp_ssim = validate(model, val_loader, \"VAL (Raw)\")\n        ema_loss, ema_mae, ema_ssim, ema_grp_mae, ema_grp_ssim = validate(ema.module, val_loader, \"VAL (EMA)\")\n\n        # Step the scheduler based on the validation MAE\n        scheduler.step(ema_mae) \n\n        # Log history\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"train_denorm_mae\"].append(train_mae)\n        history[\"val_denorm_mae\"].append(val_mae)\n        history[\"train_ssim\"].append(train_ssim)\n        history[\"val_ssim\"].append(val_ssim)\n        history[\"val_mae_group\"][epoch] = val_grp_mae\n        history[\"val_ssim_group\"][epoch] = val_grp_ssim\n        history[\"ema_val_denorm_mae\"].append(ema_mae)\n        history[\"ema_val_ssim\"].append(ema_ssim)\n        history[\"ema_val_mae_group\"][epoch] = ema_grp_mae\n        history[\"ema_val_ssim_group\"][epoch] = ema_grp_ssim\n\n        # Check for early stopping based on EMA model's performance\n        early_stopper(ema_mae, ema.module)\n        if early_stopper.early_stop:\n            print(\"Early stopping triggered.\")\n            break\n\n    print(f\"Training complete. Best EMA MAE: {early_stopper.best_score:.4f}\")\n    return history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.162393Z","iopub.execute_input":"2025-06-20T17:16:40.162651Z","iopub.status.idle":"2025-06-20T17:16:40.184407Z","shell.execute_reply.started":"2025-06-20T17:16:40.162630Z","shell.execute_reply":"2025-06-20T17:16:40.183729Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_seismic_input(val_loader, num_samples=3, save_path='seismic_input.png'):\n    \"\"\"Visualize the processed seismic input data (FFT magnitude)\"\"\"\n    with torch.no_grad():\n        for inputs, targets, _ in val_loader:\n            break\n\n    indices = np.random.choice(len(inputs), min(num_samples, len(inputs)), replace=False)\n\n    fig, axes = plt.subplots(cfg.SEISMIC_NUM_SOURCES, num_samples, figsize=(6 * num_samples, 15))\n    if num_samples == 1:\n        axes = axes.reshape(-1, 1)\n\n    for i, idx in enumerate(indices):\n        seismic = inputs[idx].cpu().numpy()  # Shape: (5, 70, 70)\n\n        for source in range(cfg.SEISMIC_NUM_SOURCES):\n            im = axes[source, i].imshow(seismic[source], cmap='viridis', aspect='auto')\n            axes[source, i].set_title(f'Sample {idx+1}, Source {source+1}\\n(FFT Magnitude)',\n                                    fontsize=12, fontweight='bold')\n            axes[source, i].set_xlabel('Receiver Position')\n            axes[source, i].set_ylabel('Frequency Bin')\n            plt.colorbar(im, ax=axes[source, i], shrink=0.8)\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.185127Z","iopub.execute_input":"2025-06-20T17:16:40.185342Z","iopub.status.idle":"2025-06-20T17:16:40.204048Z","shell.execute_reply.started":"2025-06-20T17:16:40.185324Z","shell.execute_reply":"2025-06-20T17:16:40.203466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_groupwise_val_mae(history, title=\"Validation MAE by Group\"):\n    group_names = set()\n    for epoch_data in history[\"val_mae_group\"].values():\n        group_names.update(epoch_data.keys())\n\n    group_names = sorted(group_names)\n    epochs = sorted(history[\"val_mae_group\"].keys())\n\n    for group in group_names:\n        y = [history[\"val_mae_group\"][epoch].get(group, float('nan')) for epoch in epochs]\n        plt.plot(range(1, len(y) + 1), y, label=group)\n\n    plt.title(title)\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"MAE\")\n    plt.grid(True)\n    plt.legend()\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.204728Z","iopub.execute_input":"2025-06-20T17:16:40.205142Z","iopub.status.idle":"2025-06-20T17:16:40.224563Z","shell.execute_reply.started":"2025-06-20T17:16:40.205125Z","shell.execute_reply":"2025-06-20T17:16:40.223896Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_groupwise_val_ssim(history, title=\"Validation SSIM by Group\"):\n    group_names = set()\n    for epoch_data in history[\"val_ssim_group\"].values():\n        group_names.update(epoch_data.keys())\n\n    group_names = sorted(group_names)\n    epochs = sorted(history[\"val_ssim_group\"].keys())\n\n    for group in group_names:\n        y = [history[\"val_ssim_group\"][epoch].get(group, float('nan')) for epoch in epochs]\n        plt.plot(range(1, len(y) + 1), y, label=group)\n\n    plt.title(title)\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"SSIM\")\n    plt.grid(True)\n    plt.legend()\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.225303Z","iopub.execute_input":"2025-06-20T17:16:40.225462Z","iopub.status.idle":"2025-06-20T17:16:40.240162Z","shell.execute_reply.started":"2025-06-20T17:16:40.225450Z","shell.execute_reply":"2025-06-20T17:16:40.239613Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def scan_test_files(test_dir: Path) -> List[Path]:\n    \"\"\"Scan test directory for .npy files (test samples).\"\"\"\n    if not test_dir.exists():\n        raise FileNotFoundError(f\"Test directory {test_dir} not found.\")\n    test_files = sorted(test_dir.glob(\"*.npy\"))\n    print(f\"Found {len(test_files)} test samples.\")\n    return test_files","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.240855Z","iopub.execute_input":"2025-06-20T17:16:40.241052Z","iopub.status.idle":"2025-06-20T17:16:40.258559Z","shell.execute_reply.started":"2025-06-20T17:16:40.241030Z","shell.execute_reply":"2025-06-20T17:16:40.257841Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TestDataset(Dataset):\n    \"\"\"\n    Loads a list of .npy test files and applies the full 14-channel preprocessing\n    (FFT Real/Imag + Positional Encoding) and global normalization.\n    \"\"\"\n    def __init__(self, file_paths: List[Path]):\n        self.file_paths = file_paths\n        # Pre-generate the positional encoding once\n        self.pos_encoding = generate_positional_encoding(\n            cfg.VELOCITY_MAP_HEIGHT,\n            cfg.VELOCITY_MAP_WIDTH,\n            cfg.POS_ENC_CHANNELS\n        )\n\n    def __len__(self):\n        return len(self.file_paths)\n\n    def __getitem__(self, idx):\n        path = self.file_paths[idx]\n        data = np.load(path)\n        tensor = torch.from_numpy(data).float()\n\n        # 1. Perform FFT and separate into Real/Imaginary parts\n        ffted = torch.fft.fft(tensor, dim=1)\n        real_part = ffted.real[:, :cfg.SEISMIC_NUM_RECEIVERS, :]\n        imag_part = ffted.imag[:, :cfg.SEISMIC_NUM_RECEIVERS, :]\n        seismic_fft = torch.cat([real_part, imag_part], dim=0)\n\n        # 2. Concatenate with positional encoding\n        seismic_processed = torch.cat([seismic_fft, self.pos_encoding], dim=0)\n\n        # 3. Normalize with the pre-calculated GLOBAL stats\n        norm_seismic = (seismic_processed - cfg.SEISMIC_MEAN) / cfg.SEISMIC_STD\n        \n        return norm_seismic, path.stem","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.259326Z","iopub.execute_input":"2025-06-20T17:16:40.259507Z","iopub.status.idle":"2025-06-20T17:16:40.273776Z","shell.execute_reply.started":"2025-06-20T17:16:40.259493Z","shell.execute_reply":"2025-06-20T17:16:40.273016Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_test_inference(model, test_files: List[Path], batch_size: int = 8):\n    \"\"\"Batched inference over TestDataset, updated for multi-scale model.\"\"\"\n    ds = TestDataset(test_files)\n    loader = DataLoader(\n        ds,\n        batch_size=batch_size,\n        shuffle=False,\n        num_workers=cfg.NUM_WORKERS,\n        pin_memory=cfg.PIN_MEMORY,\n    )\n\n    model.eval()\n    all_preds = []\n    all_oids  = []\n    with torch.no_grad():\n        for batch, stems in tqdm(loader, desc=\"Test inference\"):\n            batch = batch.to(cfg.DEVICE)\n            \n            # out is a list of predictions\n            out = model(batch)\n            \n            ### --- CORRECTED CODE --- ###\n            # Select the final, full-resolution prediction for submission\n            final_out = out[-1]\n            \n            # Denormalize the final prediction\n            den = denormalize_velocity_global(final_out).squeeze(1).cpu().numpy()\n            ### --- END CORRECTION --- ###\n            \n            all_preds.append(den)\n            all_oids.extend(stems)\n\n    preds = np.concatenate(all_preds, axis=0)\n    return preds, all_oids","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.274482Z","iopub.execute_input":"2025-06-20T17:16:40.274677Z","iopub.status.idle":"2025-06-20T17:16:40.293079Z","shell.execute_reply.started":"2025-06-20T17:16:40.274663Z","shell.execute_reply":"2025-06-20T17:16:40.292449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def make_submission(preds: np.ndarray, oids: List[str], filename=\"submission.csv\"):\n#     \"\"\"\n#     Build submission DataFrame: for each oid, for each y in [0..69],\n#     output only odd x‐columns (1,3,5,…).\n#     \"\"\"\n#     rows = []\n#     x_cols = list(range(1, cfg.VELOCITY_MAP_WIDTH, 2))  # [1,3,...,69]\n#     for oid, pred in zip(oids, preds):\n#         for y in range(pred.shape[0]):\n#             row_id = f\"{oid}y{y}\"\n#             vals   = pred[y, x_cols].tolist()\n#             rows.append([row_id] + vals)\n\n#     cols = [\"oid_ypos\"] + [f\"x{i}\" for i in x_cols]\n#     df   = pd.DataFrame(rows, columns=cols)\n#     df.to_csv(filename, index=False)\n#     print(f\"Saved submission to {filename}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.293724Z","iopub.execute_input":"2025-06-20T17:16:40.293926Z","iopub.status.idle":"2025-06-20T17:16:40.311749Z","shell.execute_reply.started":"2025-06-20T17:16:40.293911Z","shell.execute_reply":"2025-06-20T17:16:40.311027Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_submission(preds: np.ndarray, oids: List[str], filename=\"submission.csv\"):\n    \"\"\"\n    Builds the submission DataFrame with the corrected header and row ID format.\n    \"\"\"\n    rows = []\n    x_cols = list(range(1, cfg.VELOCITY_MAP_WIDTH, 2))\n\n    for oid, pred in zip(oids, preds):\n        for y in range(pred.shape[0]):\n            \n            ### --- FIX 1: Correct the format of the values --- ###\n            row_id = f\"{oid}_y_{y}\"\n\n            vals   = pred[y, x_cols].tolist()\n            rows.append([row_id] + vals)\n\n    ### --- FIX 2: Correct the name of the header column --- ###\n    cols = [\"oid_ypos\"] + [f\"x{i}\" for i in x_cols]\n    \n    df   = pd.DataFrame(rows, columns=cols)\n    df.to_csv(filename, index=False)\n    print(f\"Saved submission to {filename}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.312462Z","iopub.execute_input":"2025-06-20T17:16:40.312702Z","iopub.status.idle":"2025-06-20T17:16:40.326094Z","shell.execute_reply.started":"2025-06-20T17:16:40.312686Z","shell.execute_reply":"2025-06-20T17:16:40.325407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    print(f\"\\n === Starting Group-Aware Waveform Inversion Training === \\n\")\n\n    # Load file pairs\n    file_pairs = scan_files(cfg.TRAIN_PATH)\n\n    # Stratified split\n    train_pairs, val_pairs = create_stratified_split(\n        file_pairs, val_frac=cfg.VALIDATION_SPLIT\n    )\n\n    # --- Compute Normalization Statistics ---\n    print(f\"\\n === Computing normalization statistics === \\n\")\n    # 1. Group-wise stats for training/validation\n    print(\"Computing group-wise stats...\")\n    cfg.GROUP_STATS = compute_stratified_stats(train_pairs)\n\n    # 2. Global stats for test set inference\n    print(\"\\nComputing global stats for test set...\")\n    train_file_tuples = [(p['input'], p['target']) for p in train_pairs]\n    stats_ds = StatsDataset(train_file_tuples, log_transform_velocity=cfg.LOG_TRANSFORM_VELOCITY)\n    stats_loader = DataLoader(\n        stats_ds, batch_size=cfg.BATCH_SIZE, num_workers=cfg.NUM_WORKERS,\n        pin_memory=cfg.PIN_MEMORY\n    )\n    (\n        cfg.SEISMIC_MEAN, cfg.SEISMIC_STD,\n        cfg.VELOCITY_MEAN, cfg.VELOCITY_STD\n    ) = compute_stats_gpu(stats_loader, sample_fraction=1.0)\n    print(f\"GLOBAL SEISMIC: μ={cfg.SEISMIC_MEAN:.4f}, σ={cfg.SEISMIC_STD:.4f}\")\n    print(f\"GLOBAL VELOCITY: μ={cfg.VELOCITY_MEAN:.4f}, σ={cfg.VELOCITY_STD:.4f}\")\n\n    # Create datasets using group-aware stats\n    train_dataset = WaveformInversionDataset(\n        train_pairs, group_stats=cfg.GROUP_STATS,\n        augment=True, cache_in_memory=False, log_transform_velocity=cfg.LOG_TRANSFORM_VELOCITY\n    )\n    val_dataset = WaveformInversionDataset(\n        val_pairs, group_stats=cfg.GROUP_STATS,\n        light_augment=True, cache_in_memory=False, log_transform_velocity=cfg.LOG_TRANSFORM_VELOCITY\n    )\n\n    # Dataloaders\n    g = torch.Generator()\n    g.manual_seed(cfg.RANDOM_SEED)\n    train_loader = DataLoader(\n        train_dataset, batch_size=cfg.BATCH_SIZE, shuffle=True,\n        num_workers=cfg.NUM_WORKERS, pin_memory=cfg.PIN_MEMORY,\n        persistent_workers=cfg.PERSISTENT_WORKERS, prefetch_factor=cfg.PREFETCH_FACTOR,\n        worker_init_fn=seed_worker, generator=g\n    )\n    val_loader = DataLoader(\n        val_dataset, batch_size=cfg.BATCH_SIZE, shuffle=False,\n        num_workers=cfg.NUM_WORKERS, pin_memory=cfg.PIN_MEMORY,\n        persistent_workers=cfg.PERSISTENT_WORKERS, prefetch_factor=cfg.PREFETCH_FACTOR\n    )\n\n    # Optional: visualize some preprocessed seismic\n    print(f\"\\n === Visualizing some preprocessed seismic waveform samples === \\n\")\n    visualize_seismic_input(val_loader, num_samples=2)\n\n    print(f\"\\n === Preparing to train the model === \\n\")\n\n    # Initialize model & optimizer\n    model = UNet(\n        n_channels=cfg.UNET_INPUT_CHANNELS,\n        n_classes=cfg.UNET_OUTPUT_CHANNELS,\n        bilinear=cfg.BILINEAR,\n        base_channels=cfg.BASE_CHANNELS\n    ).to(cfg.DEVICE)\n\n    summary(\n        model,\n        input_size=(cfg.BATCH_SIZE, cfg.UNET_INPUT_CHANNELS, 70, 70),\n        col_names=[\"input_size\", \"output_size\", \"num_params\", \"trainable\"],\n        depth=4,\n        verbose=1\n    )\n\n    # optimizer = optim.Lookahead(torch.optim.AdamW(model.parameters(), lr=cfg.LEARNING_RATE))\n    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.LEARNING_RATE, weight_decay=cfg.WEIGHT_DECAY)\n\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', factor=0.2, patience=3)\n\n\n    # Train the model\n    history = train_model(model, train_loader, val_loader, optimizer, scheduler, cfg.NUM_EPOCHS)\n\n    # --- Post-Training Analysis ---\n    print(f\"\\nTraining complete. Loading best model for evaluation.\\n\")\n    model.load_state_dict(torch.load(\"best_model.pth\"))\n    torch.save(model.state_dict(), \"final_model.pth\") # Save a final copy\n    print(f\"Best model loaded and saved to 'final_model.pth'.\\n\")\n\n    # Visualizations\n    print(f\"\\n === Preparing visualizations === \\n\")\n    plot_model_metrics(history, save_path='training_curves.png')\n    metrics = plot_scatter_and_residuals(model, val_loader, save_path='residuals.png')\n    plot_groupwise_val_mae(history)\n    plot_groupwise_val_ssim(history)\n    visualize_predictions(model, val_loader, cfg.GROUP_STATS, num_samples=3)\n\n    # --- Inference ---\n    print(f\"\\n === Starting Test Set Inference === \\n\")\n    test_files = scan_test_files(cfg.TEST_PATH)\n    preds, oids = run_test_inference(model, test_files, batch_size=cfg.BATCH_SIZE)\n\n    # Make submission file\n    make_submission(preds, oids, filename=\"submission.csv\")\n\n    return model, history, metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.326823Z","iopub.execute_input":"2025-06-20T17:16:40.327088Z","iopub.status.idle":"2025-06-20T17:16:40.349252Z","shell.execute_reply.started":"2025-06-20T17:16:40.327067Z","shell.execute_reply":"2025-06-20T17:16:40.348618Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    # To prevent issues in environments like Jupyter\n    try:\n        from numba import cuda\n        cuda.select_device(0)\n        cuda.close()\n    except:\n        pass\n\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-20T17:16:40.349947Z","iopub.execute_input":"2025-06-20T17:16:40.350488Z","execution_failed":"2025-06-20T17:19:32.032Z"}},"outputs":[],"execution_count":null}]}