{"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":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n\"\"\"\nScript for training a FWIGAN (WGAN-GP) model for Full Waveform Inversion,\nusing data sourced solely from Kaggle input directories, with data augmentation.\nIncludes placeholder normalization - **USER MUST UPDATE NORMALIZATION VALUES**.\n\"\"\"\n\n!pip install webdataset -q # Install quietly\n\n# %% Imports\nimport csv\nimport gc\nimport glob\nimport os\nimport random\nimport shutil\nimport sys\nfrom pathlib import Path\nimport time # For timing epochs\n\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.amp # For Automatic Mixed Precision\nimport torch.autograd as autograd # For gradient penalty\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms.functional as TF # For Augmentation\nimport webdataset as wds\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm.auto import tqdm\n\n# %% Configuration\nclass cfg:\n    \"\"\"Configuration parameters for the workflow.\"\"\"\n\n    # --- Paths ---\n    kaggle_train_dir = \"/kaggle/input/waveform-inversion/train_samples\"\n    kaggle_test_dir = \"/kaggle/input/waveform-inversion/test\"\n    shard_output_dir = \"/kaggle/working/sharded_data\"\n    working_dir = \"/kaggle/working/\"\n    submission_file = os.path.join(working_dir, \"submission.csv\")\n    model_save_dir = os.path.join(working_dir, \"models\") # Directory to save models\n\n    # --- Dataset Params ---\n    dataset_name = \"fwi_kaggle_only_augmented\"\n    # !!! IMPORTANT: Define your data normalization parameters !!!\n    # !!! Determine these from your *entire* training dataset !!!\n    VEL_MIN = 1400.0 # Placeholder min velocity\n    VEL_MAX = 4600.0 # Placeholder max velocity\n    SEIS_NORM_MODE = 'minmax_sample' # 'minmax_sample', 'std_sample', or None\n\n    # --- Sharding Params ---\n    maxsize = 1e9  # Approx 1 GB per shard\n    force_shard_creation = False\n\n    # --- Splitting & Loading Params ---\n    num_used_shards = None  # Use all available\n    test_size = 0.1  # Proportion for validation split\n    batch_size = 8 # GANs are memory intensive, start lower\n    num_workers = 2\n\n    # --- Augmentation Params ---\n    apply_augmentation = True\n    aug_hflip_prob = 0.5\n    # Noise added *before* normalization in this setup\n    aug_seis_noise_std = 0.01 # Std dev relative to original seismic range\n\n    # --- Model params (Generator - U-Net like) ---\n    unet_in_channels = 5\n    unet_out_channels = 1\n    unet_init_features = 32 # Initial features for Generator\n    unet_depth = 5\n    unet_bilinear = True\n\n    # --- Model params (Discriminator) ---\n    # Input channels = vel_map (1) + processed_seismic (5)\n    disc_in_channels = unet_out_channels + unet_in_channels\n    disc_init_features = 64 # Initial features for Discriminator\n\n    # --- Training params (WGAN-GP) ---\n    n_epochs = 150 # GANs often need more epochs\n    lr_g = 1e-4 # Learning rate for Generator\n    lr_d = 1e-4 # Learning rate for Discriminator\n    b1 = 0.5    # Adam beta1 (Common GAN value)\n    b2 = 0.999  # Adam beta2\n    lambda_gp = 10 # Gradient penalty coefficient\n    lambda_l1 = 100 # L1 content loss coefficient\n    n_critic = 5   # Train Discriminator n_critic times per Generator update\n\n    plot_every_n_epochs = 10 # Plot history less frequently\n    save_every_n_epochs = 10 # Save models periodically\n\n    # --- Misc ---\n    seed = 42\n    use_cuda = torch.cuda.is_available()\n    device = torch.device(\"cuda\" if use_cuda else \"cpu\")\n    # Use float16 on CUDA, bfloat16 on CPU (if available) for AMP\n    autocast_dtype = torch.float16 if use_cuda else (torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float32)\n\n# %% Helper Functions\ndef set_seed(seed=cfg.seed):\n    \"\"\"Sets seed for reproducibility across libraries.\"\"\"\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if cfg.use_cuda:\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n        # You might disable deterministic for performance in GANs if needed\n        # torch.backends.cudnn.deterministic = True\n        # torch.backends.cudnn.benchmark = False\n    print(f\"Seed set to {seed}\")\n\ndef find_best_model(model_dir=cfg.model_save_dir, model_prefix=\"generator_epoch\", suffix=\".pth\"):\n    \"\"\"Finds the model file with the highest epoch number.\"\"\"\n    best_epoch = -1\n    best_model_path = None\n    pattern = os.path.join(model_dir, f\"{model_prefix}_*_loss_*{suffix}\") # Keep loss for potential future use\n    all_model_files = glob.glob(os.path.join(model_dir, f\"{model_prefix}_*{suffix}\"))\n\n    if not all_model_files:\n        print(f\"W: No models matching pattern '{model_prefix}_*{suffix}' found in {model_dir}.\")\n        return None\n\n    parsed_models = []\n    for f in all_model_files:\n        try:\n            epoch_str = f.split(model_prefix + \"_\")[-1].split(\"_\")[0]\n            epoch = int(epoch_str)\n            parsed_models.append((epoch, f))\n        except (ValueError, IndexError, AttributeError):\n            print(f\"W: Couldn't parse epoch from filename: {os.path.basename(f)}\")\n\n    if parsed_models:\n        parsed_models.sort(key=lambda x: x[0], reverse=True) # Sort descending by epoch\n        best_epoch, best_model_path = parsed_models[0]\n        print(f\"Found latest model by epoch: {os.path.basename(best_model_path)} (Epoch: {best_epoch})\")\n    else:\n        # Fallback: Select most recently modified if parsing failed\n        print(\"W: No epochs parsed. Selecting most recently modified model.\")\n        best_model_path = max(all_model_files, key=os.path.getmtime, default=None)\n        if best_model_path:\n             print(f\"Using most recent modification time: {os.path.basename(best_model_path)}\")\n\n    return best_model_path\n\n# %% Data Normalization Functions (PLACEHOLDERS - UPDATE VALUES)\ndef normalize_vel(vel_tensor):\n    \"\"\"Normalizes velocity tensor to [-1, 1] using global min/max.\"\"\"\n    # !!! UPDATE cfg.VEL_MIN and cfg.VEL_MAX with your dataset's actual values !!!\n    if cfg.VEL_MAX == cfg.VEL_MIN:\n        print(\"W: VEL_MAX equals VEL_MIN, cannot normalize velocity.\")\n        return vel_tensor\n    # Formula: 2 * (x - min) / (max - min) - 1\n    return 2.0 * (vel_tensor - cfg.VEL_MIN) / (cfg.VEL_MAX - cfg.VEL_MIN) - 1.0\n\ndef unnormalize_vel(norm_vel_tensor):\n    \"\"\"Un-normalizes velocity tensor from [-1, 1] to original range.\"\"\"\n    # !!! Uses cfg.VEL_MIN and cfg.VEL_MAX !!!\n    if cfg.VEL_MAX == cfg.VEL_MIN:\n        print(\"W: VEL_MAX equals VEL_MIN, cannot unnormalize velocity.\")\n        return norm_vel_tensor\n    # Formula: ((y + 1) * (max - min) / 2) + min\n    return ((norm_vel_tensor + 1.0) * (cfg.VEL_MAX - cfg.VEL_MIN) / 2.0) + cfg.VEL_MIN\n\ndef normalize_seis(seis_tensor):\n    \"\"\"Normalizes seismic tensor based on cfg.SEIS_NORM_MODE.\"\"\"\n    if cfg.SEIS_NORM_MODE == 'minmax_sample':\n        # Normalize each sample (image) in the batch independently to [-1, 1]\n        min_val = torch.amin(seis_tensor, dim=(-1, -2, -3), keepdim=True)\n        max_val = torch.amax(seis_tensor, dim=(-1, -2, -3), keepdim=True)\n        range_val = max_val - min_val\n        # Add epsilon to avoid division by zero for constant samples\n        range_val = torch.where(range_val == 0, torch.tensor(1e-6, device=seis_tensor.device), range_val)\n        return 2.0 * (seis_tensor - min_val) / range_val - 1.0\n    elif cfg.SEIS_NORM_MODE == 'std_sample':\n        # Normalize each sample to mean 0, std 1\n        mean_val = torch.mean(seis_tensor, dim=(-1, -2, -3), keepdim=True)\n        std_val = torch.std(seis_tensor, dim=(-1, -2, -3), keepdim=True)\n        # Add epsilon to avoid division by zero\n        std_val = torch.where(std_val == 0, torch.tensor(1e-6, device=seis_tensor.device), std_val)\n        return (seis_tensor - mean_val) / std_val\n    elif cfg.SEIS_NORM_MODE is None:\n        return seis_tensor # No normalization\n    else:\n        raise ValueError(f\"Unknown SEIS_NORM_MODE: {cfg.SEIS_NORM_MODE}\")\n\n# %% WebDataset Preprocessing Functions\ndef search_data_path(target_dirs, root_dir, shuffle=True, seed=cfg.seed):\n    \"\"\"Finds input/output .npy file pairs within subdirectories.\"\"\"\n    # (Keep this function as is from your original code)\n    files = []\n    root_path = Path(root_dir)\n    if not root_path.is_dir():\n        print(f\"W: Root directory not found: {root_path}\")\n        return []\n\n    print(f\"Searching for data families {target_dirs} in root: {root_path}\")\n    total_pairs_found = 0\n    for target_dir in target_dirs:\n        data_dir = root_path / target_dir\n        if not data_dir.is_dir():\n            continue\n\n        in_files, out_files = [], []\n        data_subdir = data_dir / \"data\"\n        model_subdir = data_dir / \"model\"\n\n        if data_subdir.is_dir() and model_subdir.is_dir():\n            in_files = sorted(data_subdir.glob(\"*.npy\"))\n            out_files = sorted(model_subdir.glob(\"*.npy\"))\n        else:\n            in_files = sorted(data_dir.glob(\"seis*.npy\"))\n            out_files = sorted(data_dir.glob(\"vel*.npy\"))\n\n        if not in_files or len(in_files) != len(out_files):\n            if in_files or out_files:\n                print(\n                    f\"W: Mismatch or missing files in {data_dir} (in:{len(in_files)}, out:{len(out_files)}). Skipping.\"\n                )\n            continue\n\n        current_pairs = list(zip(in_files, out_files))\n        files.extend(current_pairs)\n        total_pairs_found += len(current_pairs)\n\n    print(f\"Found {len(files)} total valid pairs across specified families.\")\n    if shuffle and files:\n        print(f\"Shuffling {len(files)} pairs (seed={seed}).\")\n        rng = np.random.default_rng(seed)\n        rng.shuffle(files)\n\n    return files\n\n\ndef generate_sample(in_file, out_file=None, base_dir=None):\n    \"\"\"\n    Loads data from .npy files, prepares dicts for WebDataset, converts to float16.\n    Handles errors during loading gracefully.\n    \"\"\"\n    # (Keep this function as is from your original code)\n    # Note: It loads data but doesn't normalize here. Normalization happens later.\n    data = []\n    seis = None\n    vel = None\n    try:\n        if out_file is None:\n            print(\"W: generate_sample called without out_file (test mode?), not implemented.\")\n            return []\n        else:\n            try:\n                seis = np.load(in_file, mmap_mode=\"r\")\n            except Exception as e:\n                print(f\"E: Load fail for input {in_file.name}: {e}\")\n                return []\n\n            try:\n                vel = np.load(out_file, mmap_mode=\"r\")\n            except Exception as e:\n                print(f\"E: Load fail for output {out_file.name}: {e}\")\n                if seis is not None: del seis\n                return []\n\n            n_samples = 0\n            if seis.ndim == 4 and vel.ndim == 4:\n                if seis.shape[0] != vel.shape[0]:\n                    print(f\"W: Batch size mismatch in {in_file.name} ({seis.shape[0]}) vs {out_file.name} ({vel.shape[0]})\")\n                    del seis, vel\n                    return []\n                n_samples = seis.shape[0]\n            elif seis.ndim == 3 and vel.ndim == 3:\n                n_samples = 1\n            else:\n                 raise ValueError(f\"Unexpected dims: seis {seis.shape}, vel {vel.shape} in {in_file.name}\")\n\n            if n_samples == 0:\n                print(f\"W: Found 0 samples in pair: {in_file.name}, {out_file.name}\")\n                del seis, vel\n                return []\n\n            common_part = f\"{in_file.parent.name}_{in_file.stem}\"\n            if base_dir:\n                try:\n                    relative_path = in_file.relative_to(base_dir)\n                    common_part = \"_\".join(relative_path.parts).replace(\".npy\", \"\")\n                    common_part = common_part.replace(os.sep, \"_\").replace(\"\\\\\", \"_\")\n                except ValueError:\n                    pass\n\n            for i in range(n_samples):\n                key = f\"{common_part}_{i}\"\n                s_sample = (seis[i].copy().astype(np.float16) if seis.ndim == 4 else seis.copy().astype(np.float16))\n                v_sample = (vel[i].copy().astype(np.float16) if vel.ndim == 4 else vel.copy().astype(np.float16))\n                data.append({\"__key__\": key, \"sample_id.txt\": key, \"seis.npy\": s_sample, \"vel.npy\": v_sample})\n\n            del seis\n            del vel\n\n    except Exception as e:\n        print(f\"E: Error during sample generation for {in_file.name}: {e}\")\n        if seis is not None:\n            try: del seis\n            except NameError: pass\n        if vel is not None:\n            try: del vel\n            except NameError: pass\n        return []\n\n    return data\n\n\n# %% WebDataset Loading Functions\ndef get_shard_paths(root_dir, dataset_name, stage, num_shards=None, test_size=cfg.test_size, seed=cfg.seed):\n    \"\"\"Gets list of shard paths, optionally selects subset, optionally splits train/val.\"\"\"\n    # (Keep this function as is from your original code)\n    source_dir_name = f\"train_{dataset_name}\"\n    dataset_dir = Path(root_dir) / source_dir_name\n    print(f\"Looking for shards for stage '{stage}' in: {dataset_dir}\")\n\n    if not dataset_dir.is_dir():\n        print(f\"W: Shard directory not found: {dataset_dir}\")\n        return (None, None) if stage == \"train\" else None\n\n    shard_paths = sorted([str(p) for p in dataset_dir.glob(\"*.tar\")])\n\n    if not shard_paths:\n        print(f\"W: No .tar shards found in {dataset_dir}.\")\n        return (None, None) if stage == \"train\" else None\n\n    print(f\"Found {len(shard_paths)} total shards.\")\n    selected_paths = shard_paths\n    available_count = len(shard_paths)\n    if num_shards is not None:\n        if 0 < num_shards < available_count:\n            print(f\"Selecting {num_shards} shards randomly (seed={seed}).\")\n            rng = np.random.default_rng(seed)\n            indices = rng.choice(available_count, size=num_shards, replace=False)\n            selected_paths = sorted([shard_paths[i] for i in indices])\n        elif num_shards >= available_count:\n            print(f\"Requested {num_shards} or more shards, using all {available_count} available.\")\n        else:\n            print(f\"W: Invalid num_shards ({num_shards}). Using all {available_count} shards.\")\n    print(f\"Using {len(selected_paths)} selected shards for stage '{stage}'.\")\n\n    if stage == \"train\":\n        count = len(selected_paths)\n        print(f\"Splitting {count} selected shards (test_size={test_size}, seed={seed})\")\n        try:\n            if not (0 <= test_size < 1): raise ValueError(\"test_size must be in [0, 1)\")\n            if count <= 1 or test_size == 0:\n                reason = \"only 1 shard\" if count <= 1 else \"test_size is 0\"\n                print(f\"W: Cannot split for validation ({reason}). Assigning all to train.\")\n                return sorted(selected_paths), []\n            else:\n                trn_paths, val_paths = train_test_split(selected_paths, test_size=test_size, random_state=seed, shuffle=True)\n                trn_paths.sort()\n                val_paths.sort()\n                print(f\"# Train shards: {len(trn_paths)}, # Val shards: {len(val_paths)}\")\n                return trn_paths, val_paths\n        except Exception as e:\n            print(f\"E: Failed to split shards: {e}\")\n            return None, None\n    else:\n        print(f\"# Shards returned for stage '{stage}': {len(selected_paths)}\")\n        return sorted(selected_paths)\n\n\ndef get_dataset(paths, stage, seed=cfg.seed):\n    \"\"\"Creates WebDataset object. Applies augmentations and NORMALIZATION.\"\"\"\n    if not paths:\n        print(f\"W: No shard paths provided for stage '{stage}'. Cannot create dataset.\")\n        return None\n\n    print(f\"Creating WebDataset for stage '{stage}' from {len(paths)} shards.\")\n    is_train = stage == \"train\"\n    map_handler = wds.warn_and_continue\n\n    try:\n        dataset = wds.WebDataset(\n            paths, nodesplitter=wds.split_by_node, shardshuffle=is_train, seed=seed\n        )\n        dataset = dataset.decode(handler=map_handler) # Decode standard types\n\n        def map_train_val(sample):\n            \"\"\"Inner function to process, augment, and NORMALIZE samples.\"\"\"\n            key_info = sample.get(\"__key__\", \"N/A\")\n            try:\n                required = [\"sample_id.txt\", \"seis.npy\", \"vel.npy\"]\n                if not all(k in sample for k in required):\n                    raise KeyError(f\"Missing required keys in sample {key_info}\")\n\n                sid = sample[\"sample_id.txt\"]\n                # Convert to float32 tensors first\n                s_np = np.asarray(sample[\"seis.npy\"]).astype(np.float32)\n                v_np = np.asarray(sample[\"vel.npy\"]).astype(np.float32)\n                seis_tensor = torch.from_numpy(s_np)\n                vel_tensor = torch.from_numpy(v_np)\n\n                # --- Augmentation Block (Applied BEFORE Normalization potentially) ---\n                if is_train and cfg.apply_augmentation:\n                    # 1. Add Gaussian Noise to Seismic Data (relative to original range)\n                    if cfg.aug_seis_noise_std > 0:\n                        noise = torch.randn_like(seis_tensor) * cfg.aug_seis_noise_std\n                        seis_tensor.add_(noise) # In-place addition\n\n                    # 2. Horizontal Flip\n                    if torch.rand(1).item() < cfg.aug_hflip_prob:\n                        seis_tensor = TF.hflip(seis_tensor)\n                        vel_tensor = TF.hflip(vel_tensor)\n\n                # --- NORMALIZATION ---\n                # !!! Ensure your normalization functions handle the tensor shapes correctly !!!\n                seis_tensor_norm = normalize_seis(seis_tensor)\n                vel_tensor_norm = normalize_vel(vel_tensor)\n\n                return {\"sample_id\": sid, \"seis\": seis_tensor_norm, \"vel\": vel_tensor_norm}\n\n            except Exception as map_e:\n                print(f\"E: Map function failed for sample {key_info}: {map_e}\")\n                raise map_e # Let handler decide fate\n\n        if stage in [\"train\", \"val\"]:\n            dataset = dataset.map(map_train_val, handler=map_handler)\n\n        if is_train:\n            dataset = dataset.shuffle(1000) # Buffer size\n\n        return dataset\n\n    except Exception as e:\n        print(f\"E: Error creating WebDataset pipeline for stage '{stage}': {e}\")\n        return None\n\n\n# %% Kaggle TestSet Loading (Directly from .npy)\nclass KaggleTestDataset(Dataset):\n    \"\"\"Loads Kaggle test set .npy files and applies SEISMIC normalization.\"\"\"\n    def __init__(self, test_files_dir):\n        self.test_files_dir = Path(test_files_dir)\n        self.test_files = []\n        try:\n            if not self.test_files_dir.is_dir():\n                raise FileNotFoundError(f\"Kaggle test directory missing: {self.test_files_dir}\")\n            self.test_files = sorted(list(self.test_files_dir.glob(\"*.npy\")))\n            print(f\"Found {len(self.test_files)} '.npy' files in Kaggle test dir: {self.test_files_dir}\")\n            if not self.test_files: print(f\"W: No .npy files found in {self.test_files_dir}.\")\n        except Exception as e:\n            print(f\"E: Error accessing Kaggle test directory {self.test_files_dir}: {e}\")\n\n    def __len__(self):\n        return len(self.test_files)\n\n    def __getitem__(self, index):\n        if not self.test_files or index >= len(self.test_files):\n            raise IndexError(f\"Index {index} out of bounds ({len(self.test_files)} files).\")\n        test_file_path = self.test_files[index]\n        try:\n            # Load numpy array, convert to float32 tensor\n            data_np = np.load(test_file_path).astype(np.float32)\n            data_tensor = torch.from_numpy(data_np)\n            # Apply SEISMIC normalization consistent with training\n            data_tensor_norm = normalize_seis(data_tensor)\n            original_id = test_file_path.stem\n            return data_tensor_norm, original_id\n        except Exception as e:\n            raise IOError(f\"Error loading/normalizing Kaggle test file: {test_file_path}\") from e\n\n\n# %% U-Net Blocks (DoubleConv, Down, Up, OutConv - Keep as is)\nclass DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels, mid_channels=None):\n        super().__init__()\n        if not mid_channels: mid_channels = out_channels\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(mid_channels), nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels), nn.ReLU(inplace=True))\n    def forward(self, x): return self.double_conv(x)\n\nclass Down(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.maxpool_conv = nn.Sequential(nn.MaxPool2d(2), DoubleConv(in_channels, out_channels))\n    def forward(self, x): return self.maxpool_conv(x)\n\nclass Up(nn.Module):\n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super().__init__(); self.bilinear = bilinear\n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode=\"bilinear\", align_corners=False)\n            self.conv = DoubleConv(in_channels + out_channels, out_channels, mid_channels=out_channels)\n        else:\n            self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)\n            self.conv = DoubleConv(in_channels // 2 + out_channels, out_channels)\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        diffY = x2.size(2) - x1.size(2); diffX = x2.size(3) - x1.size(3)\n        x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2])\n        x = torch.cat([x2, x1], dim=1); return self.conv(x)\n\nclass OutConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__(); self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)\n    def forward(self, x): return self.conv(x)\n\n\n# %% Generator Definition (Adapted U-Net)\nclass GeneratorUNet(nn.Module):\n    \"\"\"Generator using U-Net architecture, outputs normalized velocity map.\"\"\"\n    def __init__( self, n_channels=cfg.unet_in_channels, n_classes=cfg.unet_out_channels,\n                 init_features=cfg.unet_init_features, depth=cfg.unet_depth, bilinear=cfg.unet_bilinear):\n        super().__init__()\n        self.n_channels = n_channels; self.n_classes = n_classes\n        self.bilinear = bilinear; self.depth = depth\n        self.initial_pool = nn.AvgPool2d(kernel_size=(14, 1), stride=(14, 1))\n        self.encoder_blocks = nn.ModuleList(); self.inc = DoubleConv(n_channels, init_features)\n        self.encoder_blocks.append(self.inc); current_features = init_features\n        for _ in range(depth):\n            down_block = Down(current_features, current_features * 2)\n            self.encoder_blocks.append(down_block); current_features *= 2\n        bottleneck_features = current_features\n        self.decoder_blocks = nn.ModuleList(); current_features = bottleneck_features\n        for _ in range(depth):\n            up_block = Up(current_features, current_features // 2, bilinear)\n            self.decoder_blocks.append(up_block); current_features //= 2\n        self.outc = OutConv(current_features, n_classes)\n        self.final_activation = nn.Tanh() # Output in [-1, 1] matching normalization\n        self.processed_seismic = None # To store condition for Discriminator\n\n    def _pad_or_crop(self, x, target_h=70, target_w=70):\n        _, _, h, w = x.shape\n        if h < target_h: pad_top = (target_h - h) // 2; pad_bottom = target_h - h - pad_top; x = F.pad(x, (0, 0, pad_top, pad_bottom)); h = target_h\n        if w < target_w: pad_left = (target_w - w) // 2; pad_right = target_w - w - pad_left; x = F.pad(x, (pad_left, pad_right, 0, 0)); w = target_w\n        if h > target_h: crop_top = (h - target_h) // 2; x = x[:, :, crop_top : crop_top + target_h, :]; h = target_h\n        if w > target_w: crop_left = (w - target_w) // 2; x = x[:, :, :, crop_left : crop_left + target_w]; w = target_w\n        return x\n\n    def forward(self, x_seismic):\n        x_pooled = self.initial_pool(x_seismic)\n        x_resized = self._pad_or_crop(x_pooled, target_h=70, target_w=70)\n        self.processed_seismic = x_resized # Store for D conditioning\n        skip_connections = []; xi = x_resized\n        for i, block in enumerate(self.encoder_blocks):\n            xi = block(xi)\n            if i < len(self.encoder_blocks) - 1: skip_connections.append(xi)\n        xu = xi\n        for i, block in enumerate(self.decoder_blocks):\n            skip = skip_connections[len(skip_connections) - 1 - i]; xu = block(xu, skip)\n        logits = self.outc(xu); output = self.final_activation(logits)\n        return output\n\n# %% Discriminator Definition (Example CNN)\n\n# %% Discriminator Definition (Corrected Final Layer)\nclass Discriminator(nn.Module):\n    \"\"\"Discriminator network for FWIGAN (Conditional).\"\"\"\n    def __init__(self, in_channels=cfg.disc_in_channels, init_features=cfg.disc_init_features):\n        super().__init__()\n        def discriminator_block(in_filters, out_filters, bn=True):\n            block = [ nn.Conv2d(in_filters, out_filters, kernel_size=4, stride=2, padding=1),\n                      nn.LeakyReLU(0.2, inplace=True)]\n            if bn:\n                 block.append(nn.BatchNorm2d(out_filters)) # Consider InstanceNorm/LayerNorm or removal if issues persist\n            return block\n\n        nf = init_features\n        self.model = nn.Sequential(\n            *discriminator_block(in_channels, nf, bn=False), # (B, 6, 70, 70) -> (B, 64, 35, 35)\n            *discriminator_block(nf, nf * 2),                # -> (B, 128, 17, 17) ~ Corrected Size\n            *discriminator_block(nf * 2, nf * 4),            # -> (B, 256, 8, 8)  ~ Corrected Size\n            *discriminator_block(nf * 4, nf * 8),            # -> (B, 512, 4, 4)  ~ Corrected Size\n            # *** FIXED KERNEL SIZE HERE ***\n            # Final layer outputs a single score (no sigmoid for WGAN)\n            nn.Conv2d(nf * 8, 1, kernel_size=4, stride=1, padding=0) # Output: (B, 1, 1, 1)\n        )\n\n    def forward(self, velocity_map, seismic_condition):\n        # Concatenate along channel dimension\n        img_input = torch.cat((velocity_map, seismic_condition), dim=1)\n        return self.model(img_input) # Output is raw score\n\n\n\n\n\n# %% Weight Initialization Function\ndef weights_init_normal(m):\n    classname = m.__class__.__name__\n    if classname.find(\"Conv\") != -1:\n        try: torch.nn.init.normal_(m.weight.data, 0.0, 0.02)\n        except: pass # handles bias items etc.\n    elif classname.find(\"BatchNorm\") != -1:\n        torch.nn.init.normal_(m.weight.data, 1.0, 0.02)\n        torch.nn.init.constant_(m.bias.data, 0.0)\n\n# %% WGAN-GP Gradient Penalty Function\ndef compute_gradient_penalty(D, real_samples_vel, fake_samples_vel, seismic_cond, device):\n    \"\"\"Calculates the gradient penalty loss for WGAN GP\"\"\"\n    batch_size = real_samples_vel.size(0)\n    alpha = torch.rand(batch_size, 1, 1, 1, device=device) # Shape (B, 1, 1, 1)\n    interpolates_vel = (alpha * real_samples_vel.data + ((1 - alpha) * fake_samples_vel.data)).requires_grad_(True)\n    # Use the same seismic condition for interpolated samples\n    d_interpolates = D(interpolates_vel, seismic_cond.data) # Pass condition\n\n    # Use torch.ones_like for fake gradients to match output shape\n    fake = torch.ones_like(d_interpolates, device=device, requires_grad=False)\n\n    gradients = autograd.grad(\n        outputs=d_interpolates, inputs=interpolates_vel, grad_outputs=fake,\n        create_graph=True, retain_graph=True, only_inputs=True,\n    )[0]\n    gradients = gradients.view(batch_size, -1)\n    gradient_penalty = ((gradients.norm(2, dim=1) - 1) ** 2).mean()\n    return gradient_penalty\n\n\n# %% Main Execution Script\nprint(\"--- Starting Full FWIGAN Workflow ---\")\nstart_time = time.time()\nset_seed(cfg.seed)\nprint(f\"Device: {cfg.device}\")\nprint(f\"Using PyTorch version: {torch.__version__}\")\nif cfg.use_cuda: print(f\"CUDA available: {torch.cuda.get_device_name(0)}\")\nprint(f\"AMP dtype: {cfg.autocast_dtype}\")\nprint(f\"Velocity Norm Range: [{cfg.VEL_MIN}, {cfg.VEL_MAX}] -> [-1, 1]\")\nprint(f\"Seismic Norm Mode: {cfg.SEIS_NORM_MODE}\")\n\n\n# ==============================================================================\n# Cleanup Code\n# ==============================================================================\nprint(\"\\n--- Cleaning up previous run artifacts ---\")\nPath(cfg.model_save_dir).mkdir(parents=True, exist_ok=True) # Ensure model dir exists\npaths_to_clean = [cfg.shard_output_dir]\n# Clean previous models (adjust pattern if needed)\nmodel_patterns = [\n    os.path.join(cfg.model_save_dir, \"generator_epoch_*.pth\"),\n    os.path.join(cfg.model_save_dir, \"discriminator_epoch_*.pth\")\n]\nfor pattern in model_patterns: paths_to_clean.extend(glob.glob(pattern))\npaths_to_clean.append(os.path.join(cfg.working_dir, \"training_history.png\"))\npaths_to_clean.append(cfg.submission_file)\n\nfor path_str in paths_to_clean:\n    path_obj = Path(path_str)\n    try:\n        if path_obj.is_dir(): shutil.rmtree(path_obj, ignore_errors=True)\n        elif path_obj.is_file(): path_obj.unlink(missing_ok=True)\n    except Exception as e: print(f\"W: Error during cleanup of {path_obj}: {e}\")\nprint(\"--- Cleanup finished ---\")\ngc.collect()\n\n# ==============================================================================\n# SECTION 0/1: Sharding from Kaggle Data Only\n# ==============================================================================\nprint(\"\\n--- 0/1. Sharding from Kaggle Data Only ---\")\n# (Keep this section largely as is from your original code)\nshard_stage_dir = Path(cfg.shard_output_dir) / f\"train_{cfg.dataset_name}\"\nkaggle_train_root = Path(cfg.kaggle_train_dir)\nneeds_creation = True; total_samples_written = 0\ntry:\n    if shard_stage_dir.exists() and any(shard_stage_dir.glob(\"*.tar\")):\n        if cfg.force_shard_creation:\n            print(f\"Forcing shard creation. Removing existing shards in {shard_stage_dir}\")\n            shutil.rmtree(shard_stage_dir)\n        else:\n            print(f\"Found existing shards at: {shard_stage_dir}. Skipping creation.\")\n            needs_creation = False\n    Path(cfg.shard_output_dir).mkdir(parents=True, exist_ok=True)\n    shard_stage_dir.mkdir(parents=True, exist_ok=True)\n\n    if needs_creation:\n        print(f\"Starting shard creation from {kaggle_train_root} into {shard_stage_dir}\")\n        if not kaggle_train_root.is_dir(): raise FileNotFoundError(f\"Kaggle train dir not found: {kaggle_train_root}\")\n        families = [d.name for d in kaggle_train_root.iterdir() if d.is_dir()]\n        if not families: raise FileNotFoundError(f\"No family subdirs found in {kaggle_train_root}\")\n        print(f\"Searching Kaggle data families: {families}\")\n        kaggle_file_pairs = search_data_path(families, kaggle_train_root, shuffle=True, seed=cfg.seed)\n        if not kaggle_file_pairs: raise RuntimeError(\"No valid data pairs found.\")\n\n        shard_pattern = str(shard_stage_dir / \"%06d.tar\")\n        print(f\"Writing shards using pattern {shard_pattern} (max size {cfg.maxsize / 1e9:.2f} GB)\")\n        with wds.ShardWriter(shard_pattern, maxsize=int(cfg.maxsize)) as writer:\n            common_base_dir = kaggle_train_root\n            for in_file, out_file in tqdm(kaggle_file_pairs, desc=\"Sharding Kaggle Data\", unit=\"pair\"):\n                samples_from_pair = generate_sample(Path(in_file), Path(out_file), base_dir=common_base_dir)\n                if samples_from_pair:\n                    for sample_dict in samples_from_pair: writer.write(sample_dict)\n                    total_samples_written += len(samples_from_pair)\n        print(f\"Finished writing {total_samples_written} samples to shards.\")\n    else: print(f\"Using existing shards in {shard_stage_dir}.\")\n\nexcept Exception as e: print(f\"E: Sharding process failed: {e}\"); raise\n\n# ==============================================================================\n# SECTION 2: Create DataLoaders (Using modified get_dataset with Normalization)\n# ==============================================================================\nprint(\"\\n--- 2. Creating DataLoaders from Shards (with Normalization) ---\")\ndltrain, dlvalid = None, None; val_paths_saved = []\ntry:\n    trn_paths, val_paths = get_shard_paths(cfg.shard_output_dir, cfg.dataset_name, \"train\")\n    val_paths_saved = val_paths\n    if trn_paths is None: raise RuntimeError(\"Failed to get or split shard paths.\")\n\n    # Check if shards exist if paths are empty\n    shard_check_dir = Path(cfg.shard_output_dir) / f\"train_{cfg.dataset_name}\"\n    if not trn_paths and not list(shard_check_dir.glob(\"*.tar\")):\n         raise RuntimeError(f\"No training shards found in {shard_check_dir}.\")\n    if not trn_paths: print(\"W: No shards for training.\")\n    else: print(f\"Using {len(trn_paths)} shards for training.\")\n    if not val_paths: print(\"W: No shards for validation.\")\n    else: print(f\"Using {len(val_paths)} shards for validation.\")\n\n    trn_ds = get_dataset(trn_paths, \"train\", seed=cfg.seed) if trn_paths else None\n    val_ds = get_dataset(val_paths, \"val\", seed=cfg.seed + 1) if val_paths else None # Use different seed for val\n\n    if trn_ds:\n        n_trn_w = min(cfg.num_workers, os.cpu_count(), len(trn_paths)) if trn_paths else 0\n        p_trn = n_trn_w > 0\n        dltrain = DataLoader(trn_ds.batched(cfg.batch_size), batch_size=None, shuffle=False,\n                             num_workers=n_trn_w, pin_memory=cfg.use_cuda, persistent_workers=p_trn,\n                             prefetch_factor=2 if p_trn else None)\n        print(f\"Train DataLoader created (workers={n_trn_w}, persistent={p_trn}).\")\n    if val_ds:\n        n_val_w = min(cfg.num_workers, os.cpu_count(), len(val_paths)) if val_paths else 0\n        p_val = n_val_w > 0\n        dlvalid = DataLoader(val_ds.batched(cfg.batch_size), batch_size=None, shuffle=False,\n                             num_workers=n_val_w, pin_memory=cfg.use_cuda, persistent_workers=p_val,\n                             prefetch_factor=2 if p_val else None)\n        print(f\"Validation DataLoader created (workers={n_val_w}, persistent={p_val}).\")\n\nexcept Exception as e: print(f\"E: DataLoader creation failed: {e}\"); raise\n\n# ==============================================================================\n# SECTION 3: Initialize Models, Optimizers, Losses\n# ==============================================================================\nprint(\"\\n--- 3. Initializing Models, Losses, Optimizers ---\")\ntry:\n    generator = GeneratorUNet().to(cfg.device)\n    discriminator = Discriminator().to(cfg.device)\n\n    # Initialize weights\n    generator.apply(weights_init_normal)\n    discriminator.apply(weights_init_normal)\n    print(\"Applied normal weight initialization.\")\n\n    # Optimizers (Using Adam as recommended for GANs)\n    optimizer_G = torch.optim.Adam(generator.parameters(), lr=cfg.lr_g, betas=(cfg.b1, cfg.b2))\n    optimizer_D = torch.optim.Adam(discriminator.parameters(), lr=cfg.lr_d, betas=(cfg.b1, cfg.b2))\n    print(f\"Optimizers: Adam (G: lr={cfg.lr_g}, D: lr={cfg.lr_d}, beta1={cfg.b1})\")\n\n    # Loss Functions\n    criterion_L1 = nn.L1Loss().to(cfg.device) # Content loss\n    print(f\"Losses: Adversarial (WGAN-GP), Content (L1, lambda={cfg.lambda_l1})\")\n\n    # AMP Grad Scalers\n\n    # NEW Lines (Corrected GradScaler initialization)\n# Use 'cuda' if using GPU, 'cpu' otherwise. enabled logic remains the same.\n    device_type = 'cuda' if cfg.use_cuda else 'cpu'\n    g_scaler = torch.amp.GradScaler(device_type, enabled=(cfg.use_cuda and cfg.autocast_dtype != torch.float32))\n    d_scaler = torch.amp.GradScaler(device_type, enabled=(cfg.use_cuda and cfg.autocast_dtype != torch.float32))\n    print(f\"AMP GradScaler enabled: {g_scaler.is_enabled()}\") # Keep check\n    # g_scaler = torch.cuda.amp.GradScaler(enabled=(cfg.use_cuda and cfg.autocast_dtype != torch.float32))\n    # d_scaler = torch.cuda.amp.GradScaler(enabled=(cfg.use_cuda and cfg.autocast_dtype != torch.float32))\n    # print(f\"AMP GradScaler enabled: {g_scaler.is_enabled()}\")\n\n    params_g = sum(p.numel() for p in generator.parameters() if p.requires_grad)\n    params_d = sum(p.numel() for p in discriminator.parameters() if p.requires_grad)\n    print(f\"Generator Params: {params_g:,}\")\n    print(f\"Discriminator Params: {params_d:,}\")\n\nexcept Exception as e: print(f\"E: Model/Optimizer initialization failed: {e}\"); raise\n\n# ==============================================================================\n# SECTION 4: Training Loop\n# ==============================================================================\nprint(\"\\n--- 4. Starting FWIGAN Training ---\")\nhistory = {\"epoch\": [], \"d_loss\": [], \"g_loss\": [], \"l1_loss\": [], \"val_l1_loss\": []}\nbatches_done = 0\n\nif dltrain is None:\n    print(\"E: Training cannot proceed. Train DataLoader is missing.\")\nelse:\n    try:\n        for epoch in range(1, cfg.n_epochs + 1):\n            epoch_start_time = time.time()\n            generator.train()\n            discriminator.train()\n            epoch_d_losses, epoch_g_losses, epoch_l1_losses = [], [], []\n\n            pbar_train = tqdm(dltrain, desc=f\"Train E{epoch}\", leave=False, unit=\"batch\")\n\n            for i, batch in enumerate(pbar_train):\n                if not batch or \"seis\" not in batch or \"vel\" not in batch:\n                     print(f\"W: Skipping invalid train batch {i}\")\n                     continue\n\n                real_seis = batch[\"seis\"].to(cfg.device, non_blocking=True).float()\n                real_vel = batch[\"vel\"].to(cfg.device, non_blocking=True).float()\n                current_batch_size = real_seis.size(0) # Get actual batch size\n\n                # ---------------------\n                #  Train Discriminator\n                # ---------------------\n                optimizer_D.zero_grad(set_to_none=True)\n\n                # Use AMP context manager for Discriminator forward/loss\n                with torch.amp.autocast(device_type=cfg.device.type, dtype=cfg.autocast_dtype, enabled=g_scaler.is_enabled()):\n                    # Generate fake velocity map (no grad for G here)\n                    with torch.no_grad():\n                         fake_vel = generator(real_seis)\n                    # Get the processed seismic condition used by the generator\n                    processed_seis_cond = generator.processed_seismic.detach()\n\n                    # Calculate D scores\n                    real_validity = discriminator(real_vel, processed_seis_cond)\n                    fake_validity = discriminator(fake_vel.detach(), processed_seis_cond) # Detach fake_vel\n\n                    # Gradient penalty\n                    gradient_penalty = compute_gradient_penalty(\n                        discriminator, real_vel, fake_vel, processed_seis_cond, cfg.device\n                    )\n\n                    # Adversarial loss (WGAN-GP)\n                    d_loss = -torch.mean(real_validity) + torch.mean(fake_validity) + cfg.lambda_gp * gradient_penalty\n\n                # Scale loss and backpropagate (Discriminator)\n                d_scaler.scale(d_loss).backward()\n                d_scaler.step(optimizer_D)\n                d_scaler.update()\n                epoch_d_losses.append(d_loss.item())\n\n                # -----------------\n                #  Train Generator\n                # -----------------\n                # Train Generator every n_critic discriminator iterations\n                if i % cfg.n_critic == 0:\n                    optimizer_G.zero_grad(set_to_none=True)\n\n                    # Use AMP context manager for Generator forward/loss\n                    with torch.amp.autocast(device_type=cfg.device.type, dtype=cfg.autocast_dtype, enabled=g_scaler.is_enabled()):\n                        # Generate fake velocity map (track grads for G now)\n                        gen_vel = generator(real_seis)\n                        # Get the corresponding processed seismic condition\n                        processed_seis_cond_gen = generator.processed_seismic\n\n                        # Get discriminator score for generated map\n                        gen_validity = discriminator(gen_vel, processed_seis_cond_gen)\n\n                        # Adversarial loss (aims for high D score)\n                        g_adv_loss = -torch.mean(gen_validity)\n\n                        # Content loss (L1 distance between normalized maps)\n                        g_l1_loss = criterion_L1(gen_vel, real_vel)\n\n                        # Total generator loss\n                        g_loss = g_adv_loss + cfg.lambda_l1 * g_l1_loss\n\n                    # Scale loss and backpropagate (Generator)\n                    g_scaler.scale(g_loss).backward()\n                    g_scaler.step(optimizer_G)\n                    g_scaler.update()\n                    epoch_g_losses.append(g_loss.item())\n                    epoch_l1_losses.append(g_l1_loss.item())\n\n                    batches_done += 1\n\n                # Update progress bar description\n                if i % 100 == 0: # Update less frequently\n                    pbar_train.set_postfix(\n                        D_Loss=f\"{np.mean(epoch_d_losses[-50:]):.3f}\",\n                        G_Loss=f\"{np.mean(epoch_g_losses[-10:]):.3f}\",\n                        L1=f\"{np.mean(epoch_l1_losses[-10:]):.3f}\"\n                    )\n\n            # --- End of Epoch ---\n            avg_d_loss = np.mean(epoch_d_losses) if epoch_d_losses else 0\n            avg_g_loss = np.mean(epoch_g_losses) if epoch_g_losses else 0\n            avg_l1_loss = np.mean(epoch_l1_losses) if epoch_l1_losses else 0\n            epoch_time = time.time() - epoch_start_time\n\n            print(f\"Epoch {epoch}/{cfg.n_epochs} [{epoch_time:.2f}s] - D_Loss: {avg_d_loss:.4f}, G_Loss: {avg_g_loss:.4f}, L1_Loss: {avg_l1_loss:.4f}\")\n            history[\"epoch\"].append(epoch)\n            history[\"d_loss\"].append(avg_d_loss)\n            history[\"g_loss\"].append(avg_g_loss)\n            history[\"l1_loss\"].append(avg_l1_loss)\n\n            # --- Validation Phase (Optional but recommended) ---\n            if dlvalid is not None:\n                generator.eval() # Generator only for validation\n                val_l1 = []\n                with torch.no_grad():\n                    for batch_val in tqdm(dlvalid, desc=f\"Valid E{epoch}\", leave=False):\n                         if not batch_val or \"seis\" not in batch_val or \"vel\" not in batch_val: continue\n                         seis_val = batch_val[\"seis\"].to(cfg.device).float()\n                         vel_val = batch_val[\"vel\"].to(cfg.device).float()\n                         with torch.amp.autocast(device_type=cfg.device.type, dtype=cfg.autocast_dtype, enabled=g_scaler.is_enabled()):\n                             pred_vel = generator(seis_val)\n                             # Calculate L1 loss on normalized validation data\n                             loss_l1 = criterion_L1(pred_vel, vel_val)\n                         val_l1.append(loss_l1.item())\n                avg_val_l1 = np.mean(val_l1) if val_l1 else 0\n                print(f\"Epoch {epoch} Avg Validation L1_Loss: {avg_val_l1:.4f}\")\n                history[\"val_l1_loss\"].append(avg_val_l1)\n            else:\n                history[\"val_l1_loss\"].append(None) # Append None if no validation\n\n            # --- Save Models Periodically ---\n            if epoch % cfg.save_every_n_epochs == 0 or epoch == cfg.n_epochs:\n                g_path = os.path.join(cfg.model_save_dir, f\"generator_epoch_{epoch}.pth\")\n                d_path = os.path.join(cfg.model_save_dir, f\"discriminator_epoch_{epoch}.pth\")\n                torch.save(generator.state_dict(), g_path)\n                torch.save(discriminator.state_dict(), d_path)\n                print(f\"Saved models at epoch {epoch} to {cfg.model_save_dir}\")\n\n    except KeyboardInterrupt:\n        print(\"\\n--- Training interrupted by user ---\")\n    except Exception as e:\n        print(f\"\\nE: Training loop encountered a critical error: {e}\")\n        import traceback\n        traceback.print_exc()\n    finally:\n        print(\"\\n--- Training Loop Finished ---\")\n\n\n# ==============================================================================\n# SECTION 5: Plot History\n# ==============================================================================\nprint(\"\\n--- 5. Plotting Training History ---\")\nif history[\"epoch\"]:\n    try:\n        hist_df = pd.DataFrame(history)\n        plt.figure(figsize=(15, 7))\n\n        plt.subplot(1, 2, 1) # Loss plot\n        plt.plot(hist_df[\"epoch\"], hist_df[\"d_loss\"], \"o-\", label=\"Discriminator Loss\")\n        plt.plot(hist_df[\"epoch\"], hist_df[\"g_loss\"], \"s-\", label=\"Generator Loss\")\n        plt.title(\"GAN Losses vs. Epoch\")\n        plt.xlabel(\"Epoch\"); plt.ylabel(\"Loss\"); plt.legend(); plt.grid(True, alpha=0.5)\n\n        plt.subplot(1, 2, 2) # L1 Loss plot\n        plt.plot(hist_df[\"epoch\"], hist_df[\"l1_loss\"], \"^-\", label=\"Train L1 Loss (Content)\")\n        if not hist_df[\"val_l1_loss\"].isnull().all():\n            plt.plot(hist_df[\"epoch\"], hist_df[\"val_l1_loss\"], \"v--\", label=\"Validation L1 Loss\")\n        plt.title(\"L1 Content Loss vs. Epoch\")\n        plt.xlabel(\"Epoch\"); plt.ylabel(\"L1 Loss\"); plt.legend(); plt.grid(True, alpha=0.5)\n        plt.ylim(bottom=0) # L1 should be non-negative\n\n        plt.tight_layout()\n        plot_fname = os.path.join(cfg.working_dir, \"training_history.png\")\n        plt.savefig(plot_fname)\n        print(f\"Saved history plot: {plot_fname}\")\n        plt.show()\n    except Exception as e: print(f\"E: Failed plotting training history: {e}\")\nelse: print(\"No training history recorded to plot.\")\n\n# ==============================================================================\n# SECTION 6: Error Analysis / Validation Visualization (Placeholder)\n# ==============================================================================\nprint(\"\\n--- 6. Validation Set Visualization (Example) ---\")\nbest_generator_path = find_best_model(model_prefix=\"generator_epoch\") # Find latest Generator\n\nif not best_generator_path:\n    print(\"W: No generator model found. Skipping validation visualization.\")\nelif dlvalid is None:\n    print(\"W: Validation loader unavailable. Skipping visualization.\")\nelse:\n    print(f\"Visualizing using generator: {os.path.basename(best_generator_path)}\")\n    try:\n        # Load the latest generator model found\n        vis_generator = GeneratorUNet().to(cfg.device)\n        vis_generator.load_state_dict(torch.load(best_generator_path, map_location=cfg.device))\n        vis_generator.eval()\n\n        # Get a sample batch from validation loader\n        val_batch = next(iter(dlvalid))\n        seis_val = val_batch[\"seis\"].to(cfg.device).float()\n        vel_val_norm = val_batch[\"vel\"].to(cfg.device).float() # Ground truth (normalized)\n\n        with torch.no_grad():\n            with torch.amp.autocast(device_type=cfg.device.type, dtype=cfg.autocast_dtype, enabled=g_scaler.is_enabled()):\n                pred_vel_norm = vis_generator(seis_val) # Predicted (normalized)\n\n        # Un-normalize for visualization\n        vel_val_unnorm = unnormalize_vel(vel_val_norm)\n        pred_vel_unnorm = unnormalize_vel(pred_vel_norm)\n\n        # Plot the first sample in the batch\n        idx_to_plot = 0\n        plt.figure(figsize=(12, 6))\n        plt.subplot(1, 2, 1)\n        plt.imshow(vel_val_unnorm[idx_to_plot, 0].cpu().numpy(), cmap='viridis', aspect='auto')\n        plt.title(f\"Ground Truth Velocity (Sample {idx_to_plot})\")\n        plt.colorbar(label='Velocity')\n        plt.subplot(1, 2, 2)\n        plt.imshow(pred_vel_unnorm[idx_to_plot, 0].cpu().numpy(), cmap='viridis', aspect='auto')\n        plt.title(f\"Predicted Velocity (Epoch: {best_generator_path.split('_')[-1].split('.')[0]})\")\n        plt.colorbar(label='Velocity')\n        plt.tight_layout()\n        plt.show()\n\n    except Exception as e:\n        print(f\"E: Error during validation visualization: {e}\")\n        import traceback; traceback.print_exc()\n\n# ==============================================================================\n# SECTION 7: Prediction on Kaggle Test Set\n# ==============================================================================\nprint(\"\\n--- 7. Final Prediction on Kaggle Test Set ---\")\nbest_generator_final_path = find_best_model(model_prefix=\"generator_epoch\") # Find latest Generator\n\nif not best_generator_final_path:\n    print(\"W: No generator model found. Skipping final prediction.\")\nelif not Path(cfg.kaggle_test_dir).is_dir():\n    print(f\"W: Kaggle test directory '{cfg.kaggle_test_dir}' not found. Skipping prediction.\")\nelse:\n    try:\n        print(f\"Loading generator for prediction: {os.path.basename(best_generator_final_path)}\")\n        model_pred = GeneratorUNet().to(cfg.device)\n        model_pred.load_state_dict(torch.load(best_generator_final_path, map_location=cfg.device))\n        model_pred.eval()\n\n        test_ds = KaggleTestDataset(cfg.kaggle_test_dir) # Applies seismic normalization\n        if len(test_ds) == 0:\n            print(\"W: Kaggle test dataset is empty. No submission generated.\")\n        else:\n            t_bs = max(1, cfg.batch_size // 2) # Use smaller batch size for inference if needed\n            t_nw = min(max(0, cfg.num_workers // 2), os.cpu_count())\n            dl_test = DataLoader(test_ds, batch_size=t_bs, shuffle=False, num_workers=t_nw, pin_memory=cfg.use_cuda)\n            print(f\"Test DataLoader created (bs={t_bs}, workers={t_nw})\")\n            print(f\"Writing submission file to: {cfg.submission_file}\")\n\n            rows_written = 0\n            with open(cfg.submission_file, \"wt\", newline=\"\") as csvfile:\n                x_cols = [f\"x_{i}\" for i in range(1, 70, 2)]\n                fieldnames = [\"oid_ypos\"] + x_cols\n                writer = csv.DictWriter(csvfile, fieldnames=fieldnames)\n                writer.writeheader()\n\n                pbar_test = tqdm(dl_test, desc=\"Generating Submission\", unit=\"batch\")\n                with torch.no_grad():\n                    for inputs_norm, original_ids in pbar_test: # Input is normalized seismic\n                        if isinstance(original_ids, str): original_ids = [original_ids]\n                        try:\n                            inputs_norm = inputs_norm.to(cfg.device).float()\n                            with torch.amp.autocast(device_type=cfg.device.type, dtype=cfg.autocast_dtype, enabled=g_scaler.is_enabled()):\n                                outputs_norm = model_pred(inputs_norm) # Output is normalized velocity\n\n                            # !!! UN-NORMALIZE the prediction !!!\n                            outputs_unnorm = unnormalize_vel(outputs_norm)\n\n                            # Get predictions as numpy array (B, H, W)\n                            preds = outputs_unnorm[:, 0].cpu().numpy()\n\n                            for y_pred, oid in zip(preds, original_ids): # y_pred is (H=70, W=70)\n                                for y_pos in range(y_pred.shape[0]):\n                                    vals = y_pred[y_pos, 1::2].astype(np.float32) # Extract odd columns\n                                    row = dict(zip(x_cols, vals))\n                                    row[\"oid_ypos\"] = f\"{oid}_y_{y_pos}\"\n                                    writer.writerow(row); rows_written += 1\n                        except Exception as e:\n                            print(f\"\\nE: Prediction failed for batch (OID: {original_ids[0] if original_ids else '?'}) : {e}\")\n\n            print(f\"Submission file created: {cfg.submission_file} ({rows_written} rows).\")\n            expected_rows = len(test_ds) * 70\n            if rows_written != expected_rows:\n                print(f\"W: Row count mismatch! Expected {expected_rows}, but wrote {rows_written}.\")\n\n    except Exception as e:\n        print(f\"E: Final prediction process failed critically: {e}\")\n        import traceback; traceback.print_exc()\n\n# ==============================================================================\n# End of Workflow\n# ==============================================================================\nend_time = time.time()\nprint(f\"\\n--- Full Workflow Finished in {(end_time - start_time) / 60:.2f} minutes ---\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-20T10:31:14.607477Z","iopub.execute_input":"2025-04-20T10:31:14.607721Z"}},"outputs":[],"execution_count":null}]}