{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n\"\"\"\nScript for training a U-Net model for Full Waveform Inversion,\nusing data sourced solely from Kaggle input directories, with data augmentation.\nCorrected generate_sample function.\n\"\"\"\n\n!pip install webdataset\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\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\nimport webdataset as wds\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm.auto import tqdm\n\n\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\"\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 = 16\n    num_workers = 2\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 data\n\n    # --- Model params (U-Net) ---\n    unet_in_channels = 5\n    unet_out_channels = 1\n    unet_init_features = 32\n    unet_depth = 3\n    unet_bilinear = True\n\n    # --- Training params ---\n    n_epochs = 50\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    # Continue pipeline even if some samples fail decoding/mapping\n    map_handler = wds.warn_and_continue\n\n    try:\n        dataset = wds.WebDataset(\n            paths, nodesplitter=wds.split_by_node, shardshuffle=is_train, seed=seed\n        )\n        # 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                seis_tensor = torch.from_numpy(s_np).float()\n                vel_tensor = torch.from_numpy(v_np).float()\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                return {\"sample_id\": sid, \"seis\": seis_tensor, \"vel\": vel_tensor}\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        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)\nclass DoubleConv(nn.Module):\n    \"\"\"(Convolution => [BN] => ReLU) * 2\"\"\"\n\n    def __init__(self, in_channels, out_channels, mid_channels=None):\n        super().__init__()\n        if not mid_channels:\n            mid_channels = out_channels\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(mid_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)\n\n\nclass Down(nn.Module):\n    \"\"\"Downscaling with MaxPool then DoubleConv\"\"\"\n\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.maxpool_conv = nn.Sequential(\n            nn.MaxPool2d(2), DoubleConv(in_channels, out_channels)\n        )\n\n    def forward(self, x):\n        return self.maxpool_conv(x)\n\n\nclass Up(nn.Module):\n    \"\"\"Upscaling then DoubleConv\"\"\"\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(\n                scale_factor=2, mode=\"bilinear\", align_corners=False\n            )\n            # In bilinear mode, input channels to conv is sum of skip and upsampled\n            conv_in_channels = in_channels + out_channels\n            # Mid channels in DoubleConv is explicitly set to out_channels\n            self.conv = DoubleConv(\n                conv_in_channels, out_channels, mid_channels=out_channels\n            )\n        else:\n            # Use ConvTranspose2d for learned upsampling\n            self.up = nn.ConvTranspose2d(\n                in_channels, in_channels // 2, kernel_size=2, stride=2\n            )\n            # Input channels to conv is sum of skip and halved upsampled channels\n            conv_in_channels = (in_channels // 2) + out_channels\n            self.conv = DoubleConv(conv_in_channels, out_channels)\n\n    def forward(self, x1, x2):\n        # x1 is the feature map from the layer below (needs upsampling)\n        # x2 is the skip connection from the corresponding encoder layer\n        x1 = self.up(x1)\n\n        # Pad x1 if its dimensions don't match x2 after upsampling\n        # Input is CHW\n        diffY = x2.size(2) - x1.size(2)\n        diffX = x2.size(3) - x1.size(3)\n\n        # Pad format: (padding_left, padding_right, padding_top, padding_bottom)\n        x1 = F.pad(\n            x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]\n        )\n\n        # Concatenate along the 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\"\"\"\n\n    def __init__(\n        self,\n        n_channels=cfg.unet_in_channels,\n        n_classes=cfg.unet_out_channels,\n        init_features=cfg.unet_init_features,\n        depth=cfg.unet_depth,\n        bilinear=cfg.unet_bilinear,\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 average pooling layer specific to this problem\n        self.initial_pool = nn.AvgPool2d(kernel_size=(14, 1), stride=(14, 1))\n\n        # --- Encoder ---\n        self.encoder_blocks = nn.ModuleList()\n        self.inc = DoubleConv(n_channels, init_features)\n        self.encoder_blocks.append(self.inc)\n        current_features = init_features\n        for _ in range(depth):\n            down_block = Down(current_features, current_features * 2)\n            self.encoder_blocks.append(down_block)\n            current_features *= 2\n\n        # --- Bottleneck ---\n        # The last block of the encoder acts as the bottleneck\n        bottleneck_features = current_features\n\n        # --- Decoder ---\n        self.decoder_blocks = nn.ModuleList()\n        current_features = bottleneck_features\n        for _ in range(depth):\n            # Output channels are halved at each step\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 ---\n        self.outc = OutConv(current_features, n_classes)\n\n    def _pad_or_crop(self, x, target_h=70, target_w=70):\n        \"\"\"Pads or crops input tensor x to target height and width.\"\"\"\n        _, _, h, w = x.shape\n        # Pad Height if needed\n        if h < target_h:\n            pad_top = (target_h - h) // 2\n            pad_bottom = target_h - h - pad_top\n            x = F.pad(x, (0, 0, pad_top, pad_bottom))  # Pad height only\n            h = target_h\n        # Pad Width if needed\n        if w < target_w:\n            pad_left = (target_w - w) // 2\n            pad_right = target_w - w - pad_left\n            x = F.pad(x, (pad_left, pad_right, 0, 0))  # Pad width only\n            w = target_w\n        # Crop Height if needed\n        if h > target_h:\n            crop_top = (h - target_h) // 2\n            # Use slicing to crop\n            x = x[:, :, crop_top : crop_top + target_h, :]\n            h = target_h\n        # Crop Width if needed\n        if w > target_w:\n            crop_left = (w - target_w) // 2\n            x = x[:, :, :, crop_left : crop_left + target_w]\n            w = target_w\n        return x\n\n    def forward(self, x):\n        # Initial pooling and resizing\n        x_pooled = self.initial_pool(x)\n        x_resized = self._pad_or_crop(x_pooled, target_h=70, target_w=70)\n\n        # --- Encoder Path ---\n        skip_connections = []\n        xi = x_resized\n        for i, block in enumerate(self.encoder_blocks):\n            xi = block(xi)\n            # Store intermediate feature maps for skip connections (except bottleneck)\n            if i < len(self.encoder_blocks) - 1:\n                skip_connections.append(xi)\n\n        # --- Decoder Path ---\n        # Start with the bottleneck output (last xi)\n        xu = xi\n        # Iterate through decoder blocks and corresponding skip connections in reverse\n        for i, block in enumerate(self.decoder_blocks):\n            skip = skip_connections[len(skip_connections) - 1 - i]\n            xu = block(xu, skip)\n\n        # --- Final Output ---\n        logits = self.outc(xu)\n        # Apply scaling and offset specific to the problem's target range\n        output = logits * 1000.0 + 1500.0\n        return output\n\n\n# %% Main Execution\nprint(\"--- Starting Full Workflow ---\")\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# ==============================================================================\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 ---\")\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    trn_paths, val_paths = get_shard_paths(\n        cfg.shard_output_dir,\n        cfg.dataset_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    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        dltrain = DataLoader(\n            trn_ds.batched(cfg.batch_size),\n            batch_size=None,  # Already batched by WebDataset\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.\")\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.batched(cfg.batch_size),\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.\")\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    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__}, 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()\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                    inputs = batch[\"seis\"].to(cfg.device, non_blocking=True).float()\n                    targets = batch[\"vel\"].to(cfg.device, non_blocking=True).float()\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                        outputs = model(inputs)\n                        loss = criterion(outputs, targets)\n\n                    # Backward pass and optimization step\n                    loss.backward()\n                    optimizer.step()\n                    train_losses.append(loss.item())\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            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                        inputs = batch[\"seis\"].to(cfg.device, non_blocking=True).float()\n                        targets = batch[\"vel\"].to(cfg.device, non_blocking=True).float()\n                        with torch.amp.autocast(\n                            device_type=cfg.device.type,\n                            dtype=cfg.autocast_dtype,\n                            enabled=cfg.use_cuda,\n                        ):\n                            outputs = model(inputs)\n                            loss = criterion(outputs, targets)\n                        val_losses.append(loss.item())\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\")\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\n# Keep the existing analysis code block here if needed, ensuring dlvalid exists or is recreated.\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    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        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            # Use slightly smaller batch size and fewer workers for inference if needed\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                                outputs = model_pred(inputs)\n                            # Output shape is (B, 1, H, W), get predictions (B, H, W)\n                            preds = outputs[:, 0].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,"execution":{"iopub.status.busy":"2025-04-20T06:10:10.095585Z","iopub.execute_input":"2025-04-20T06:10:10.095917Z","iopub.status.idle":"2025-04-20T06:10:38.578744Z","shell.execute_reply.started":"2025-04-20T06:10:10.095895Z","shell.execute_reply":"2025-04-20T06:10:38.577755Z"}},"outputs":[],"execution_count":null}]}