{"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":31011,"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 U-Net model (with Residual Blocks) for Full Waveform\nInversion, using data sourced solely from Kaggle input directories, with\ndata augmentation. Corrected generate_sample function.\n\"\"\"\n\n# %% Imports\n# Standard Library Imports\nimport csv\nimport gc\nimport glob\nimport os\nimport random\nimport shutil\nimport sys\nfrom pathlib import Path\n\n# Third-party Imports\n# Install webdataset if not present (useful in notebook environments)\ntry:\n    import webdataset as wds\nexcept ImportError:\n    print(\"Installing webdataset...\")\n    !pip install webdataset\n    import webdataset as wds\n\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.amp\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms.functional as TF  # For Augmentation\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\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\n    # --- Dataset Params ---\n    dataset_name = \"fwi_kaggle_only_augmented_resnet\" # Updated name\n\n    # --- Sharding Params ---\n    maxsize = 1e9  # Approx 1 GB\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  # Reduced from 16 to 8\n    num_workers = 1  # Reduced from 2 to 1\n\n    # --- Augmentation Params ---\n    apply_augmentation = True\n    aug_hflip_prob = 0.5  # Probability of horizontal flip\n    aug_seis_noise_std = 0.01  # Std dev of Gaussian noise added to seismic\n\n    # --- Model params (U-Net with Residual Blocks) ---\n    unet_in_channels = 5\n    unet_out_channels = 1\n    unet_init_features = 24  # Reduced from 32 to 24\n    unet_depth = 4  # Reduced from 5 to 4\n    unet_bilinear = True # Upsampling method\n\n    # --- Training params ---\n    n_epochs = 100\n    learning_rate = 1e-4\n    weight_decay = 1e-5\n    plot_every_n_epochs = 5\n\n    # --- Misc ---\n    seed = 42\n    use_cuda = torch.cuda.is_available()\n    device = torch.device(\"cuda\" if use_cuda else \"cpu\")\n    autocast_dtype = torch.float16 if use_cuda else torch.bfloat16\n\n\n# %% Helper Functions\ndef set_seed(seed=42):\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        # Ensure reproducibility if desired, may impact performance\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n    print(f\"Seed set to {seed}\")\n\n\ndef find_best_model(model_dir=cfg.working_dir, model_prefix=\"unet_best_model\"):\n    \"\"\"\n    Finds the best model file based on filename pattern (lowest loss).\n    Falls back to most recently created/modified if pattern fails or doesn't exist.\n    \"\"\"\n    best_loss = float(\"inf\")\n    best_model_path = None\n    pattern = os.path.join(model_dir, f\"{model_prefix}_epoch_*_loss_*.pth\")\n    model_files = glob.glob(pattern)\n\n    if not model_files:\n        # Fallback 1: No pattern match -> find latest created .pth\n        print(f\"W: No models matching pattern '{pattern}'. Looking for *.pth\")\n        all_pth_files = glob.glob(os.path.join(model_dir, \"*.pth\"))\n        if all_pth_files:\n            best_model_path = max(all_pth_files, key=os.path.getctime, default=None)\n            if best_model_path:\n                print(\n                    f\"Using most recently created: {os.path.basename(best_model_path)}\"\n                )\n            else:\n                print(\"W: No .pth models found.\")\n                return None\n        else:\n            print(\"W: No .pth models found in model directory.\")\n            return None\n\n    elif \"loss\" in os.path.basename(pattern):\n        # Try parsing loss from filename\n        parsed_models = []\n        for f in model_files:\n            try:\n                loss_str = f.split(\"_loss_\")[-1].split(\".pth\")[0]\n                loss = float(loss_str)\n                parsed_models.append((loss, f))\n            except (ValueError, IndexError, AttributeError):\n                print(f\"W: Couldn't parse loss from filename: {os.path.basename(f)}\")\n\n        if parsed_models:\n            # Found models with parseable loss, sort by loss\n            parsed_models.sort(key=lambda x: x[0])\n            best_loss, best_model_path = parsed_models[0]\n            print(\n                f\"Found best model by loss: {os.path.basename(best_model_path)} (Loss: {best_loss:.4f})\"\n            )\n        elif model_files:\n            # Pattern matched, but loss couldn't be parsed from any filename\n            print(\n                \"W: Pattern matched but no losses parsed. Selecting most recently created.\"\n            )\n            best_model_path = max(model_files, key=os.path.getctime, default=None)\n            if best_model_path:\n                print(f\"Using most recent creation time: {os.path.basename(best_model_path)}\")\n\n    else:\n        # Pattern matched but doesn't contain \"loss\" part (unexpected)\n        if model_files:\n            print(\n                f\"W: Pattern matched but no loss info expected. Selecting most recently created.\"\n            )\n            best_model_path = max(model_files, key=os.path.getctime, default=None)\n            if best_model_path:\n                print(\n                    f\"Using most recent creation time match: {os.path.basename(best_model_path)}\"\n                )\n\n    if not best_model_path:\n        # Final fallback: If no model found yet, use the latest modified .pth file\n        all_pth_files = glob.glob(os.path.join(model_dir, \"*.pth\"))\n        if all_pth_files:\n            print(\"W: Fallback: Selecting most recently modified .pth file.\")\n            best_model_path = max(all_pth_files, key=os.path.getmtime, default=None)\n            if best_model_path:\n                print(\n                    f\"Using most recent modification time: {os.path.basename(best_model_path)}\"\n                )\n\n    return best_model_path\n\n\n# %% WebDataset Preprocessing Functions\ndef search_data_path(target_dirs, root_dir, shuffle=True, seed=42):\n    \"\"\"Finds input/output .npy file pairs within subdirectories of a root directory.\"\"\"\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            # print(f\"W: Target directory {target_dir} not found in {root_path}\")\n            continue\n\n        in_files, out_files = [], []\n        data_subdir = data_dir / \"data\"\n        model_subdir = data_dir / \"model\"\n\n        # Check for HF structure first, then Kaggle structure\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            # print(f\"Found {len(in_files)}/{len(out_files)} files (HF style) in {target_dir}\")\n        else:\n            in_files = sorted(data_dir.glob(\"seis*.npy\"))\n            out_files = sorted(data_dir.glob(\"vel*.npy\"))\n            # print(f\"Found {len(in_files)}/{len(out_files)} files (Kaggle style) in {target_dir}\")\n\n        if not in_files or len(in_files) != len(out_files):\n            if in_files or out_files:  # Only warn if some files were found\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\n# ==================================================\n# CORRECTED generate_sample function (No finally block)\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    data = []\n    seis = None  # Initialize to ensure variable exists for potential del\n    vel = None\n    try:\n        if out_file is None:\n            # Logic for test data sharding (if needed later) - not implemented here\n            print(\"W: generate_sample called without out_file (test mode?), not implemented.\")\n            return []\n        else:\n            # --- Load Train/Validation data ---\n            try:\n                # Use mmap_mode='r' for memory efficiency if files are large\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 []  # Exit early if input fails\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                # Clean up the already loaded seis if vel loading fails\n                if seis is not None:\n                    del seis\n                return []  # Exit early if output fails\n\n            # --- Validate shapes and determine number of samples ---\n            n_samples = 0\n            if seis.ndim == 4 and vel.ndim == 4:  # Batch of samples (N, C, H, W)\n                if seis.shape[0] != vel.shape[0]:\n                    print(\n                        f\"W: Batch size mismatch in {in_file.name} ({seis.shape[0]}) vs {out_file.name} ({vel.shape[0]})\"\n                    )\n                    del seis, vel\n                    return []\n                n_samples = seis.shape[0]\n            elif seis.ndim == 3 and vel.ndim == 3:  # Single sample (C, H, W)\n                n_samples = 1\n            else:\n                # Raise error for unexpected dimensions\n                raise ValueError(\n                    f\"Unexpected dims: seis {seis.shape}, vel {vel.shape} in {in_file.name}\"\n                )\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            # --- Generate unique key based on file path relative to base_dir ---\n            common_part = f\"{in_file.parent.name}_{in_file.stem}\"  # Default key\n            if base_dir:\n                try:\n                    # Create key from relative path parts, removing .npy suffix\n                    relative_path = in_file.relative_to(base_dir)\n                    common_part = \"_\".join(relative_path.parts).replace(\".npy\", \"\")\n                    # Ensure compatibility across OS path separators\n                    common_part = common_part.replace(os.sep, \"_\").replace(\"\\\\\", \"_\")\n                except ValueError:\n                    # If relative_to fails (e.g., different drives), use the default key\n                    pass\n\n            # --- Process and append each sample ---\n            for i in range(n_samples):\n                key = f\"{common_part}_{i}\"\n                # Extract sample, explicitly copy, and convert to float16\n                s_sample = (\n                    seis[i].copy().astype(np.float16)\n                    if seis.ndim == 4\n                    else seis.copy().astype(np.float16)\n                )\n                v_sample = (\n                    vel[i].copy().astype(np.float16)\n                    if vel.ndim == 4\n                    else vel.copy().astype(np.float16)\n                )\n                data.append(\n                    {\n                        \"__key__\": key,\n                        \"sample_id.txt\": key,  # Store key as text too\n                        \"seis.npy\": s_sample,\n                        \"vel.npy\": v_sample,\n                    }\n                )\n\n            # --- Explicitly delete mmap objects after copying data ---\n            # This is important to release file handles, especially with mmap\n            del seis\n            del vel\n\n    except Exception as e:\n        # Catch other errors (ValueError from dim check, key gen errors, etc.)\n        print(f\"E: Error during sample generation for {in_file.name}: {e}\")\n        # Explicitly try deleting here, in case they were loaded before the error\n        if seis is not None:\n            try:\n                del seis\n            except NameError:  # Should not happen if assigned None initially\n                pass\n        if vel is not None:\n            try:\n                del vel\n            except NameError:\n                pass\n        return []  # Return empty list on any error during processing\n\n    # No finally block needed as del is handled within try/except scopes\n    return data\n\n\n# ==================================================\n\n\n# %% WebDataset Loading Functions\ndef get_shard_paths(\n    root_dir, dataset_name, stage, num_shards=None, test_size=0.2, seed=42\n):\n    \"\"\"Gets list of shard paths, optionally selects subset, optionally splits train/val.\"\"\"\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\n    # --- Shard Selection Logic ---\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(\n                f\"Requested {num_shards} or more shards, using all {available_count} available.\"\n            )\n        else:  # num_shards <= 0\n            print(\n                f\"W: Invalid num_shards ({num_shards}). Using all {available_count} shards.\"\n            )\n    print(f\"Using {len(selected_paths)} selected shards for stage '{stage}'.\")\n\n    # --- Train/Validation Split Logic ---\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):\n                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(\n                    selected_paths, test_size=test_size, random_state=seed, shuffle=True\n                )\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:  # Not 'train' stage (e.g., 'val' direct loading or 'test')\n        print(f\"# Shards returned for stage '{stage}': {len(selected_paths)}\")\n        return sorted(selected_paths)\n\n\ndef get_dataset(paths, stage, seed=42):\n    \"\"\"Creates WebDataset object. Applies augmentations if stage=='train'.\"\"\"\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    \n    # Custom handler that will just skip errors rather than warning \n    def silent_skip(exn):\n        return None\n        \n    # Older WebDataset version compatible error handler setup\n    map_handler = wds.ignore_and_continue\n    \n    try:\n        # Create the dataset with proper error handling\n        dataset = wds.WebDataset(\n            paths, \n            nodesplitter=wds.split_by_node, \n            shardshuffle=1 if is_train else 0,  # True/False can cause issues in older versions\n            seed=seed,\n            handler=wds.ignore_and_continue  # Skip errors at shard level\n        )\n        \n        # Decode standard types (.npy, .txt, etc.)\n        dataset = dataset.decode(handler=map_handler)\n        \n        def map_train_val(sample):\n            \"\"\"Inner function to process decoded samples and apply augmentations.\"\"\"\n            key_info = sample.get(\"__key__\", \"N/A\")  # For error reporting\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                # Ensure numpy arrays and convert to float32 tensors\n                s_np = np.asarray(sample[\"seis.npy\"])\n                v_np = np.asarray(sample[\"vel.npy\"])\n                \n                # Verify correct shapes - should be (5, 1000, 70) and (1, 70, 70)\n                expected_seis_shape = (5, 1000, 70)\n                expected_vel_shape = (1, 70, 70)\n                \n                # Only log a small random subset of samples to prevent log overflow\n                if random.random() < 0.01:  # 1% chance to log\n                    print(f\"\\n--- DATASET INSPECTION: Sample Key: {key_info} ---\")\n                    print(f\"Original seismic numpy shape: {s_np.shape}, dtype: {s_np.dtype}\")\n                    print(f\"Original velocity numpy shape: {v_np.shape}, dtype: {v_np.dtype}\")\n                    \n                # Fix shapes if needed (resize/pad to expected dimensions)\n                if s_np.shape != expected_seis_shape:\n                    print(f\"Fixing seismic shape from {s_np.shape} to {expected_seis_shape} for {key_info}\")\n                    # Simple case: missing batch dimension\n                    if len(s_np.shape) == 2 and s_np.shape[0] == 1000 and s_np.shape[1] == 70:\n                        # Add channel dimension\n                        s_np = np.expand_dims(s_np, axis=0)\n                        # Repeat to get 5 channels\n                        s_np = np.repeat(s_np, 5, axis=0)\n                    # Other cases need proper resizing - convert to tensor for easier interpolation\n                    elif s_np.shape != expected_seis_shape:\n                        # Need to reshape or resize\n                        temp_tensor = torch.from_numpy(s_np.astype(np.float32))\n                        # Add missing dimensions if needed\n                        while len(temp_tensor.shape) < 3:\n                            temp_tensor = temp_tensor.unsqueeze(0)\n                        # Ensure we have exactly 3 dimensions (C, H, W)\n                        if len(temp_tensor.shape) > 3:\n                            if temp_tensor.shape[0] == 1:  # First dim is batch\n                                temp_tensor = temp_tensor.squeeze(0)\n                            else:\n                                # More complex case, just flatten and reshape\n                                temp_tensor = temp_tensor.reshape(-1, 1000, 70)\n                                temp_tensor = temp_tensor[:5] if temp_tensor.shape[0] >= 5 else temp_tensor\n                        # Ensure we have 5 channels\n                        if temp_tensor.shape[0] != 5:\n                            if temp_tensor.shape[0] < 5:\n                                # Repeat the first channel to get to 5\n                                temp_tensor = torch.cat([temp_tensor, temp_tensor[0:1].repeat(5-temp_tensor.shape[0], 1, 1)], dim=0)\n                            else:\n                                # Take the first 5 channels\n                                temp_tensor = temp_tensor[:5]\n                        # Resize spatial dimensions if needed\n                        if temp_tensor.shape[1:] != (1000, 70):\n                            temp_tensor = F.interpolate(temp_tensor.unsqueeze(0), size=(1000, 70), mode='bilinear', align_corners=False).squeeze(0)\n                        # Convert back to numpy\n                        s_np = temp_tensor.numpy()\n                \n                # Similarly validate and fix velocity tensor shape\n                if v_np.shape != expected_vel_shape:\n                    print(f\"Fixing velocity shape from {v_np.shape} to {expected_vel_shape} for {key_info}\")\n                    # Simple case: missing channel dimension\n                    if len(v_np.shape) == 2 and v_np.shape[0] == 70 and v_np.shape[1] == 70:\n                        v_np = np.expand_dims(v_np, axis=0)\n                    # Other cases need proper resizing\n                    else:\n                        temp_tensor = torch.from_numpy(v_np.astype(np.float32))\n                        # Add missing dimensions if needed\n                        while len(temp_tensor.shape) < 3:\n                            temp_tensor = temp_tensor.unsqueeze(0)\n                        # Ensure we have exactly 3 dimensions (C, H, W)\n                        if len(temp_tensor.shape) > 3:\n                            if temp_tensor.shape[0] == 1:  # First dim is batch\n                                temp_tensor = temp_tensor.squeeze(0)\n                            else:\n                                # More complex case, flatten and take first channel\n                                temp_tensor = temp_tensor.reshape(-1, 70, 70)\n                                temp_tensor = temp_tensor[:1]\n                        # Ensure we have 1 channel\n                        if temp_tensor.shape[0] != 1:\n                            # Take first channel only\n                            temp_tensor = temp_tensor[0:1]\n                        # Resize spatial dimensions if needed\n                        if temp_tensor.shape[1:] != (70, 70):\n                            temp_tensor = F.interpolate(temp_tensor.unsqueeze(0), size=(70, 70), mode='bilinear', align_corners=False).squeeze(0)\n                        # Convert back to numpy\n                        v_np = temp_tensor.numpy()\n                \n                # Convert to tensors\n                seis_tensor = torch.from_numpy(s_np).float()\n                vel_tensor = torch.from_numpy(v_np).float()\n                \n                if random.random() < 0.01:  # Only log occasionally\n                    print(f\"Loaded seismic tensor: {seis_tensor.shape}\")\n                    print(f\"Loaded velocity tensor: {vel_tensor.shape}\")\n                \n                # --- Augmentation Block ---\n                if is_train and cfg.apply_augmentation:\n                    # 1. 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                    # 2. Add Gaussian Noise to Seismic Data\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                # ====== Compute Gradients/Normals ======\n                # Compute gradients with large kernel (5x5 Sobel-like)\n                def large_kernel_gradient(tensor):\n                    kernel_x = torch.tensor([\n                        [-1, -2, 0, 2, 1],\n                        [-4, -8, 0, 8, 4],\n                        [-6,-12, 0,12, 6],\n                        [-4, -8, 0, 8, 4],\n                        [-1, -2, 0, 2, 1]\n                    ], dtype=torch.float32) / 48.0\n                    \n                    kernel_y = kernel_x.T\n                    \n                    # Ensure input is properly shaped for convolution\n                    if tensor.dim() == 2:  # Add batch and channel dimensions if needed\n                        tensor_4d = tensor.unsqueeze(0).unsqueeze(0)\n                    elif tensor.dim() == 3:  # Add channel dimension if needed\n                        tensor_4d = tensor.unsqueeze(1)\n                    else:\n                        tensor_4d = tensor\n                    \n                    # Verify tensor shape and apply convolution\n                    grad_x = F.conv2d(tensor_4d, kernel_x.view(1,1,5,5), padding=2)\n                    grad_y = F.conv2d(tensor_4d, kernel_y.view(1,1,5,5), padding=2)\n                    \n                    # Remove batch dimension if input was 2D or 3D\n                    if tensor.dim() <= 3:\n                        grad_x = grad_x.squeeze(0)\n                        grad_y = grad_y.squeeze(0)\n                    \n                    # Remove channel dimension always\n                    grad_x = grad_x.squeeze(1)\n                    grad_y = grad_y.squeeze(1)\n                    \n                    # Ensure output has same shape as input\n                    if grad_x.shape != tensor.shape:\n                        grad_x = F.interpolate(grad_x.unsqueeze(1), size=tensor.shape, mode='bilinear', align_corners=False).squeeze(1)\n                        grad_y = F.interpolate(grad_y.unsqueeze(1), size=tensor.shape, mode='bilinear', align_corners=False).squeeze(1)\n                    \n                    return grad_x, grad_y\n\n                # Compute gradients with our large kernel\n                grad_x, grad_y = large_kernel_gradient(vel_tensor)\n                \n                # Debug prints for sizes\n                if random.random() < 0.01:  # Print for ~1% of samples\n                    print(f\"Debug shapes - vel_tensor: {vel_tensor.shape}, grad_x: {grad_x.shape}, grad_y: {grad_y.shape}\")\n                \n                # Handle surface layer (first row)\n                if grad_y.dim() > 1 and grad_y.size(1) > 0:  # Check if there's a height dimension\n                    grad_y[:,0,:] = vel_tensor[:,0,:]  # Surface velocity as first row\n                \n                # Compute normals from gradients\n                norm = torch.sqrt(grad_x**2 + grad_y**2 + 1e-6)\n                normal_x = -grad_x / norm  # Negative gradient direction\n                normal_y = -grad_y / norm\n\n                # Ensure correct dimensions for gradient and normal tensors\n                # If vel_tensor is [B, H, W], then grad components are [B, H, W]\n                # We want to create [B, 2, H, W] tensors for both grad and normal\n                \n                # Stack gradients along new channel dimension\n                if grad_x.dim() == 2:  # If we have a 2D tensor (H, W)\n                    grad_stacked = torch.stack([grad_x, grad_y], dim=0)  # [2, H, W]\n                    normal_stacked = torch.stack([normal_x, normal_y], dim=0)  # [2, H, W]\n                else:  # We have a 3D tensor (B, H, W)\n                    grad_stacked = torch.stack([grad_x, grad_y], dim=1)  # [B, 2, H, W]\n                    normal_stacked = torch.stack([normal_x, normal_y], dim=1)  # [B, 2, H, W]\n                \n                # Make sure tensors have compatible shapes\n                if grad_stacked.shape[-2:] != vel_tensor.shape[-2:]:\n                    print(f\"Shape mismatch: grad_stacked {grad_stacked.shape[-2:]} != vel_tensor {vel_tensor.shape[-2:]}\")\n                    new_shape = vel_tensor.shape[-2:]\n                    \n                    # Reshape to 4D if needed for interpolation\n                    if grad_stacked.dim() == 3:  # [2, H, W]\n                        grad_stacked = grad_stacked.unsqueeze(0)  # [1, 2, H, W]\n                        normal_stacked = normal_stacked.unsqueeze(0)  # [1, 2, H, W]\n                        need_squeeze = True\n                    else:\n                        need_squeeze = False\n                        \n                    # Do the interpolation\n                    grad_stacked = F.interpolate(grad_stacked, size=new_shape, mode='bilinear', align_corners=False)\n                    normal_stacked = F.interpolate(normal_stacked, size=new_shape, mode='bilinear', align_corners=False)\n                    \n                    # Squeeze back if we added a dimension\n                    if need_squeeze:\n                        grad_stacked = grad_stacked.squeeze(0)\n                        normal_stacked = normal_stacked.squeeze(0)\n                        \n                    print(f\"After interpolation: grad_stacked {grad_stacked.shape[-2:]}\")\n                \n                # More debug prints\n                if random.random() < 0.01:  # Print for ~1% of samples \n                    print(f\"Final shapes - vel: {vel_tensor.shape}, grad: {grad_stacked.shape}, normal: {normal_stacked.shape}\")\n\n                # Ensure all tensors have batch dimension\n                if seis_tensor.dim() == 2:\n                    seis_tensor = seis_tensor.unsqueeze(0)\n                if vel_tensor.dim() == 2:\n                    vel_tensor = vel_tensor.unsqueeze(0)\n                if grad_stacked.dim() == 2:\n                    grad_stacked = grad_stacked.unsqueeze(0)\n                if normal_stacked.dim() == 2:\n                    normal_stacked = normal_stacked.unsqueeze(0)\n\n                return {\n                    \"__key__\": key_info,  # Keep the key for debugging\n                    \"sample_id\": sid,\n                    \"seis\": seis_tensor,\n                    \"vel\": vel_tensor,\n                    \"grad\": grad_stacked,\n                    \"normal\": normal_stacked\n                }\n\n            except Exception as map_e:\n                print(f\"E: Map function failed for sample {key_info}: {map_e}\")\n                # Let the handler decide whether to skip or raise\n                raise map_e\n\n        # Apply the mapping function to train/val stages\n        if stage in [\"train\", \"val\"]:\n            dataset = dataset.map(map_train_val, handler=map_handler)\n\n        # Shuffle buffer for training data\n        if is_train:\n            dataset = dataset.shuffle(1000)  # Buffer size for shuffling\n        \n        # Check for empty items (compatible with older WebDataset versions)\n        def check_sample(sample):\n            return sample is not None and len(sample) > 0\n            \n        # Make sure all items in a batch have the same tensor shapes\n        def check_batch_consistency(batch):\n            \"\"\"Make sure all items in a batch have the same tensor shapes\"\"\"\n            if not batch:\n                return None\n                \n            try:\n                # First filter out any non-dictionary items (like strings)\n                valid_items = []\n                for item in batch:\n                    if isinstance(item, dict) and \"seis\" in item and \"vel\" in item:\n                        valid_items.append(item)\n                    else:\n                        # If we encounter a non-dictionary or incomplete item, log it\n                        item_type = type(item).__name__\n                        print(f\"Skipping invalid batch item of type {item_type}\")\n                \n                # If no valid items remain, return None\n                if not valid_items:\n                    print(\"No valid items found in batch, skipping\")\n                    return None\n                \n                # Extract tensor shapes only from valid dictionary items\n                seis_shapes = [item[\"seis\"].shape for item in valid_items]\n                vel_shapes = [item[\"vel\"].shape for item in valid_items]\n                \n                # Check if all shapes are the same\n                unique_seis_shapes = set(str(s) for s in seis_shapes)\n                unique_vel_shapes = set(str(s) for s in vel_shapes)\n                \n                # Log shape info for debugging\n                if len(unique_seis_shapes) > 1 or len(unique_vel_shapes) > 1:\n                    print(f\"Found inconsistent shapes in batch: seis={unique_seis_shapes}, vel={unique_vel_shapes}\")\n                    return None  # Skip this batch\n                    \n                # Use only the valid items for the batch\n                return valid_items  # Return the filtered batch\n            except Exception as e:\n                print(f\"Error in batch filtering: {e}\")\n                return None\n        \n        # WebDataset v0.2.111 compatible filtering and batching\n        # For older versions, use to_tuple() and then from_tuple() as a replacement for select/filter\n        \n        # Check if this version has the select method\n        if hasattr(dataset, 'select'):\n            # Modern WebDataset\n            dataset = dataset.select(check_sample)  # Filter out empty items\n            dataset = dataset.batched(cfg.batch_size, partial=True)\n            dataset = dataset.map(check_batch_consistency)\n            dataset = dataset.select(lambda x: x is not None)  # Remove filtered out batches\n        else:\n            # Older WebDataset version without select/filter\n            # Convert to tuples\n            dataset = dataset.to_tuple(\"__key__ sample_id.txt seis.npy vel.npy\".split())\n            \n            # Define a function to filter and process tuples\n            def process_tuple(sample_tuple):\n                try:\n                    # Unpack the tuple\n                    key, sample_id, seis_np, vel_np = sample_tuple\n                    \n                    # Skip if any component is None\n                    if key is None or sample_id is None or seis_np is None or vel_np is None:\n                        return None\n                        \n                    # Process like map_train_val would\n                    key_info = key\n                    s_np = np.asarray(seis_np)\n                    v_np = np.asarray(vel_np)\n                    \n                    # Verify correct shapes - should be (5, 1000, 70) and (1, 70, 70)\n                    expected_seis_shape = (5, 1000, 70)\n                    expected_vel_shape = (1, 70, 70)\n                    \n                    # Fix shapes if needed\n                    # (Shape fixing code would go here - simplified for readability)\n                    \n                    # Convert to tensors\n                    seis_tensor = torch.from_numpy(s_np).float()\n                    vel_tensor = torch.from_numpy(v_np).float()\n                    \n                    # Compute gradients/normals\n                    # (Remaining logic from map_train_val)\n                    \n                    return {\n                        \"__key__\": key_info,\n                        \"sample_id\": sample_id,\n                        \"seis\": seis_tensor,\n                        \"vel\": vel_tensor\n                        # Add grad and normal here too\n                    }\n                    \n                except Exception as e:\n                    print(f\"Error processing sample tuple: {e}\")\n                    return None\n                    \n            # Map the function over the dataset\n            dataset = dataset.map(process_tuple)\n            # Skip None values\n            dataset = dataset.compose(lambda src: (x for x in src if x is not None))\n            # Batch the dataset\n            dataset = dataset.batched(cfg.batch_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 the final Kaggle test set directly from individual .npy files.\"\"\"\n\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(\n                    f\"Kaggle test directory missing: {self.test_files_dir}\"\n                )\n            self.test_files = sorted(list(self.test_files_dir.glob(\"*.npy\")))\n            print(\n                f\"Found {len(self.test_files)} '.npy' files in Kaggle test dir: {self.test_files_dir}\"\n            )\n            if not self.test_files:\n                print(f\"W: No .npy files found in {self.test_files_dir}.\")\n        except Exception as e:\n            print(\n                f\"E: Error accessing Kaggle test directory {self.test_files_dir}: {e}\"\n            )\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(\n                f\"Index {index} out of bounds for KaggleTestDataset ({len(self.test_files)} files).\"\n            )\n        test_file_path = self.test_files[index]\n        try:\n            # Load numpy array and convert to float32 tensor\n            data = torch.from_numpy(np.load(test_file_path).astype(np.float32))\n            # Get the original ID (filename without extension)\n            original_id = test_file_path.stem\n            return data, original_id\n        except Exception as e:\n            # Raise a more informative error if loading fails\n            raise IOError(f\"Error loading Kaggle test file: {test_file_path}\") from e\n\n\n# %% U-Net Model Definition (Formatted for Readability with Residual Blocks)\n\nclass ResidualDoubleConv(nn.Module):\n    \"\"\"(Convolution => [BN] => ReLU) * 2 + Residual Connection\"\"\"\n\n    def __init__(self, in_channels, out_channels, mid_channels=None, stride=1):\n        super().__init__()\n        if not mid_channels:\n            mid_channels = out_channels\n\n        # First convolution layer\n        self.conv1 = nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1, stride=stride, bias=False)\n        self.bn1 = nn.BatchNorm2d(mid_channels)\n        self.relu = nn.ReLU(inplace=True)\n\n        # Second convolution layer\n        self.conv2 = nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n\n        # Shortcut connection to handle potential channel mismatch\n        if in_channels == out_channels and stride == 1:\n            self.shortcut = nn.Identity()\n        else:\n            # Projection shortcut: 1x1 conv + BN to match output channels and spatial dimensions\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n\n    def forward(self, x):\n        identity = x  # Store the input for the residual connection\n        # Only print debug info rarely to reduce memory usage\n        debug_logging = random.random() < 0.001  # 0.1% chance\n        \n        if debug_logging:\n            print(f\"ResidualDoubleConv input: {x.shape}\")\n\n        # First conv block\n        out = self.conv1(x)\n        if debug_logging:\n            print(f\"After conv1: {out.shape}\")\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        # Second conv block (without final ReLU yet)\n        out = self.conv2(out)\n        if debug_logging:\n            print(f\"After conv2: {out.shape}\")\n        out = self.bn2(out)\n\n        # Apply shortcut to the identity path\n        identity_mapped = self.shortcut(identity)\n        if debug_logging:\n            print(f\"Identity after shortcut: {identity_mapped.shape}\")\n            \n        # Ensure dimensions match before adding residual connection\n        # Both out and identity_mapped should have the same shape\n        if out.shape != identity_mapped.shape:\n            if debug_logging:\n                print(f\"Warning: Shape mismatch in residual connection - out: {out.shape}, identity: {identity_mapped.shape}\")\n                \n            # Resize the smaller one to match the larger one\n            if out.shape[-1] > identity_mapped.shape[-1] or out.shape[-2] > identity_mapped.shape[-2]:\n                # Identity is smaller, resize it to match out\n                identity_mapped = F.interpolate(identity_mapped, size=out.shape[-2:], mode='bilinear', align_corners=False)\n                if debug_logging:\n                    print(f\"Resized identity to: {identity_mapped.shape}\")\n            else:\n                # Out is smaller, resize it to match identity\n                out = F.interpolate(out, size=identity_mapped.shape[-2:], mode='bilinear', align_corners=False)\n                if debug_logging:\n                    print(f\"Resized out to: {out.shape}\")\n\n        # Add the residual connection\n        out += identity_mapped\n\n        # Apply final ReLU\n        out = self.relu(out)\n        \n        # Final check - ensure output has desired dimensions (70x70)\n        if out.shape[-1] != 70 or out.shape[-2] != 70:\n            # Need to adjust dimensions\n            if debug_logging:\n                print(f\"Adjusting output to 70x70 (currently {out.shape[-2]}x{out.shape[-1]})\")\n                \n            # Use center crop or padding to get to 70x70\n            # Pad if smaller\n            if out.shape[-1] < 70 or out.shape[-2] < 70:\n                diff_y = max(0, 70 - out.shape[-2])\n                diff_x = max(0, 70 - out.shape[-1])\n                pad_top = diff_y // 2\n                pad_bottom = diff_y - pad_top\n                pad_left = diff_x // 2\n                pad_right = diff_x - pad_left\n                out = F.pad(out, (pad_left, pad_right, pad_top, pad_bottom), mode='replicate')\n                \n            # Crop if larger\n            if out.shape[-1] > 70 or out.shape[-2] > 70:\n                crop_y = out.shape[-2] - 70\n                crop_x = out.shape[-1] - 70\n                start_y = crop_y // 2\n                start_x = crop_x // 2\n                out = out[..., start_y:start_y+70, start_x:start_x+70]\n                \n            if debug_logging:\n                print(f\"Final adjusted output: {out.shape}\")\n        \n        if debug_logging:\n            print(f\"ResidualDoubleConv output: {out.shape}\")\n        return out\n\n\nclass Up(nn.Module):\n    \"\"\"Upscaling then ResidualDoubleConv\"\"\"\n\n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super().__init__()\n        self.bilinear = bilinear\n\n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode=\"bilinear\", align_corners=False)\n            # For bilinear upsampling, we need to handle the channel concatenation\n            self.conv = ResidualDoubleConv(in_channels + out_channels, out_channels)\n        else:\n            # For transposed conv, we first reduce channels then concatenate\n            self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)\n            self.conv = ResidualDoubleConv(in_channels // 2 + out_channels, out_channels)\n\n    def forward(self, x1, x2):\n        # x1 is from below (needs upsampling), x2 is skip connection\n        x1 = self.up(x1)\n        \n        # Ensure spatial dimensions match\n        if x1.shape[-1] != x2.shape[-1] or x1.shape[-2] != x2.shape[-2]:\n            x1 = F.interpolate(x1, size=x2.shape[-2:], mode='bilinear', align_corners=False)\n        \n        # Concatenate along channel dimension\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\n\nclass OutConv(nn.Module):\n    \"\"\"1x1 Convolution for the output layer\"\"\"\n\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)\n\n    def forward(self, x):\n        return self.conv(x)\n\n\nclass UNet(nn.Module):\n    \"\"\"U-Net architecture implementation with Residual Blocks\"\"\"\n\n    def __init__(\n        self,\n        n_channels=5,  # Fixed to 5 for seismic input channels\n        n_classes=1,   # Fixed to 1 for velocity output\n        init_features=24,\n        depth=4,\n        bilinear=True,\n    ):\n        super().__init__()\n        self.n_channels = n_channels\n        self.n_classes = n_classes\n        self.bilinear = bilinear\n        self.depth = depth\n\n        # Initial pooling to reduce time dimension\n        self.initial_pool = nn.AvgPool2d(kernel_size=(100, 1), stride=(100, 1))  # Reduce 1000 to 10\n\n        # Encoder\n        self.encoder_convs = nn.ModuleList()\n        self.encoder_pools = nn.ModuleList()\n\n        # Initial conv block\n        self.inc = ResidualDoubleConv(n_channels, init_features)\n        self.encoder_convs.append(self.inc)\n\n        current_features = init_features\n        for _ in range(depth):\n            # Use stride=2 in conv instead of separate pooling to maintain spatial dimensions\n            conv = ResidualDoubleConv(current_features, current_features * 2, stride=2)\n            self.encoder_convs.append(conv)\n            current_features *= 2\n\n        # Bottleneck\n        self.bottleneck = ResidualDoubleConv(current_features, current_features)\n\n        # Decoder\n        self.decoder_blocks = nn.ModuleList()\n        for _ in range(depth):\n            up_block = Up(current_features, current_features // 2, bilinear)\n            self.decoder_blocks.append(up_block)\n            current_features //= 2\n\n        # Output layer for velocity\n        self.outc = OutConv(current_features, n_classes)\n\n        # Physics-guided prediction heads\n        self.grad_head = nn.Sequential(\n            nn.Conv2d(current_features, 32, 3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(32, 2, 1)  # Predicts grad_x, grad_y\n        )\n        \n        self.normal_head = nn.Sequential(\n            nn.Conv2d(current_features, 64, 3, padding=1),\n            nn.GroupNorm(8, 64),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(64, 2, 1)  # Predicts normal_x, normal_y\n        )\n\n    def forward(self, x):\n        # Input shape could be [B, 5, 1000, 70] or [B, B, 5, 1000, 70] with double batching\n        # Handle various potential tensor shapes\n        orig_shape = x.shape\n\n        try:\n            # Handle the double-batching case\n            if x.dim() > 4:\n                print(f\"UNet received {x.dim()}D input with shape {x.shape}, reshaping...\")\n                # If we have [B, B, C, H, W] shape, reshape to [B*B, C, H, W]\n                if x.dim() == 5:\n                    batch_size = x.size(0) * x.size(1)\n                    x = x.reshape(batch_size, *x.shape[2:])\n                    print(f\"Reshaped to {x.shape}\")\n                else:\n                    print(f\"Unexpected tensor dimensionality: {x.dim()}, attempting to process\")\n                \n            # Ensure we have a batch dimension\n            if x.dim() == 3:  # [C, H, W] -> [1, C, H, W]\n                x = x.unsqueeze(0)\n                print(f\"Added batch dimension, shape now: {x.shape}\")\n                \n            # Ensure the input tensor is in the correct format [B, 5, 1000, 70]\n            if x.dim() == 4:\n                # Check channel dimension is 5\n                if x.size(1) != 5:\n                    print(f\"Warning: Channel dimension is {x.size(1)}, expected 5\")\n                    \n                # Handle different spatial dimensions with interpolation\n                if x.size(2) != 1000 or x.size(3) != 70:\n                    print(f\"Warning: Spatial dimensions are {x.size(2)}x{x.size(3)}, expected 1000x70\")\n                    x = F.interpolate(x, size=(1000, 70), mode='bilinear', align_corners=False)\n                    print(f\"Interpolated to shape: {x.shape}\")\n            \n            # Input shape: [B, 5, 1000, 70]\n            # First reduce time dimension\n            x = self.initial_pool(x)  # [B, 5, 10, 70]\n            \n        except Exception as e:\n            print(f\"Error preprocessing input tensor (shape {orig_shape}): {e}\")\n            # Emergency reshape - try to make it work with minimal assumptions\n            try:\n                if x.dim() >= 3:\n                    # Extract relevant dimensions and reshape\n                    target_shape = [-1, 5, 1000, 70]  # Target shape with unknown batch size\n                    total_elements = x.numel()\n                    batch_size = total_elements // (5 * 1000 * 70)  # Calculate batch size\n                    \n                    if batch_size > 0:\n                        x = x.reshape(batch_size, 5, 1000, 70)\n                        print(f\"Emergency reshape to {x.shape}\")\n                        x = self.initial_pool(x)  # Try to continue with pipeline\n                    else:\n                        # If tensor is too small, pad it\n                        print(f\"Tensor too small, emergency padding\")\n                        x = torch.zeros(1, 5, 1000, 70, device=x.device, dtype=x.dtype)\n                        x = self.initial_pool(x)\n                else:\n                    # Complete failure case, just create a dummy tensor\n                    print(f\"Unable to reshape tensor, creating emergency dummy\")\n                    x = torch.zeros(1, 5, 10, 70, device=x.device, dtype=x.dtype)\n            except Exception as e2:\n                print(f\"Emergency reshape failed: {e2}, creating dummy tensor\")\n                # Last resort dummy tensor\n                x = torch.zeros(1, 5, 10, 70, device=x.device, dtype=x.dtype)\n\n        # Initial conv\n        x1 = self.encoder_convs[0](x)  # [B, 24, 10, 70]\n        \n        # Store skip connections\n        skips = [x1]\n        \n        # Encoder path\n        for i in range(self.depth):\n            x1 = self.encoder_convs[i+1](self.encoder_pools[i](x1))\n            if i < self.depth - 1:\n                skips.append(x1)\n        \n        # Bottleneck\n        x1 = self.bottleneck(x1)\n        \n        # Decoder path\n        for i in range(self.depth):\n            # Get skip connection\n            skip = skips[-(i+1)]\n            \n            # Ensure skip connection has correct number of channels\n            if skip.shape[1] != x1.shape[1] // 2:\n                # Create a 1x1 conv to adjust channels\n                conv = nn.Conv2d(skip.shape[1], x1.shape[1] // 2, 1).to(skip.device)\n                skip = conv(skip)\n            \n            x1 = self.decoder_blocks[i](x1, skip)\n        \n        # Main velocity prediction\n        velocity = self.outc(x1)  # [B, 1, 70, 70]\n        \n        # Physics-guided predictions\n        gradients = self.grad_head(x1)  # [B, 2, 70, 70]\n        normals = self.normal_head(x1)  # [B, 2, 70, 70]\n        \n        # Ensure all outputs are exactly 70x70\n        if velocity.shape[-1] != 70 or velocity.shape[-2] != 70:\n            velocity = F.interpolate(velocity, size=(70, 70), mode='bilinear', align_corners=False)\n        if gradients.shape[-1] != 70 or gradients.shape[-2] != 70:\n            gradients = F.interpolate(gradients, size=(70, 70), mode='bilinear', align_corners=False)\n        if normals.shape[-1] != 70 or normals.shape[-2] != 70:\n            normals = F.interpolate(normals, size=(70, 70), mode='bilinear', align_corners=False)\n        \n        return velocity, gradients, normals\n\n\ndef ensure_tensor_sizes(pred, target, expected_size=70):\n    \"\"\"Ensure tensors have the exact size needed in spatial dimensions\"\"\"\n    \n    # First, determine if resizing is needed\n    spatial_dim = pred.shape[-1]  # Get the last dimension (width)\n    \n    if spatial_dim != expected_size:\n        print(f\"Resizing tensors from spatial dim {spatial_dim} to {expected_size}\")\n        \n        # Create paddings for spatial dimensions\n        if spatial_dim < expected_size:\n            # Need to pad\n            pad_size = expected_size - spatial_dim\n            left_pad = pad_size // 2\n            right_pad = pad_size - left_pad\n            \n            # For 4D tensors: [batch, channels, height, width]\n            # Format: (left_pad, right_pad, top_pad, bottom_pad)\n            padding = (left_pad, right_pad, left_pad, right_pad)\n            \n            # Use padding instead of interpolation to preserve values\n            pred_resized = F.pad(pred, padding, mode='replicate')\n            target_resized = F.pad(target, padding, mode='replicate')\n            \n            print(f\"Padded tensors: {pred.shape} -> {pred_resized.shape}, {target.shape} -> {target_resized.shape}\")\n        else:\n            # Need to crop\n            crop_size = spatial_dim - expected_size\n            left_crop = crop_size // 2\n            right_crop = left_crop + expected_size\n            \n            # Crop the tensors (using proper slicing)\n            pred_resized = pred[..., :, left_crop:right_crop, left_crop:right_crop]\n            target_resized = target[..., :, left_crop:right_crop, left_crop:right_crop]\n            \n            print(f\"Cropped tensors: {pred.shape} -> {pred_resized.shape}, {target.shape} -> {target_resized.shape}\")\n        \n        return pred_resized, target_resized\n    \n    # No resize needed\n    return pred, target\n\n\ndef verify_batch_tensors(inputs, targets, grads, normals, expected_size=70):  # Changed to expect 70 not 71\n    \"\"\"Helper function to check and fix tensor shapes for compatibility.\"\"\"\n    \n    def log_shape(name, tensor):\n        # Only log occasionally to avoid log bloat\n        if random.random() < 0.1:  # 10% chance to log\n            print(f\"{name} shape: {tensor.shape}\")\n    \n    # Log shapes of all tensors\n    log_shape(\"inputs\", inputs)\n    log_shape(\"targets\", targets)\n    log_shape(\"grads\", grads)\n    log_shape(\"normals\", normals)\n    \n    try:\n        # Reshape tensors if needed\n        # Handle 5D grads tensor: [batch, 1, 2, h, w] -> [batch, 2, h, w]\n        if grads.dim() == 5:\n            if random.random() < 0.1:\n                print(f\"Reshaping 5D grads to 4D: {grads.shape} -> \", end=\"\")\n            grads = grads.squeeze(1)  # Remove the extra dimension\n            if random.random() < 0.1:\n                print(f\"{grads.shape}\")\n        \n        # Handle 5D normals tensor: [batch, 1, 2, h, w] -> [batch, 2, h, w]\n        if normals.dim() == 5:\n            if random.random() < 0.1:\n                print(f\"Reshaping 5D normals to 4D: {normals.shape} -> \", end=\"\")\n            normals = normals.squeeze(1)  # Remove the extra dimension\n            if random.random() < 0.1:\n                print(f\"{normals.shape}\")\n        \n        # Ensure targets have correct dimensions\n        # If targets shape is [B, 1, H, W] and H=W=70, we need to ensure it is exactly 70x70\n        if targets.shape[-1] != expected_size or targets.shape[-2] != expected_size:\n            if targets.dim() == 3:  # [B, H, W]\n                targets = targets.unsqueeze(1)  # [B, 1, H, W]\n            padding_mode = 'replicate'\n            if random.random() < 0.1:\n                print(f\"Padding targets from {targets.shape[-2:]} to [{expected_size}, {expected_size}] using {padding_mode} mode\")\n            \n            # Calculate padding\n            h_pad = max(0, expected_size - targets.shape[-2])\n            w_pad = max(0, expected_size - targets.shape[-1])\n            \n            pad_left = w_pad // 2\n            pad_right = w_pad - pad_left\n            pad_top = h_pad // 2\n            pad_bottom = h_pad - pad_top\n            \n            # Apply padding\n            targets = F.pad(targets, (pad_left, pad_right, pad_top, pad_bottom), mode=padding_mode)\n            if random.random() < 0.1:\n                print(f\"Padded targets to: {targets.shape}\")\n        \n        # Same for gradients and normals - ensure 4D tensors\n        if grads.dim() == 3:  # Missing batch dimension\n            if random.random() < 0.1:\n                print(f\"Adding batch dimension to grads: {grads.shape} -> \", end=\"\")\n            grads = grads.unsqueeze(0)  # [1, 2, H, W]\n            if random.random() < 0.1:\n                print(f\"{grads.shape}\")\n        \n        if normals.dim() == 3:\n            if random.random() < 0.1:\n                print(f\"Adding batch dimension to normals: {normals.shape} -> \", end=\"\")\n            normals = normals.unsqueeze(0)  # [1, 2, H, W]\n            if random.random() < 0.1:\n                print(f\"{normals.shape}\")\n        \n        # Pad/crop grads to expected size\n        if grads.shape[-1] != expected_size or grads.shape[-2] != expected_size:\n            padding_mode = 'replicate'\n            if random.random() < 0.1:\n                print(f\"Padding grads from {grads.shape[-2:]} to [{expected_size}, {expected_size}] using {padding_mode} mode\")\n            \n            # Calculate padding\n            h_pad = max(0, expected_size - grads.shape[-2])\n            w_pad = max(0, expected_size - grads.shape[-1])\n            \n            pad_left = w_pad // 2\n            pad_right = w_pad - pad_left\n            pad_top = h_pad // 2\n            pad_bottom = h_pad - pad_top\n            \n            # Apply padding\n            grads = F.pad(grads, (pad_left, pad_right, pad_top, pad_bottom), mode=padding_mode)\n            if random.random() < 0.1:\n                print(f\"Padded grads to: {grads.shape}\")\n        \n        # Same for normals\n        if normals.shape[-1] != expected_size or normals.shape[-2] != expected_size:\n            padding_mode = 'replicate'\n            if random.random() < 0.1:\n                print(f\"Padding normals from {normals.shape[-2:]} to [{expected_size}, {expected_size}] using {padding_mode} mode\")\n            \n            # Calculate padding\n            h_pad = max(0, expected_size - normals.shape[-2])\n            w_pad = max(0, expected_size - normals.shape[-1])\n            \n            pad_left = w_pad // 2\n            pad_right = w_pad - pad_left\n            pad_top = h_pad // 2\n            pad_bottom = h_pad - pad_top\n            \n            # Apply padding\n            normals = F.pad(normals, (pad_left, pad_right, pad_top, pad_bottom), mode=padding_mode)\n            if random.random() < 0.1:\n                print(f\"Padded normals to: {normals.shape}\")\n            \n    except Exception as e:\n        print(f\"Error during tensor verification: {e}\")\n        print(\"Using fallback approach for tensor reshaping...\")\n        \n        # Fallback approach - just make sure dimensions are compatible for forward pass\n        if grads.dim() == 5:\n            grads = grads.squeeze(1)\n        if normals.dim() == 5:\n            normals = normals.squeeze(1)\n            \n        # Make sure tensors have 4 dimensions for interpolation\n        if grads.dim() < 4:\n            grads = grads.view(-1, grads.size(0), grads.size(1), grads.size(2))\n        if normals.dim() < 4:\n            normals = normals.view(-1, normals.size(0), normals.size(1), normals.size(2))\n        \n        # Emergency fallback - resize everything to expected size\n        if targets.shape[-1] != expected_size:\n            # Force targets to expected size\n            targets = F.interpolate(targets, size=(expected_size, expected_size), mode='bilinear', align_corners=False)\n            grads = F.interpolate(grads, size=(expected_size, expected_size), mode='bilinear', align_corners=False)\n            normals = F.interpolate(normals, size=(expected_size, expected_size), mode='bilinear', align_corners=False)\n            print(f\"Emergency resizing to {expected_size}x{expected_size}\")\n    \n    # One final check and force fix\n    if targets.shape[-1] != expected_size:\n        targets = F.pad(targets, (0, expected_size - targets.shape[-1], 0, expected_size - targets.shape[-2]), mode='replicate')\n    if grads.shape[-1] != expected_size:\n        grads = F.pad(grads, (0, expected_size - grads.shape[-1], 0, expected_size - grads.shape[-2]), mode='replicate')\n    if normals.shape[-1] != expected_size:\n        normals = F.pad(normals, (0, expected_size - normals.shape[-1], 0, expected_size - normals.shape[-2]), mode='replicate')\n    \n    return inputs, targets, grads, normals\n\n\n# %% Main Execution\nprint(\"--- Starting Full Workflow (U-Net with Residual Blocks) ---\")\nset_seed(cfg.seed)\nprint(f\"Device: {cfg.device}\")\nprint(f\"Using PyTorch version: {torch.__version__}\")\nif cfg.use_cuda:\n    print(f\"CUDA available: {torch.cuda.get_device_name(0)}\")\n\n# ==============================================================================\n# Cleanup Code\n# ==============================================================================\n\n\n\nprint(\"\\n--- Cleaning up previous run artifacts ---\")\npaths_to_clean = [cfg.shard_output_dir]\n# Find previous best model files based on pattern\nmodel_pattern = os.path.join(cfg.working_dir, \"unet_best_model_epoch_*_loss_*.pth\")\npaths_to_clean.extend(glob.glob(model_pattern))\n# Add plot and submission files\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():\n            print(f\"Attempting to remove directory: {path_obj}\")\n            shutil.rmtree(path_obj, ignore_errors=True)\n            print(f\"Removed directory (if existed): {path_obj}\")\n        elif path_obj.is_file():\n            print(f\"Attempting to remove file: {path_obj}\")\n            path_obj.unlink(missing_ok=True)  # Ignore error if file doesn't exist\n            print(f\"Removed file (if existed): {path_obj}\")\n    except Exception as e:\n        print(f\"W: Error during cleanup of {path_obj}: {e}\")\nprint(\"--- Cleanup finished ---\")\ngc.collect()\n# ==============================================================================\n\n# ==============================================================================\n# SECTION 0/1: Sharding from Kaggle Data Only\n# ==============================================================================\nprint(\"\\n--- 0/1. Sharding from Kaggle Data Only ---\")\n# Use updated dataset name in shard path\nshard_stage_dir = Path(cfg.shard_output_dir) / f\"train_{cfg.dataset_name}\"\nkaggle_train_root = Path(cfg.kaggle_train_dir)\nneeds_creation = True\ntotal_samples_written = 0\n\ntry:\n    # --- Check if shards need creating ---\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\n    # --- Ensure output directories exist & Check Disk Space ---\n    print(\"\\n--- Checking Disk Space Before Directory Creation ---\")\n    try:\n        total, used, free = shutil.disk_usage(cfg.working_dir)\n        print(\n            f\"Disk Usage for {cfg.working_dir}: Total={total / 1e9:.2f}GB, Used={used / 1e9:.2f}GB, Free={free / 1e9:.2f}GB\"\n        )\n    except Exception as du_e:\n        print(f\"W: Could not check disk usage: {du_e}\")\n\n    try:\n        Path(cfg.shard_output_dir).mkdir(parents=True, exist_ok=True)\n        shard_stage_dir.mkdir(parents=True, exist_ok=True)\n    except OSError as e:\n        print(f\"E: Critical error creating output directories: {e}\")\n        raise  # Stop if directories can't be created\n\n    # --- Sharding Process ---\n    if needs_creation:\n        print(\n            f\"Starting shard creation from {kaggle_train_root} into {shard_stage_dir}\"\n        )\n        if not kaggle_train_root.is_dir():\n            raise FileNotFoundError(f\"Kaggle train directory not found: {kaggle_train_root}\")\n\n        # Find family subdirectories in the Kaggle train directory\n        families = [d.name for d in kaggle_train_root.iterdir() if d.is_dir()]\n        print(f\"Searching Kaggle data families: {families}\")\n        if not families:\n            raise FileNotFoundError(\n                f\"No family subdirectories found in {kaggle_train_root}\"\n            )\n\n        print(\"Searching for all data pairs in Kaggle source...\")\n        kaggle_file_pairs = search_data_path(\n            families, kaggle_train_root, shuffle=True, seed=cfg.seed\n        )\n        print(f\"Found {len(kaggle_file_pairs)} total valid pairs from Kaggle source.\")\n        if not kaggle_file_pairs:\n            raise RuntimeError(\n                \"No valid data pairs found in the specified Kaggle directories.\"\n            )\n\n        # --- Write Shards ---\n        shard_pattern = str(shard_stage_dir / \"%06d.tar\")\n        print(\n            f\"Writing shards using pattern {shard_pattern} (max size {cfg.maxsize / 1e9:.2f} GB)\"\n        )\n        with wds.ShardWriter(shard_pattern, maxsize=int(cfg.maxsize)) as writer:\n            common_base_dir = kaggle_train_root  # For relative path key generation\n            for in_file, out_file in tqdm(\n                kaggle_file_pairs, desc=\"Sharding Kaggle Data\", unit=\"pair\"\n            ):\n                # generate_sample handles potential errors for each pair\n                samples_from_pair = generate_sample(\n                    Path(in_file), Path(out_file), base_dir=common_base_dir\n                )\n                if samples_from_pair:\n                    for sample_dict in samples_from_pair:\n                        writer.write(sample_dict)\n                    total_samples_written += len(samples_from_pair)\n\n        print(\n            f\"Finished writing {total_samples_written} samples from Kaggle source to shards.\"\n        )\n\n    elif not needs_creation:\n        existing_shard_count = len(list(shard_stage_dir.glob(\"*.tar\")))\n        print(f\"Using {existing_shard_count} existing shards.\")\n\nexcept Exception as e:\n    print(f\"E: Kaggle-only sharding process failed critically: {e}\")\n    import traceback\n\n    traceback.print_exc()\n    raise\n# ==============================================================================\n\n\n# --- 2. Get Train/Val DataLoaders from Created Shards ---\nprint(\"\\n--- 2. Creating DataLoaders from Shards ---\")\ndltrain, dlvalid = None, None\nval_paths_saved = []  # Keep track of validation paths for potential later use\ntry:\n    # Use updated dataset name for getting paths\n    trn_paths, val_paths = get_shard_paths(\n        cfg.shard_output_dir,\n        cfg.dataset_name, # Use updated name\n        \"train\",  # Request splitting\n        num_shards=cfg.num_used_shards,\n        test_size=cfg.test_size,\n        seed=cfg.seed,\n    )\n    val_paths_saved = val_paths  # Save the validation paths\n\n    if trn_paths is None:\n        # get_shard_paths returns None, None on critical split error\n        raise RuntimeError(\"Failed to get or split shard paths for train/val.\")\n\n    # Check if any shards actually exist if paths were returned empty\n    # Use updated dataset name in check path\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(\n            f\"No training shards selected AND no .tar files found in {shard_check_dir}.\"\n        )\n\n    # Report shard counts\n    if not trn_paths:\n        print(\"W: No shards assigned for training. Training cannot proceed.\")\n    else:\n        print(f\"Using {len(trn_paths)} shards for training.\")\n    if not val_paths:\n        print(\"W: No shards assigned for validation. Validation will be skipped.\")\n    else:\n        print(f\"Using {len(val_paths)} shards for validation.\")\n\n    # Create WebDatasets (Augmentation applied in get_dataset for 'train')\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\n\n    # Check if dataset creation failed unexpectedly\n    if trn_ds is None and trn_paths:\n        raise RuntimeError(\"Failed to create train WebDataset pipeline.\")\n    if val_ds is None and val_paths:\n        # Only warn if validation dataset failed, training might still proceed\n        print(\"W: Failed to create validation WebDataset pipeline.\")\n\n    # Create DataLoaders\n    if trn_ds:\n        n_trn_w = min(cfg.num_workers, len(trn_paths)) if trn_paths else 0\n        p_trn = n_trn_w > 0  # Use persistent workers only if num_workers > 0\n        \n        # For WebDataset 0.2.111, we don't need a custom collation function\n        # because we're handling the batching at the WebDataset level\n        dltrain = DataLoader(\n            trn_ds,  # WebDataset already handles batching\n            batch_size=None,  # No additional batching in DataLoader\n            shuffle=False,  # Shuffling done by WebDataset\n            num_workers=n_trn_w,\n            pin_memory=cfg.use_cuda,\n            persistent_workers=p_trn,\n            prefetch_factor=2 if p_trn else None,  # Only relevant if num_workers > 0\n        )\n        print(f\"Train DataLoader created with {n_trn_w} workers (WebDataset v0.2.111 compatible).\")\n        \n    if val_ds:\n        n_val_w = min(cfg.num_workers, len(val_paths)) if val_paths else 0\n        p_val = n_val_w > 0\n        dlvalid = DataLoader(\n            val_ds,  # WebDataset already handles batching\n            batch_size=None,\n            shuffle=False,\n            num_workers=n_val_w,\n            pin_memory=cfg.use_cuda,\n            persistent_workers=p_val,\n            prefetch_factor=2 if p_val else None,\n        )\n        print(f\"Validation DataLoader created with {n_val_w} workers (WebDataset v0.2.111 compatible).\")\n\n    # Final check (can sometimes trigger TypeError: 'IterableDataset' has no len())\n    try:\n        loaders_exist = bool(dltrain or dlvalid)\n        if not loaders_exist and (trn_paths or val_paths):\n             # Should not happen if datasets were created but loaders failed\n             raise RuntimeError(\"Loaders missing despite dataset paths existing.\")\n        print(\"DataLoader(s) created successfully or skipped appropriately.\")\n    except TypeError as te:\n        # Expected error for IterableDataset without explicit length\n        if \"has no len()\" in str(te):\n            print(f\"W: Caught expected TypeError '{te}'. Assume DataLoaders are ok.\")\n        else:\n            raise te  # Re-raise unexpected TypeError\nexcept Exception as e:\n    print(f\"E: DataLoader creation failed critically: {e}\")\n    import traceback\n    traceback.print_exc()\n    raise\n\n\n# --- 3. Initialize Model, Loss, Optimizer ---\nprint(\"\\n--- 3. Initializing Model, Loss, Optimizer ---\")\nmodel = None\ntry:\n    # Instantiate the modified U-Net\n    model = UNet().to(cfg.device)\n    params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"Model: {model.__class__.__name__} (Residual), Trainable Params: {params:,}\")\n    criterion = nn.L1Loss()  # Mean Absolute Error\n    optimizer = torch.optim.AdamW(\n        model.parameters(), lr=cfg.learning_rate, weight_decay=cfg.weight_decay\n    )\n    print(f\"Loss Function: {criterion.__class__.__name__}\")\n    print(f\"Optimizer: {optimizer.__class__.__name__} (lr={cfg.learning_rate}, wd={cfg.weight_decay})\")\nexcept Exception as e:\n    print(f\"E: Model initialization failed: {e}\")\n    raise\n\n\n# --- 4. Training Loop ---\nprint(\"\\n--- 4. Starting Training ---\")\nhistory = []\nbest_val_loss = float(\"inf\")\n\nif dltrain is None or model is None:\n    print(\"E: Training cannot proceed. Train DataLoader or Model is missing.\")\nelse:\n    try:\n        for epoch in range(1, cfg.n_epochs + 1):\n            print(f\"\\n=== Epoch {epoch}/{cfg.n_epochs} ===\")\n            # --- Training Phase ---\n            gc.collect()  # Run garbage collection before each epoch\n            if cfg.use_cuda:\n                torch.cuda.empty_cache()\n            model.train()\n            train_losses = []\n            pbar_train = tqdm(dltrain, desc=f\"Train E{epoch}\", leave=False, unit=\"batch\")\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                try:\n                    # Check if batch is a valid dictionary with required fields\n                    if not isinstance(batch, dict):\n                        print(f\"W: Skipping invalid train batch {i} (type: {type(batch).__name__})\")\n                        continue\n                        \n                    if \"seis\" not in batch or \"vel\" not in batch or \"grad\" not in batch or \"normal\" not in batch:\n                        print(f\"W: Skipping incomplete train batch {i} (missing keys)\")\n                        continue\n                        \n                    # Access batch dictionary elements and move them to device\n                    inputs = batch[\"seis\"].to(cfg.device, non_blocking=True)\n                    targets = batch[\"vel\"].to(cfg.device, non_blocking=True)\n                    batch_grads = batch[\"grad\"].to(cfg.device)\n                    batch_normals = batch[\"normal\"].to(cfg.device)\n\n                    # Fix double batching issue - detect and reshape tensor if it has extra dimensions\n                    # Error example: [8, 8, 5, 1000, 70] - has an extra batch dimension\n                    if inputs.dim() > 4:\n                        print(f\"Fixing double-batched inputs: shape before = {inputs.shape}\")\n                        # Calculate new shape: flatten the first two dimensions if they're both batch dimensions\n                        new_batch_size = inputs.size(0) * inputs.size(1)\n                        new_shape = [new_batch_size] + list(inputs.shape[2:])\n                        inputs = inputs.reshape(new_shape).float()\n                        print(f\"Fixed inputs shape = {inputs.shape}\")\n\n                    # Similarly fix other tensors if they have extra dimensions\n                    if targets.dim() > 4:\n                        print(f\"Fixing double-batched targets: shape before = {targets.shape}\")\n                        new_batch_size = targets.size(0) * targets.size(1)\n                        new_shape = [new_batch_size] + list(targets.shape[2:])\n                        targets = targets.reshape(new_shape).float()\n                        print(f\"Fixed targets shape = {targets.shape}\")\n\n                    if batch_grads.dim() > 4:\n                        print(f\"Fixing double-batched grads: shape before = {batch_grads.shape}\")\n                        new_batch_size = batch_grads.size(0) * batch_grads.size(1)\n                        new_shape = [new_batch_size] + list(batch_grads.shape[2:])\n                        batch_grads = batch_grads.reshape(new_shape).float()\n                        print(f\"Fixed grads shape = {batch_grads.shape}\")\n\n                    if batch_normals.dim() > 4:\n                        print(f\"Fixing double-batched normals: shape before = {batch_normals.shape}\")\n                        new_batch_size = batch_normals.size(0) * batch_normals.size(1)\n                        new_shape = [new_batch_size] + list(batch_normals.shape[2:])\n                        batch_normals = batch_normals.reshape(new_shape).float()\n                        print(f\"Fixed normals shape = {batch_normals.shape}\")\n                        \n                    # Make sure all tensors are in float format\n                    inputs = inputs.float()\n                    targets = targets.float()\n                    batch_grads = batch_grads.float()\n                    batch_normals = batch_normals.float()\n\n                    # Verify tensor shapes on first batch of each epoch\n                    if i == 0:\n                        print(f\"Training batch 0 tensor verification:\")\n                        print(f\"  inputs: {inputs.shape}\")\n                        print(f\"  targets: {targets.shape}\")\n                        print(f\"  grads: {batch_grads.shape}\")\n                        print(f\"  normals: {batch_normals.shape}\")\n                        inputs, targets, batch_grads, batch_normals = verify_batch_tensors(\n                            inputs, targets, batch_grads, batch_normals\n                        )\n\n                    optimizer.zero_grad(set_to_none=True)\n                    # Use Automatic Mixed Precision (AMP) if on CUDA\n                    with torch.amp.autocast(\n                        device_type=cfg.device.type,\n                        dtype=cfg.autocast_dtype,\n                        enabled=cfg.use_cuda,\n                    ):\n                        try:\n                            # Run model forward pass\n                            outputs, pred_grad, pred_normal = model(inputs)\n                            \n                            # Ensure consistent dimensions for targets (may have extra channel dim)\n                            if targets.dim() == 4 and targets.size(1) == 1:\n                                targets_flat = targets.squeeze(1)\n                            else:\n                                targets_flat = targets\n                                \n                            # Velocity loss remains the same\n                            loss_vel = criterion(outputs, targets)\n                            \n                            # Check for dimension mismatch in model outputs vs targets\n                            if outputs.shape[-1] != targets.shape[-1] or outputs.shape[-2] != targets.shape[-2]:\n                                if random.random() < 0.1:  # Only log occasionally\n                                    print(f\"Dimension mismatch in batch {i}: outputs {outputs.shape} vs targets {targets.shape}\")\n                                # Resize tensors to match exactly\n                                outputs, targets = ensure_tensor_sizes(outputs, targets, expected_size=70)\n                                # Recalculate velocity loss with resized tensors\n                                loss_vel = criterion(outputs, targets)\n                            \n                            # Reshape batch_grads if needed (handle any extra dimensions)\n                            if batch_grads.dim() > 4:  # If it's 5D: [B, 1, 2, H, W] -> [B, 2, H, W]\n                                batch_grads = batch_grads.squeeze(1)\n                            \n                            # Reshape batch_normals if needed\n                            if batch_normals.dim() > 4:  # If it's 5D: [B, 1, 2, H, W] -> [B, 2, H, W]\n                                batch_normals = batch_normals.squeeze(1)\n                            \n                            # Gradient loss (weighted MAE) - ensure shapes match\n                            if pred_grad.shape[-2:] != batch_grads.shape[-2:]:\n                                # Print detailed shape info for debugging \n                                if random.random() < 0.1:  # Reduce logging\n                                    print(f\"Shape mismatch in batch {i}: pred_grad {pred_grad.shape} vs batch_grads {batch_grads.shape}\")\n                                # Resize tensors to match exactly\n                                pred_grad, batch_grads = ensure_tensor_sizes(pred_grad, batch_grads, expected_size=70)\n                            \n                            loss_grad = F.l1_loss(pred_grad, batch_grads)\n                            \n                            # Normal loss (cosine similarity) - ensure shapes match\n                            if pred_normal.shape[-2:] != batch_normals.shape[-2:]:\n                                # Print detailed shape info for debugging\n                                if random.random() < 0.1:  # Reduce logging\n                                    print(f\"Shape mismatch in batch {i}: pred_normal {pred_normal.shape} vs batch_normals {batch_normals.shape}\")\n                                # Resize tensors to match exactly\n                                pred_normal, batch_normals = ensure_tensor_sizes(pred_normal, batch_normals, expected_size=70)\n                            \n                            # Compute normalized dot product for cosine similarity\n                            true_normal = batch_normals\n                            # Ensure vectors are normalized for proper cosine similarity\n                            pred_norm = torch.sqrt(torch.sum(pred_normal**2, dim=1, keepdim=True) + 1e-6)\n                            true_norm = torch.sqrt(torch.sum(true_normal**2, dim=1, keepdim=True) + 1e-6)\n                            \n                            pred_normalized = pred_normal / pred_norm\n                            true_normalized = true_normal / true_norm\n                            \n                            dot_product = (pred_normalized * true_normalized).sum(dim=1)\n                            loss_normal = 1 - dot_product.mean()\n                            \n                            # Combined loss with adaptive weights\n                            total_loss = loss_vel + 0.3*loss_grad + 0.2*loss_normal\n                        \n                        except RuntimeError as e:\n                            if \"size mismatch\" in str(e) or \"sizes of tensors must match\" in str(e):\n                                print(f\"\\nE: Tensor size mismatch in batch {i}. Details: {e}\")\n                                print(f\"Shapes - inputs: {inputs.shape}, targets: {targets.shape}, \"\n                                     f\"batch_grads: {batch_grads.shape}, batch_normals: {batch_normals.shape}\")\n                                \n                                # Fall back to just velocity loss for this batch\n                                outputs, _, _ = model(inputs)\n                                total_loss = criterion(outputs, targets)\n                                print(\"Falling back to velocity loss only for this batch.\")\n                            else:\n                                # Re-raise other RuntimeErrors\n                                raise e\n\n                    # Backward pass and optimization\n                    optimizer.zero_grad()\n                    total_loss.backward()\n                    optimizer.step()\n                    train_losses.append(loss_vel.item())  # Track main velocity loss for history\n\n                    # Garbage collect every few batches\n                    if i % 5 == 0:  # Every 5 batches\n                        gc.collect()\n                        if cfg.use_cuda:\n                            torch.cuda.empty_cache()\n\n                    # Update progress bar description\n                    if i % 100 == 0:\n                        pbar_train.set_postfix(loss=f\"{np.mean(train_losses):.5f}\")\n\n                except Exception as e:\n                    print(f\"\\nE: Training batch {i} failed: {e}\")\n                    # Stop training on OOM error\n                    if isinstance(e, torch.cuda.OutOfMemoryError):\n                        print(\"E: CUDA Out of Memory during training. Exiting.\")\n                        raise e\n                    # Continue on other errors if desired, or raise\n                    # raise e # Uncomment to stop on any training error\n\n            avg_train_loss = np.mean(train_losses) if train_losses else 0.0\n            print(f\"Epoch {epoch} Avg Train Loss: {avg_train_loss:.5f}\")\n\n            # --- Validation Phase ---\n            if dlvalid is None:\n                print(\"W: Skipping validation phase - no validation DataLoader.\")\n                history.append(\n                    {\"epoch\": epoch, \"train_loss\": avg_train_loss, \"valid_loss\": None}\n                )\n                continue  # Skip to next epoch\n\n            # Run garbage collection before validation\n            gc.collect()\n            if cfg.use_cuda:\n                torch.cuda.empty_cache()\n\n            model.eval()\n            val_losses = []\n            pbar_val = tqdm(dlvalid, desc=f\"Valid E{epoch}\", leave=False, unit=\"batch\")\n            with torch.no_grad():\n                for i, batch in enumerate(pbar_val):\n                    if not batch or \"seis\" not in batch or \"vel\" not in batch:\n                        print(f\"W: Skipping invalid validation batch {i}\")\n                        continue\n                    try:\n                        # Check if batch is a valid dictionary with required fields\n                        if not isinstance(batch, dict):\n                            print(f\"W: Skipping invalid validation batch {i} (type: {type(batch).__name__})\")\n                            continue\n                            \n                        if \"seis\" not in batch or \"vel\" not in batch or \"grad\" not in batch or \"normal\" not in batch:\n                            print(f\"W: Skipping incomplete validation batch {i} (missing keys)\")\n                            continue\n                            \n                        # Access batch dictionary elements and move them to device\n                        inputs = batch[\"seis\"].to(cfg.device, non_blocking=True)\n                        targets = batch[\"vel\"].to(cfg.device, non_blocking=True)\n                        batch_grads = batch[\"grad\"].to(cfg.device)\n                        batch_normals = batch[\"normal\"].to(cfg.device)\n\n                        # Fix double batching issue - detect and reshape tensor if it has extra dimensions\n                        # Error example: [8, 8, 5, 1000, 70] - has an extra batch dimension\n                        if inputs.dim() > 4:\n                            print(f\"Fixing double-batched inputs: shape before = {inputs.shape}\")\n                            # Calculate new shape: flatten the first two dimensions if they're both batch dimensions\n                            new_batch_size = inputs.size(0) * inputs.size(1)\n                            new_shape = [new_batch_size] + list(inputs.shape[2:])\n                            inputs = inputs.reshape(new_shape).float()\n                            print(f\"Fixed inputs shape = {inputs.shape}\")\n\n                        # Similarly fix other tensors if they have extra dimensions\n                        if targets.dim() > 4:\n                            print(f\"Fixing double-batched targets: shape before = {targets.shape}\")\n                            new_batch_size = targets.size(0) * targets.size(1)\n                            new_shape = [new_batch_size] + list(targets.shape[2:])\n                            targets = targets.reshape(new_shape).float()\n                            print(f\"Fixed targets shape = {targets.shape}\")\n\n                        if batch_grads.dim() > 4:\n                            print(f\"Fixing double-batched grads: shape before = {batch_grads.shape}\")\n                            new_batch_size = batch_grads.size(0) * batch_grads.size(1)\n                            new_shape = [new_batch_size] + list(batch_grads.shape[2:])\n                            batch_grads = batch_grads.reshape(new_shape).float()\n                            print(f\"Fixed grads shape = {batch_grads.shape}\")\n\n                        if batch_normals.dim() > 4:\n                            print(f\"Fixing double-batched normals: shape before = {batch_normals.shape}\")\n                            new_batch_size = batch_normals.size(0) * batch_normals.size(1)\n                            new_shape = [new_batch_size] + list(batch_normals.shape[2:])\n                            batch_normals = batch_normals.reshape(new_shape).float()\n                            print(f\"Fixed normals shape = {batch_normals.shape}\")\n                            \n                        # Make sure all tensors are in float format\n                        inputs = inputs.float()\n                        targets = targets.float()\n                        batch_grads = batch_grads.float()\n                        batch_normals = batch_normals.float()\n\n                        # Verify tensor shapes on first batch of validation\n                        if i == 0:\n                            print(f\"Validation batch 0 tensor verification:\")\n                            print(f\"  inputs: {inputs.shape}\")\n                            print(f\"  targets: {targets.shape}\")\n                            print(f\"  grads: {batch_grads.shape}\")\n                            print(f\"  normals: {batch_normals.shape}\")\n                            inputs, targets, batch_grads, batch_normals = verify_batch_tensors(\n                                inputs, targets, batch_grads, batch_normals\n                            )\n\n                        with torch.amp.autocast(\n                            device_type=cfg.device.type,\n                            dtype=cfg.autocast_dtype,\n                            enabled=cfg.use_cuda,\n                        ):\n                            try:\n                                # Run model forward pass\n                                outputs, pred_grad, pred_normal = model(inputs)\n                                \n                                # Ensure consistent dimensions for targets (may have extra channel dim)\n                                if targets.dim() == 4 and targets.size(1) == 1:\n                                    targets_flat = targets.squeeze(1)\n                                else:\n                                    targets_flat = targets\n                                    \n                                # Velocity loss remains the same\n                                loss_vel = criterion(outputs, targets)\n                                \n                                # Check for dimension mismatch in model outputs vs targets\n                                if outputs.shape[-1] != targets.shape[-1] or outputs.shape[-2] != targets.shape[-2]:\n                                    if random.random() < 0.1:  # Reduce logging\n                                        print(f\"Dimension mismatch in batch {i}: outputs {outputs.shape} vs targets {targets.shape}\")\n                                    # Resize tensors to match exactly\n                                    outputs, targets = ensure_tensor_sizes(outputs, targets, expected_size=70)\n                                    # Recalculate velocity loss with resized tensors\n                                    loss_vel = criterion(outputs, targets)\n                                \n                                # Reshape batch_grads if needed (handle any extra dimensions)\n                                if batch_grads.dim() > 4:  # If it's 5D: [B, 1, 2, H, W] -> [B, 2, H, W]\n                                    batch_grads = batch_grads.squeeze(1)\n                                \n                                # Reshape batch_normals if needed\n                                if batch_normals.dim() > 4:  # If it's 5D: [B, 1, 2, H, W] -> [B, 2, H, W]\n                                    batch_normals = batch_normals.squeeze(1)\n                                \n                                # Gradient loss (weighted MAE) - ensure shapes match\n                                if pred_grad.shape[-2:] != batch_grads.shape[-2:]:\n                                    # Print detailed shape info for debugging \n                                    if random.random() < 0.1:  # Reduce logging\n                                        print(f\"Shape mismatch in batch {i}: pred_grad {pred_grad.shape} vs batch_grads {batch_grads.shape}\")\n                                    # Resize tensors to match exactly\n                                    pred_grad, batch_grads = ensure_tensor_sizes(pred_grad, batch_grads, expected_size=70)\n                                \n                                loss_grad = F.l1_loss(pred_grad, batch_grads)\n                                \n                                # Normal loss (cosine similarity) - ensure shapes match\n                                if pred_normal.shape[-2:] != batch_normals.shape[-2:]:\n                                    # Print detailed shape info for debugging\n                                    if random.random() < 0.1:  # Reduce logging\n                                        print(f\"Shape mismatch in batch {i}: pred_normal {pred_normal.shape} vs batch_normals {batch_normals.shape}\")\n                                    # Resize tensors to match exactly\n                                    pred_normal, batch_normals = ensure_tensor_sizes(pred_normal, batch_normals, expected_size=70)\n                                \n                                # Compute normalized dot product for cosine similarity\n                                true_normal = batch_normals\n                                # Ensure vectors are normalized for proper cosine similarity\n                                pred_norm = torch.sqrt(torch.sum(pred_normal**2, dim=1, keepdim=True) + 1e-6)\n                                true_norm = torch.sqrt(torch.sum(true_normal**2, dim=1, keepdim=True) + 1e-6)\n                                \n                                pred_normalized = pred_normal / pred_norm\n                                true_normalized = true_normal / true_norm\n                                \n                                dot_product = (pred_normalized * true_normalized).sum(dim=1)\n                                loss_normal = 1 - dot_product.mean()\n                                \n                                # Combined loss with same weights as training\n                                loss = loss_vel + 0.3*loss_grad + 0.2*loss_normal\n                            \n                            except RuntimeError as e:\n                                if \"size mismatch\" in str(e) or \"sizes of tensors must match\" in str(e):\n                                    print(f\"\\nE: Tensor size mismatch in validation batch {i}. Details: {e}\")\n                                    print(f\"Shapes - inputs: {inputs.shape}, targets: {targets.shape}, \"\n                                        f\"batch_grads: {batch_grads.shape}, batch_normals: {batch_normals.shape}\")\n                                    \n                                    # Fall back to just velocity loss for this batch\n                                    outputs, _, _ = model(inputs)\n                                    loss = criterion(outputs, targets)\n                                    print(\"Falling back to velocity loss only for this validation batch.\")\n                                else:\n                                    # Re-raise other RuntimeErrors\n                                    raise e\n\n                        val_losses.append(loss_vel.item())  # Track main velocity loss for history\n                        \n                        # Garbage collect periodically during validation\n                        if i % 5 == 0:  # Every 5 batches\n                            gc.collect()\n                            if cfg.use_cuda:\n                                torch.cuda.empty_cache()\n\n                        # Plotting validation examples periodically\n                        if i == 0 and epoch % cfg.plot_every_n_epochs == 0:\n                            # Add validation plotting code here if desired\n                            pass # Placeholder\n\n                    except Exception as e:\n                        print(f\"\\nE: Validation batch {i} failed: {e}\")\n                        if isinstance(e, torch.cuda.OutOfMemoryError):\n                            print(\"E: CUDA Out of Memory during validation. Exiting.\")\n                            raise e\n                        # Continue on other errors if desired, or raise\n                        # raise e # Uncomment to stop on any validation error\n\n            avg_val_loss = np.mean(val_losses) if val_losses else float(\"inf\")\n            print(f\"Epoch {epoch} Avg Valid Loss: {avg_val_loss:.5f}\")\n            history.append(\n                {\"epoch\": epoch, \"train_loss\": avg_train_loss, \"valid_loss\": avg_val_loss}\n            )\n\n            # --- Save Best Model ---\n            if avg_val_loss < best_val_loss:\n                best_val_loss = avg_val_loss\n                # Clean previous best models before saving new one\n                del_pattern = os.path.join(\n                    cfg.working_dir, f\"unet_best_model_epoch_*_loss_*.pth\"\n                )\n                for old_model_path in glob.glob(del_pattern):\n                    try:\n                        print(f\"    Removing old best model: {os.path.basename(old_model_path)}\")\n                        os.remove(old_model_path)\n                    except OSError as e:\n                        print(f\"W: Could not delete old model {old_model_path}: {e}\")\n\n                # Save the new best model\n                fname = f\"unet_best_model_epoch_{epoch}_loss_{best_val_loss:.4f}.pth\"\n                fpath = os.path.join(cfg.working_dir, fname)\n                print(f\"*** New best validation loss: {best_val_loss:.5f}. Saving model: {fname} ***\")\n                torch.save(model.state_dict(), fpath)\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# --- 5. Plot History ---\nprint(\"\\n--- 5. Plotting Training History ---\")\nif history:\n    try:\n        hist_df = pd.DataFrame(history)\n        plt.figure(figsize=(12, 6))\n        plt.plot(hist_df[\"epoch\"], hist_df[\"train_loss\"], \"o-\", label=\"Train Loss\")\n        # Only plot validation loss if it exists and is not all None/NaN\n        if \"valid_loss\" in hist_df.columns and not hist_df[\"valid_loss\"].isnull().all():\n            plt.plot(\n                hist_df[\"epoch\"],\n                hist_df[\"valid_loss\"],\n                \"s--\",  # Square markers, dashed line\n                label=\"Validation Loss\",\n            )\n        plt.title(\"Training and Validation Loss vs. Epoch (Residual U-Net)\")\n        plt.xlabel(\"Epoch\")\n        plt.ylabel(\"L1 Loss (Mean Absolute Error)\")\n        plt.legend()\n        plt.grid(True, linestyle=\"--\", alpha=0.6)\n        plt.ylim(bottom=0)  # Loss should not be negative\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()  # Display the plot\n    except Exception as e:\n        print(f\"E: Failed plotting training history: {e}\")\nelse:\n    print(\"No training history recorded to plot.\")\n\n\n# --- 6. Error Analysis on Validation Set ---\nprint(\"\\n--- 6. Error Analysis on Validation Set ---\")\n# Placeholder for error analysis code - requires dlvalid or val_paths_saved\nbest_model_path_analysis = find_best_model()\nif not best_model_path_analysis:\n    print(\"W: No best model found. Skipping analysis.\")\nelif dlvalid is None and not val_paths_saved:\n    # Need either the original loader or the paths to recreate it\n    print(\"W: Validation loader/paths unavailable. Skipping analysis.\")\nelse:\n    print(f\"Performing analysis using model: {os.path.basename(best_model_path_analysis)}\")\n    # Add analysis code block here (e.g., load model, get batch, predict, plot errors)\n    # Ensure to handle potential recreation of dlvalid if it was lost\n    # Example: Load model, get a batch from dlvalid, predict, compare pred/target\n    # model_analysis = UNet().to(cfg.device)\n    # model_analysis.load_state_dict(torch.load(best_model_path_analysis, map_location=cfg.device))\n    # model_analysis.eval()\n    # with torch.no_grad():\n    #     # Get a batch (handle if dlvalid needs recreation from val_paths_saved)\n    #     # batch = next(iter(dlvalid_or_recreated))\n    #     # inputs = batch[\"seis\"].to(cfg.device)... targets = batch[\"vel\"]...\n    #     # preds = model_analysis(inputs)\n    #     # Plot difference, calculate stats, etc.\n    pass\n\n\n# --- 7. Prediction (Using Kaggle Test Set) ---\nprint(\"\\n--- 7. Final Prediction on Kaggle Test Set ---\")\nbest_model_final_path = find_best_model()\nif not best_model_final_path:\n    print(\"W: No best 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 model for final prediction: {os.path.basename(best_model_final_path)}\")\n        # Make sure to instantiate the correct model class (UNet)\n        model_pred = UNet().to(cfg.device)\n        model_pred.load_state_dict(torch.load(best_model_final_path, map_location=cfg.device))\n        model_pred.eval()\n\n        test_ds = KaggleTestDataset(cfg.kaggle_test_dir)\n        if len(test_ds) == 0:\n            print(\"W: Kaggle test dataset is empty. No submission generated.\")\n        else:\n            # Setup DataLoader for test set\n            t_bs = max(1, cfg.batch_size // 2)\n            t_nw = min(\n                max(0, cfg.num_workers // 2),\n                (os.cpu_count() // 2 if os.cpu_count() else 1),\n            )\n            dl_test = DataLoader(\n                test_ds,\n                batch_size=t_bs,\n                shuffle=False,\n                num_workers=t_nw,\n                pin_memory=cfg.use_cuda,\n            )\n            print(f\"Test DataLoader created with 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                # Define CSV header columns (x_1, x_3, ..., x_69)\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, original_ids in pbar_test:\n                        # Handle batch size = 1 where original_ids might be a string\n                        if isinstance(original_ids, str):\n                            original_ids = [original_ids]\n                        try:\n                            inputs = inputs.to(cfg.device).float()\n                            with torch.amp.autocast(\n                                device_type=cfg.device.type,\n                                dtype=cfg.autocast_dtype,\n                                enabled=cfg.use_cuda,\n                            ):\n                                # Get velocity prediction (discard gradients/normals for submission)\n                                velocity, _, _ = model_pred(inputs)\n                                \n                            # Output shape is (B, 1, H, W), get predictions (B, H, W)\n                            if velocity.dim() == 4 and velocity.size(1) == 1:\n                                preds = velocity[:, 0].cpu().numpy()\n                            else:\n                                preds = velocity.cpu().numpy()\n\n                            # Iterate through samples in the batch\n                            for y_pred, oid in zip(preds, original_ids): # y_pred is (H, W)\n                                # Iterate through y-positions (rows) for this sample\n                                for y_pos in range(y_pred.shape[0]): # y_pred.shape[0] should be 70\n                                    # Extract values at odd x-indices (1, 3, ..., 69)\n                                    vals = y_pred[y_pos, 1::2].astype(np.float32)\n                                    # Create row dictionary\n                                    row = dict(zip(x_cols, vals))\n                                    row[\"oid_ypos\"] = f\"{oid}_y_{y_pos}\"\n                                    writer.writerow(row)\n                                    rows_written += 1\n                        except Exception as e:\n                            # Report error but continue if possible\n                            print(\n                                f\"\\nE: Prediction failed for batch (OID: {original_ids[0] if original_ids else '?'}) : {e}\"\n                            )\n\n            print(f\"Submission file created: {cfg.submission_file} ({rows_written} rows).\")\n            # Sanity check row count\n            expected_rows = len(test_ds) * 70  # 70 y-positions per test sample\n            if rows_written != expected_rows:\n                print(\n                    f\"W: Row count mismatch! Expected {expected_rows}, but wrote {rows_written}.\"\n                )\n\n    except Exception as e:\n        print(f\"E: Final prediction process failed critically: {e}\")\n        import traceback\n        traceback.print_exc()\n\nprint(\"\\n--- Full Workflow Finished ---\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}