{"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"},{"sourceId":362945,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":301398,"modelId":319856}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🧠 Predictions Analysis: Understanding Model Behavior Beyond the MAE\n\nThis notebook is focused **entirely on analyzing the predictions** made by a trained neural network for seismic velocity inversion. The base model is adapted from [**GWI UNet with float16 dataset**](https://www.kaggle.com/code/egortrushin/gwi-unet-with-float16-dataset) by Egor Trushin — The model itself is not the focus here.\n\nInstead, **our goal is to deeply analyze what the model gets right, what it misses, where it struggles, and why.** We examine predictions using a battery of visual and statistical diagnostics:\n\n- Residual histograms (to reveal bias and variance)\n- Depth-wise error curves (to show performance by vertical layer)\n- Spectral residual plots (to diagnose frequency-specific shortcomings)\n- Heatmaps of pixel-wise MAE (to highlight spatial blind spots)\n- Best/Median/Worst case comparisons (to give concrete examples)\n- Cross-family metrics summaries and interactive tables\n\n---\n\n### 🔍 The `FWIPredictionAnalyzer` Class\n\nTo make this analysis easy and reusable, we created a self-contained class: `FWIPredictionAnalyzer`.\n\nYou can drop this class into your own PyTorch project and run powerful post-hoc analysis on **any model that maps seismic data → velocity maps**. It is plug-and-play:\n- Just point it to a folder of `.npy` files\n- Provide your PyTorch model (already loaded)\n- Run `.analyze()` — it handles the rest\n\nIt computes structured metrics like **MAE**, **RMSE**, **SSIM**, and **PSNR**, generates plots, saves all outputs, and even styles interactive tables for easier insight. It's a lab in a class.\n\n---\n\n### 💡 Why This Matters\n\nLooking only at leaderboard MAE is like grading an essay by word count. This analyzer breaks down performance **by geology type, by depth, by frequency content, and by spatial location**—so you can **pinpoint failure modes and plan improvements**.\n\nIf you care about model reliability, interpretability, or squeezing out extra performance in competitions, this analysis is essential.\n\n---\n","metadata":{"_uuid":"90432583-d9c4-4416-affc-38f6a117b090","_cell_guid":"d5e3d704-5852-4baf-be45-0a0f9d28cbae","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import os\nimport datetime\nimport random\nimport time\nimport csv\nfrom pathlib import Path\n\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\n\n##############################################\n# Configuration (Soft-Coded)\n##############################################\nconfig = {\n\"seed\": 42,\n\"test_data_path\": \"/kaggle/input/waveform-inversion/test\",\n\"batch_size\": 16,\n\"read_weights\": \"/kaggle/input/unet-resnet-ep06/pytorch/ep38/1/best_model.pth\",   # Path to your trained model weights\n\"model\": {\n    \"name\": \"UNet\",\n    \"unet_params\": {\n        \"init_features\": 32,\n        \"depth\": 5\n    }\n}\n}\n\n##############################################\n# Utility Functions\n##############################################\ndef format_time(elapsed):\n    \"\"\"Take a time in seconds and return a string hh:mm:ss.\"\"\"\n    elapsed_rounded = int(round(elapsed))\n    return str(datetime.timedelta(seconds=elapsed_rounded))\n\ndef seed_everything(seed_value: int) -> None:\n    \"\"\"Set a global random seed for reproducible results.\"\"\"\n    random.seed(seed_value)\n    np.random.seed(seed_value)\n    torch.manual_seed(seed_value)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed_value)\n        torch.cuda.manual_seed_all(seed_value)\n    if torch.backends.cudnn.is_available():\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n\n##############################################\n# Datasets\n##############################################\nclass TestDataset(Dataset):\n    \"\"\"\n    No more pre-initialization of all memmaps. Instead, we\n    load each file as needed in __getitem__.\n    \"\"\"\n    def __init__(self, test_files):\n        self.test_files = test_files\n    \n    def __len__(self):\n        return len(self.test_files)\n    \n    def __getitem__(self, i):\n        fpath = self.test_files[i]\n        # Load the entire array directly from disk\n        arr = np.load(fpath)  \n        arr_t = torch.tensor(arr, dtype=torch.float32)\n        return arr_t, fpath.stem\n\n##############################################\n# Model Definition (UNet + Components)\n##############################################\nclass ResidualDoubleConv(nn.Module):\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    \n        self.conv1 = nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(mid_channels)\n        self.relu = nn.ReLU(inplace=True)\n    \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        if in_channels == out_channels:\n            self.shortcut = nn.Identity()\n        else:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n    \n    def forward(self, x):\n        identity = x\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n    \n        out = self.conv2(out)\n        out = self.bn2(out)\n    \n        out += self.shortcut(identity)\n        out = self.relu(out)\n        return out\n\nclass Up(nn.Module):\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            self.conv = ResidualDoubleConv(in_channels + out_channels, out_channels)\n        else:\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 = self.up(x1)\n        diffY = x2.size(2) - x1.size(2)\n        diffX = x2.size(3) - x1.size(3)\n        x1 = F.pad(\n            x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]\n        )\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\nclass OutConv(nn.Module):\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\nclass UNet(nn.Module):\n    def __init__(\n        self,\n        n_channels=5,\n        n_classes=1,\n        init_features=32,\n        depth=5,\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        # Example: input shape might be [B, 5, 1000, 70]\n        self.initial_pool = nn.AvgPool2d(kernel_size=(14, 1), stride=(14, 1))\n    \n        self.encoder_convs = nn.ModuleList()\n        self.encoder_pools = nn.ModuleList()\n    \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            conv = ResidualDoubleConv(current_features, current_features * 2)\n            pool = nn.MaxPool2d(2)\n            self.encoder_convs.append(conv)\n            self.encoder_pools.append(pool)\n            current_features *= 2\n    \n        self.bottleneck = ResidualDoubleConv(current_features, current_features)\n    \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        self.outc = OutConv(current_features, n_classes)\n    \n    def _pad_or_crop(self, x, target_h=70, target_w=70):\n        \"\"\"\n        Utility to ensure each feature map is 70x70 \n        (padding or cropping as necessary).\n        \"\"\"\n        _, _, h, w = x.shape\n        # Pad or crop height\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))\n            h = target_h\n        elif h > target_h:\n            crop_top = (h - target_h) // 2\n            x = x[:, :, crop_top:crop_top + target_h, :]\n            h = target_h\n    \n        # Pad or crop width\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))\n            w = target_w\n        elif w > target_w:\n            crop_left = (w - target_w) // 2\n            x = x[:, :, :, crop_left:crop_left + target_w]\n            w = target_w\n    \n        return x\n    \n    def forward(self, x):\n        x_pooled = self.initial_pool(x)  # ~ [B, 5, 71, 70]\n        x_resized = self._pad_or_crop(x_pooled, 70, 70)\n    \n        skip_connections = []\n        xi = x_resized\n        xi = self.encoder_convs[0](xi)  # first conv\n        skip_connections.append(xi)\n    \n        for i in range(self.depth):\n            xi = self.encoder_convs[i + 1](xi)\n            skip_connections.append(xi)\n            xi = self.encoder_pools[i](xi)\n    \n        xi = self.bottleneck(xi)\n    \n        xu = xi\n        for i, block in enumerate(self.decoder_blocks):\n            skip_index = self.depth - 1 - i\n            skip = skip_connections[skip_index]\n            xu = block(xu, skip)\n    \n        logits = self.outc(xu)\n        # Scale the output, as done in training\n        output = logits * 1000.0 + 1500.0\n        return output","metadata":{"_uuid":"a806e10c-91c6-491e-a732-6a6d7e2edc5f","_cell_guid":"dc50e9d2-c258-4c64-b9b0-a04728ee17ea","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-05-13T02:08:15.440057Z","iopub.execute_input":"2025-05-13T02:08:15.440220Z","iopub.status.idle":"2025-05-13T02:08:19.459148Z","shell.execute_reply.started":"2025-05-13T02:08:15.440203Z","shell.execute_reply":"2025-05-13T02:08:19.458339Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"##############################################\n# Main: Inference Only\n##############################################\n\n# Seed for reproducibility\nseed_everything(config[\"seed\"])\n\n# Device\nif torch.cuda.is_available():\n    device = torch.device(\"cuda\")\n    print(\"Using CUDA:\", device)\nelse:\n    device = torch.device(\"cpu\")\n    print(\"Using CPU:\", device)\n\n# Build model\nmodel_params = config[\"model\"][\"unet_params\"]\nmodel = UNet(**model_params)\nmodel.to(device)\n\n# Load trained weights\nif config[\"read_weights\"] is not None:\n    print(\"Loading weights from:\", config[\"read_weights\"])\n    model.load_state_dict(torch.load(config[\"read_weights\"], map_location=device))\nelse:\n    raise ValueError(\"No model weights provided. Please set config['read_weights'].\")\n\nmodel.eval()\n\n# Prepare test dataset\ntest_path = Path(config[\"test_data_path\"])\ntest_files = sorted(list(test_path.glob(\"*.npy\")))\nif not test_files:\n    raise RuntimeError(f\"No .npy test files found in {test_path}\")\n\nds_test = TestDataset(test_files)\n\ndl_test = DataLoader(\n    ds_test,\n    batch_size=config[\"batch_size\"],\n    num_workers=0,  \n    pin_memory=False,\n    drop_last=False\n)","metadata":{"_uuid":"f1fa537b-2c7c-4ddd-a196-6ba221fb0bb2","_cell_guid":"44be2706-50e8-46f1-91b7-cb389fbd0372","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-05-13T02:08:19.460354Z","iopub.execute_input":"2025-05-13T02:08:19.460997Z","iopub.status.idle":"2025-05-13T02:08:24.040748Z","shell.execute_reply.started":"2025-05-13T02:08:19.460971Z","shell.execute_reply":"2025-05-13T02:08:24.040172Z"},"jupyter":{"outputs_hidden":false},"_kg_hide-input":true,"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## PREDICTIONS Analysis Section","metadata":{"_uuid":"fdfe3668-facc-4e53-a291-7dfdd0db2352","_cell_guid":"23695a46-cc38-4294-aa09-a731e11be311","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"## REVEAL AND COPY THE CELL BELOW TO USE IN YOUR OWN PROJECTS.  \n### SEE VERY EASY USAGE CELL NEXT.","metadata":{}},{"cell_type":"code","source":"%pip install itables -q # Uncomment and run if itables is not installed\n%pip install plotly -q # Ensure plotly is installed\n\nimport os\nimport numpy as np\nimport random\nimport torch\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport plotly.express as px # Import Plotly Express\nimport seaborn as sns\nfrom skimage.metrics import structural_similarity as ssim\nfrom skimage.metrics import peak_signal_noise_ratio as psnr\nfrom tqdm.notebook import tqdm # Use notebook version for better rendering in Jupyter/Kaggle\nimport warnings\nfrom itables import init_notebook_mode, show # Import itables\nfrom pandas.io.formats.style import Styler # Import Styler for table formatting\nimport math # For ceiling function in grid layout\n\nfrom typing import Dict, List, Tuple, Optional, Any, Set, Callable\n\n# Initialize itables for interactive DataFrames\ninit_notebook_mode(all_interactive=True)\n\n\nclass FWIPredictionAnalyzer:\n    \"\"\"\n    Scans a root directory for subfolders (classes), identifies seismic/velocity\n    .npy file pairs containing multiple samples/scenarios. Allows analyzing a\n    subset of these samples using a PyTorch model. Computes various metrics\n    (MAE, RMSE, SSIM, PSNR), generates summary statistics per class, and\n    produces a range of plots for model evaluation.\n\n    Enhancements:\n    - Displays summary metrics in an interactive, styled table (itable).\n    - Ensures consistent dataset ordering (flatfault-a, -b, flatvel-a, -b, curves, styles).\n    - Uses full metric names (e.g., \"Root Mean Squared Error\").\n    - Creates a 2x3 diagnostic grid plot per dataset *followed immediately* by best/median/worst examples.\n    - Optionally displays boxplots interactively using Plotly, while saving static versions using Seaborn.\n    - Saves all metrics (CSV) and plots (PNG) to a specified output directory.\n    - Displays plots inline in the notebook environment *in addition* to saving them.\n    - Addresses seaborn FutureWarning by handling infinite values explicitly.\n    - Styles the summary itable with column-wise color gradients.\n    - Provides functions to generate cross-family comparison grids for various diagnostics.\n    - Allows suppressing common warnings during initialization/output.\n    - **FIXED**: Correctly labels \"Best\" and \"Worst\" samples based on metric value.\n    \"\"\"\n    # --- Class Constants ---\n    # 1. Define desired class order (roughly by complexity)\n    DESIRED_ORDER = [\"flatfault-a\", \"flatfault-b\", \"flatvel-a\", \"flatvel-b\", \"curves\", \"styles\"]\n    FULL_METRIC_NAMES = {\n        \"mae\": \"Mean Absolute Error\",\n        \"rmse\": \"Root Mean Squared Error\",\n        \"ssim\": \"Structural Similarity Index\",\n        \"psnr\": \"Peak Signal-to-Noise Ratio\",\n    }\n    # Reverse mapping for convenience\n    METRIC_ABBREVIATIONS = {v: k for k, v in FULL_METRIC_NAMES.items()}\n    # Identify metrics where higher is better for styling and sorting logic\n    HIGHER_IS_BETTER_METRICS = {\"ssim\", \"psnr\"} # Used for styling and sorting direction\n\n\n    def __init__(\n        self,\n        root_dir: str,\n        model: torch.nn.Module,\n        device: str = \"cpu\",\n        max_samples_per_class: int = 25,\n        viz_engine: str = \"interactive\", # \"interactive\" enables Plotly boxplots display\n        metrics_to_compute: Set[str] = {\"mae\", \"rmse\", \"ssim\", \"psnr\"},\n        output_dir: str = \"/kaggle/working/fwi_analysis_output\",\n        suppress_output_warnings: bool = True, # 5. Added warning suppression flag\n    ):\n        \"\"\"\n        Args:\n            root_dir (str): Path to the directory containing subfolders (classes).\n            model (torch.nn.Module): A PyTorch model predicting velocity from seismic.\n            device (str): \"cpu\" or \"cuda\" for model inference.\n            max_samples_per_class (int): Max number of samples to analyze from each class.\n            viz_engine (str): \"interactive\" for Plotly boxplots display, else static (Seaborn).\n            metrics_to_compute (Set[str]): Set of scalar metric abbreviations ('mae', 'rmse', 'ssim', 'psnr').\n            output_dir (str): Directory to save generated plots and metrics CSVs. Created if it doesn't exist.\n            suppress_output_warnings (bool): If True, suppresses common warnings from libraries like skimage during output generation.\n        \"\"\"\n        self.root_dir = root_dir\n        self.model = model.eval().to(device)\n        self.device = device\n        self.max_samples_per_class = max_samples_per_class\n        self.viz_engine = viz_engine.lower().strip()\n        self.metrics_to_compute = {m.lower() for m in metrics_to_compute}\n        self.suppress_output_warnings = suppress_output_warnings\n\n        # Apply warning suppression based on the flag\n        if self.suppress_output_warnings:\n            warnings.filterwarnings(\"ignore\", message=\"Inputs have mismatched dtype\", category=UserWarning, module='skimage')\n            warnings.filterwarnings(\"ignore\", message=\"Setting data_range based on im_true\", category=UserWarning, module='skimage')\n            # Add other common warnings if needed, e.g., from matplotlib or seaborn\n            warnings.filterwarnings(\"ignore\", category=FutureWarning, module='seaborn') # General future warnings from seaborn\n            warnings.filterwarnings(\"ignore\", category=UserWarning, module='matplotlib') # Common matplotlib user warnings\n            warnings.simplefilter(action='ignore', category=FutureWarning) # Broader future warning suppression\n            warnings.simplefilter(action='ignore', category=UserWarning) # Broader user warning suppression\n        else:\n            # Reset to default behavior if suppression is off\n            warnings.resetwarnings() # Or set specific ones back to 'default' if needed\n\n        self.output_dir = output_dir\n        try:\n            os.makedirs(self.output_dir, exist_ok=True)\n            print(f\"Output will be saved to: {self.output_dir}\")\n        except OSError as e:\n            print(f\"Error creating output directory {self.output_dir}: {e}. Using current directory.\")\n            self.output_dir = \".\"\n\n        self._pairs: Dict[str, List[Tuple[str, str]]] = self._scan_pairs()\n        if not self._pairs:\n             print(f\"Warning: No valid (seismic, velocity) file pairs found in subfolders of {root_dir}\")\n\n        self.sample_details: Dict[str, List[Dict[str, Any]]] = {}\n        self.summary_stats: Dict[str, Dict[str, float]] = {}\n        self.all_samples_df: Optional[pd.DataFrame] = None\n        self.metrics_summary_table_df: Optional[pd.DataFrame] = None\n\n\n    def _scan_pairs(self) -> Dict[str, List[Tuple[str, str]]]:\n        \"\"\"Scans root_dir for subfolders and matches seis/vel files.\"\"\"\n        # (Logic remains the same - assuming it works as intended)\n        pairs: Dict[str, List[Tuple[str, str]]] = {}\n        if not os.path.isdir(self.root_dir):\n            print(f\"Error: Root directory '{self.root_dir}' not found.\")\n            return pairs\n\n        for folder_name in os.listdir(self.root_dir):\n            ds_dir = os.path.join(self.root_dir, folder_name)\n            if not os.path.isdir(ds_dir):\n                continue\n\n            matched_pairs = []\n            data_dir = os.path.join(ds_dir, \"data\")\n            model_dir = os.path.join(ds_dir, \"model\")\n\n            # Case 1: data/ and model/ subfolders\n            if os.path.isdir(data_dir) and os.path.isdir(model_dir):\n                try:\n                    data_files = [f for f in os.listdir(data_dir) if f.startswith(\"data\") and f.endswith(\".npy\")]\n                    for f in data_files:\n                        suffix = f[len(\"data\"):-4]\n                        seis_path = os.path.join(data_dir, f)\n                        vel_path = os.path.join(model_dir, f\"model{suffix}.npy\")\n                        if os.path.exists(vel_path):\n                            matched_pairs.append((seis_path, vel_path))\n                except OSError as e:\n                    print(f\"Warning: Could not read files in {data_dir} or {model_dir}. Skipping. Error: {e}\")\n                    continue\n\n            # Case 2: Direct folder with seis*.npy and vel*.npy\n            else:\n                try:\n                    files = os.listdir(ds_dir)\n                    seis_fs = sorted([f for f in files if f.startswith('seis') and f.endswith('.npy')])\n                    vel_map = {vf: os.path.join(ds_dir, vf) for vf in files if vf.startswith('vel') and vf.endswith('.npy')}\n\n                    for sf in seis_fs:\n                        # Corrected suffix extraction to handle potential variations like seis_001.npy -> vel_001.npy\n                        suffix_part = sf[len('seis'):-4] # Extracts '_001' or similar\n                        expected_vf_name = f\"vel{suffix_part}.npy\" # Constructs 'vel_001.npy'\n                        corresponding_vel_path = vel_map.get(expected_vf_name)\n\n                        if corresponding_vel_path and os.path.exists(corresponding_vel_path):\n                             matched_pairs.append((os.path.join(ds_dir, sf), corresponding_vel_path))\n                        # else: # Optional: Add a warning if a seis file doesn't have a matching vel file\n                        #     print(f\"Warning: No matching velocity file found for {sf} (expected {expected_vf_name}) in {ds_dir}\")\n\n\n                except OSError as e:\n                    print(f\"Warning: Could not read files in {ds_dir}. Skipping. Error: {e}\")\n                    continue\n\n            if matched_pairs:\n                pairs[folder_name] = sorted(matched_pairs)\n\n        return pairs\n\n    @property\n    def datasets(self) -> List[str]:\n        \"\"\"Returns the list of discovered dataset class names (subfolders).\"\"\"\n        return list(self._pairs.keys())\n\n    @property\n    def sorted_datasets(self) -> List[str]:\n        \"\"\"Returns dataset names sorted according to DESIRED_ORDER.\"\"\"\n        found_datasets = set(self.datasets)\n        # Ensure DESIRED_ORDER only includes datasets actually found\n        ordered_list = [ds for ds in self.DESIRED_ORDER if ds in found_datasets]\n        # Include remaining datasets (not in DESIRED_ORDER) alphabetically\n        remaining = sorted([ds for ds in found_datasets if ds not in self.DESIRED_ORDER])\n        return ordered_list + remaining\n\n\n    def _get_full_metric_name(self, metric_abbr: str) -> str:\n        \"\"\"Returns the full metric name for a given abbreviation.\"\"\"\n        return self.FULL_METRIC_NAMES.get(metric_abbr.lower(), metric_abbr.upper())\n\n    def _compute_metrics(self, pred: np.ndarray, true: np.ndarray) -> Dict[str, float]:\n        \"\"\"Computes requested scalar metrics for one sample. Assumes pred and true are 2D.\"\"\"\n        metrics = {}\n        pred = pred.astype(np.float32)\n        true = true.astype(np.float32)\n\n        if \"mae\" in self.metrics_to_compute:\n            metrics[\"mae\"] = float(np.mean(np.abs(pred - true)))\n        if \"rmse\" in self.metrics_to_compute:\n            metrics[\"rmse\"] = float(np.sqrt(np.mean((pred - true)**2)))\n\n        # Additional check for 2D needed for SSIM/PSNR\n        if pred.ndim != 2 or true.ndim != 2:\n            if \"ssim\" in self.metrics_to_compute: metrics[\"ssim\"] = np.nan\n            if \"psnr\" in self.metrics_to_compute: metrics[\"psnr\"] = np.nan\n            return metrics # Cannot compute SSIM/PSNR on non-2D data\n\n        # --- SSIM / PSNR Calculation (only if 2D) ---\n        data_range = np.nanmax(true) - np.nanmin(true) # Use nanmax/min\n        if not np.isfinite(data_range) or data_range <= 0:\n             # Fallback if range is zero or non-finite (e.g., constant image)\n             finite_true = true[np.isfinite(true)]\n             if finite_true.size > 0:\n                 # Use max absolute value as range, or 1.0 if all zero\n                 data_range = max(np.max(np.abs(finite_true)), 1e-6) # Ensure non-zero\n             else:\n                 data_range = 1.0 # Default if no finite values at all\n\n        # Ensure data_range is positive for ssim/psnr\n        safe_data_range = max(data_range, 1e-6)\n\n        # Ensure inputs are finite for skimage metrics by replacing NaN/inf\n        # A simple replacement with 0 might not be ideal, but is necessary for skimage\n        pred_finite = np.nan_to_num(pred, nan=0.0, posinf=np.nanmax(pred[np.isfinite(pred)]) if np.any(np.isfinite(pred)) else 0, neginf=np.nanmin(pred[np.isfinite(pred)]) if np.any(np.isfinite(pred)) else 0)\n        true_finite = np.nan_to_num(true, nan=0.0, posinf=np.nanmax(true[np.isfinite(true)]) if np.any(np.isfinite(true)) else 0, neginf=np.nanmin(true[np.isfinite(true)]) if np.any(np.isfinite(true)) else 0)\n\n\n        if \"ssim\" in self.metrics_to_compute:\n            try:\n                 metrics[\"ssim\"] = float(ssim(true_finite, pred_finite, data_range=safe_data_range))\n            except ValueError as e:\n                 # Catch potential issues like mismatched shapes (shouldn't happen here) or other skimage errors\n                 # print(f\"Warning: SSIM calculation failed: {e}. Setting SSIM to NaN.\")\n                 metrics[\"ssim\"] = np.nan\n        if \"psnr\" in self.metrics_to_compute:\n             try:\n                 # Use np.errstate to temporarily ignore divide-by-zero warnings for PSNR calculation\n                 # This happens if pred_finite == true_finite (MSE=0)\n                 with np.errstate(divide='ignore', invalid='ignore'):\n                    metric_psnr = psnr(true_finite, pred_finite, data_range=safe_data_range)\n\n                 # psnr returns inf if images are identical (MSE=0). Replace inf with NaN for statistical analysis.\n                 metrics[\"psnr\"] = float(metric_psnr) if np.isfinite(metric_psnr) else np.nan\n             except ValueError as e:\n                 # Catch potential issues\n                 # print(f\"Warning: PSNR calculation failed: {e}. Setting PSNR to NaN.\")\n                 metrics[\"psnr\"] = np.nan\n\n        return metrics\n\n\n    def analyze(self) -> Dict[str, Dict[str, float]]:\n        \"\"\"\n        Analyzes samples, runs inference, computes metrics, stores results, generates report.\n        Uses the sorted dataset order internally.\n        \"\"\"\n        if not self._pairs:\n            print(\"No dataset classes found or no valid pairs. Analysis cannot proceed.\")\n            return {}\n\n        # 1. Use sorted_datasets for processing order\n        ordered_datasets = self.sorted_datasets\n        print(f\"Starting analysis for {len(ordered_datasets)} classes in order: {ordered_datasets}\")\n        print(f\"Max samples per class: {self.max_samples_per_class}\")\n        print(f\"Metrics to compute: {[self._get_full_metric_name(m) for m in sorted(list(self.metrics_to_compute))]}\") # Sort metrics for consistent display\n\n        self.sample_details = {ds_name: [] for ds_name in ordered_datasets} # Initialize based on sorted order\n\n        # --- Sample Counting (using sorted order) ---\n        total_samples_to_process = 0\n        samples_per_class_map = {}\n        for ds_name in ordered_datasets: # Use sorted order\n             count = 0\n             if ds_name not in self._pairs: continue\n             try:\n                 for seis_path, vel_path in self._pairs[ds_name]:\n                     num_in_seis = self._get_num_samples_in_file(seis_path)\n                     num_in_vel = self._get_num_samples_in_file(vel_path)\n                     if num_in_seis == num_in_vel: count += num_in_seis\n                     else: pass # Warning printed inside _get_num_samples... if error\n             except Exception as e: print(f\"Error counting samples for {ds_name}: {e}\")\n             actual_samples_for_class = min(self.max_samples_per_class, count)\n             samples_per_class_map[ds_name] = actual_samples_for_class\n             total_samples_to_process += actual_samples_for_class\n        print(f\"Estimated total samples to analyze: {total_samples_to_process}\")\n        if total_samples_to_process == 0:\n            print(\"Warning: No samples will be processed.\")\n            # return {} # Optional early exit\n\n        # --- Sample Processing Loop (using sorted order) ---\n        with tqdm(total=total_samples_to_process, desc=\"Analyzing Samples\") as pbar:\n            for ds_name in ordered_datasets: # Use sorted order\n                 if ds_name not in self._pairs: continue\n\n                 files_in_class = self._pairs[ds_name]\n                 samples_collected_for_class = 0\n                 max_samples_for_this_class = samples_per_class_map.get(ds_name, 0)\n\n                 for seis_path, vel_path in files_in_class:\n                     if samples_collected_for_class >= max_samples_for_this_class: break\n                     try:\n                         # Load data\n                         seis_data_full = np.load(seis_path)\n                         vel_maps_full = np.load(vel_path)\n                         num_samples_in_file = seis_data_full.shape[0]\n                         if num_samples_in_file != vel_maps_full.shape[0]: continue # Skip inconsistent files\n\n                         sample_indices_in_file = list(range(num_samples_in_file))\n                         # Optional: random.shuffle(sample_indices_in_file)\n\n                         for sample_idx in sample_indices_in_file:\n                             if samples_collected_for_class >= max_samples_for_this_class: break\n\n                             # Extract sample\n                             seis_np = seis_data_full[sample_idx]\n                             vel_np_raw = vel_maps_full[sample_idx]\n\n                             # Ensure 2D GT\n                             if vel_np_raw.ndim == 3 and vel_np_raw.shape[0] == 1: vel_np = vel_np_raw.squeeze(0)\n                             elif vel_np_raw.ndim == 2: vel_np = vel_np_raw\n                             else: continue # Skip invalid GT shape\n\n                             # Inference\n                             X_torch = torch.from_numpy(seis_np).float().unsqueeze(0).to(self.device)\n                             with torch.no_grad(): pred = self.model(X_torch)\n\n                             # Ensure 2D Pred - handle different output shapes robustly\n                             try:\n                                 if pred.ndim == 4 and pred.shape[0] == 1 and pred.shape[1] == 1:\n                                     # Common case: [1, 1, H, W] -> [H, W]\n                                     pred_np = pred.squeeze(0).squeeze(0).cpu().numpy()\n                                 elif pred.ndim == 3 and pred.shape[0] == 1:\n                                     # Case: [1, H, W] -> [H, W]\n                                     pred_np = pred.squeeze(0).cpu().numpy()\n                                 elif pred.ndim == 2:\n                                     # Case: [H, W] -> [H, W] (already correct)\n                                     pred_np = pred.cpu().numpy()\n                                 else:\n                                     # Attempt general squeeze and hope for the best, or reshape\n                                     pred_np_squeezed = pred.squeeze().cpu().numpy()\n                                     if pred_np_squeezed.ndim == 2:\n                                          pred_np = pred_np_squeezed\n                                     else: # If squeezing didn't work, try reshaping (less safe)\n                                         try:\n                                             pred_np = pred.cpu().numpy().reshape(vel_np.shape)\n                                         except ValueError:\n                                              # print(f\"Warning: Skipping sample ({ds_name}, {os.path.basename(vel_path)}, {sample_idx}) due to incompatible prediction shape {pred.shape} vs GT shape {vel_np.shape}\")\n                                              continue # Give up if reshape fails\n\n                                 # Final check after processing\n                                 if pred_np.ndim != 2:\n                                     raise ValueError(f\"Processed Pred shape {pred_np.shape} is not 2D.\")\n\n                             except Exception as e:\n                                 # print(f\"Warning: Skipping sample ({ds_name}, {os.path.basename(vel_path)}, {sample_idx}) due to error processing prediction shape: {e}\")\n                                 continue # Skip problematic prediction shape\n\n                             # Shape check between processed pred and GT\n                             if pred_np.shape != vel_np.shape:\n                                # print(f\"Warning: Skipping sample ({ds_name}, {os.path.basename(vel_path)}, {sample_idx}) due to shape mismatch after processing: Pred {pred_np.shape} vs GT {vel_np.shape}\")\n                                continue\n\n                             # Metrics\n                             metrics = self._compute_metrics(pred_np, vel_np)\n\n                             # Store\n                             sample_record = {\n                                 \"pred\": pred_np, \"true\": vel_np,\n                                 \"seis_path\": seis_path, \"vel_path\": vel_path,\n                                 \"sample_idx_in_file\": sample_idx, \"dataset\": ds_name,\n                                 **metrics\n                             }\n                             # Use the correct key based on sorted order\n                             self.sample_details[ds_name].append(sample_record)\n                             samples_collected_for_class += 1\n                             pbar.update(1)\n\n                     except FileNotFoundError: continue # Skip missing files\n                     except Exception as e: # Catch other potential errors during file processing\n                         print(f\"\\nError processing file pair ({os.path.basename(seis_path)}, {os.path.basename(vel_path)}) for {ds_name}: {e}\")\n                         continue # Continue to next file pair\n\n        pbar.close()\n        print(\"\\nAnalysis complete. Generating report and saving outputs...\")\n        if total_samples_to_process == 0:\n             print(\"Warning: Zero samples were processed.\")\n\n        # Generate report uses the sorted order internally now\n        self.generate_report() # Includes saving and plotting\n        print(\"Report generation and saving finished.\")\n        return self.summary_stats\n\n\n    def _get_num_samples_in_file(self, npy_filepath: str) -> int:\n        \"\"\"Safely get sample count (first dimension) from .npy file header.\"\"\"\n        try:\n            with open(npy_filepath, 'rb') as f:\n                 version = np.lib.format.read_magic(f)\n                 shape, fortran_order, dtype = np.lib.format._read_array_header(f, version)\n                 if isinstance(shape, tuple) and len(shape) > 0:\n                     return shape[0]\n                 else:\n                     # Handle scalar or 0-dim case - assuming 1 sample if shape is ()\n                     if shape == (): return 1\n                     # print(f\"Warning: Unexpected array shape '{shape}' in {npy_filepath}. Assuming 0 samples.\")\n                     return 0\n        except FileNotFoundError:\n             # print(f\"Warning: File not found for sample counting: {npy_filepath}\") # Reduce noise\n             return 0\n        except Exception as e:\n            # print(f\"Warning: Could not read header of {npy_filepath}: {e}. Assuming 0 samples.\") # Reduce noise\n            return 0\n\n    def _save_figure(self, fig, filename_base: str, directory: Optional[str] = None, dpi=150):\n        \"\"\"Helper function to save matplotlib figures. Does NOT close the figure.\"\"\"\n        if directory is None: directory = self.output_dir\n        filepath = os.path.join(directory, f\"{filename_base}.png\")\n        try:\n            fig.savefig(filepath, dpi=dpi, bbox_inches='tight')\n            # print(f\"Saved plot: {filepath}\") # Optional confirmation\n        except Exception as e:\n            print(f\"Error saving figure {filepath}: {e}\")\n\n    def _summarize_stats(self) -> None:\n        \"\"\"Calculates summary stats, creates DataFrames, saves CSVs. Uses sorted order.\"\"\"\n        print(\"\\n--- Calculating Summary Statistics ---\")\n        all_samples_data_for_df = []\n        summary_table_data = []\n        self.summary_stats = {}\n\n        # 1. Use sorted_datasets for consistent row order in summary\n        for ds_name in self.sorted_datasets:\n            details_list = self.sample_details.get(ds_name, [])\n            count = len(details_list)\n            self.summary_stats[ds_name] = {\"count\": count}\n            summary_row = {\"Dataset\": ds_name, \"Sample Count\": count}\n\n            if not details_list:\n                # Still add row to summary table for consistency, but with NaNs/zeros\n                 for metric in sorted(list(self.metrics_to_compute)): # Sort metrics for consistent column order\n                    full_name = self._get_full_metric_name(metric)\n                    summary_row[f\"Mean {full_name}\"] = np.nan\n                    summary_row[f\"Median {full_name}\"] = np.nan\n                    summary_row[f\"Std Dev {full_name}\"] = np.nan\n                 summary_table_data.append(summary_row)\n                 continue # Skip rest for this empty dataset\n\n            metrics_data = {metric: [d.get(metric, np.nan) for d in details_list]\n                            for metric in self.metrics_to_compute}\n\n            # Sort metrics for consistent column order\n            for metric in sorted(list(self.metrics_to_compute)):\n                values = metrics_data[metric]\n                full_name = self._get_full_metric_name(metric)\n                # Handle potential inf values before calculating stats\n                valid_values = [v for v in values if np.isfinite(v)] # Exclude NaN and Inf\n\n                if not valid_values:\n                    mean_val, median_val, std_val = np.nan, np.nan, np.nan\n                    min_val, max_val = np.nan, np.nan # Keep for internal stats\n                else:\n                    mean_val = float(np.mean(valid_values))\n                    median_val = float(np.median(valid_values))\n                    std_val = float(np.std(valid_values))\n                    min_val = float(np.min(valid_values))\n                    max_val = float(np.max(valid_values))\n\n                self.summary_stats[ds_name].update({\n                    f\"mean_{metric}\": mean_val, f\"median_{metric}\": median_val,\n                    f\"min_{metric}\": min_val, f\"max_{metric}\": max_val,\n                    f\"std_{metric}\": std_val,\n                })\n                summary_row[f\"Mean {full_name}\"] = mean_val\n                summary_row[f\"Median {full_name}\"] = median_val\n                summary_row[f\"Std Dev {full_name}\"] = std_val\n\n            summary_table_data.append(summary_row)\n\n            for d in details_list:\n                record = {\n                    \"Dataset\": ds_name,\n                    \"Sample Index in File\": d[\"sample_idx_in_file\"],\n                    \"Seismic File\": os.path.basename(d[\"seis_path\"]),\n                    \"Velocity File\": os.path.basename(d[\"vel_path\"]),\n                    # Sort metrics for consistent column order\n                    **{self._get_full_metric_name(m): d.get(m, np.nan) for m in sorted(list(self.metrics_to_compute))}\n                }\n                all_samples_data_for_df.append(record)\n\n        # Create and save DataFrames\n        if all_samples_data_for_df:\n             # Define column order for per-sample CSV based on sorted metrics\n             sample_cols = [\"Dataset\", \"Sample Index in File\", \"Seismic File\", \"Velocity File\"] + \\\n                           [self._get_full_metric_name(m) for m in sorted(list(self.metrics_to_compute))]\n             self.all_samples_df = pd.DataFrame(all_samples_data_for_df)[sample_cols] # Apply column order\n             csv_path = os.path.join(self.output_dir, 'all_sample_metrics.csv')\n             try:\n                 self.all_samples_df.to_csv(csv_path, index=False)\n                 print(f\"Saved per-sample metrics to: {csv_path}\")\n             except Exception as e: print(f\"Error saving per-sample metrics CSV: {e}\")\n        else:\n             self.all_samples_df = pd.DataFrame()\n             print(\"\\nNo data collected for the per-sample DataFrame.\")\n\n        if summary_table_data:\n             # Define column order for summary CSV based on sorted metrics\n             cols = [\"Dataset\", \"Sample Count\"]\n             for metric in sorted(list(self.metrics_to_compute)): # Sort metrics\n                 full_name = self._get_full_metric_name(metric)\n                 cols.extend([f\"Mean {full_name}\", f\"Median {full_name}\", f\"Std Dev {full_name}\"])\n             # Ensure DataFrame is created with the correct row order (from sorted_datasets loop)\n             # Reindex based on sorted_datasets to guarantee order even if some datasets were empty\n             self.metrics_summary_table_df = pd.DataFrame(summary_table_data).set_index('Dataset')\n             # Filter out potential datasets in DESIRED_ORDER that were not found/processed\n             final_order = [ds for ds in self.sorted_datasets if ds in self.metrics_summary_table_df.index]\n             self.metrics_summary_table_df = self.metrics_summary_table_df.loc[final_order].reset_index()\n\n             self.metrics_summary_table_df = self.metrics_summary_table_df[cols] # Apply column order\n\n             csv_path_summary = os.path.join(self.output_dir, 'metrics_summary_by_dataset.csv')\n             try:\n                 self.metrics_summary_table_df.to_csv(csv_path_summary, index=False)\n                 print(f\"Saved summary metrics table to: {csv_path_summary}\")\n             except Exception as e: print(f\"Error saving summary metrics CSV: {e}\")\n        else:\n             self.metrics_summary_table_df = pd.DataFrame()\n             print(\"\\nNo data collected for the summary metrics table.\")\n\n\n    def _style_summary_table(self, styler: Styler) -> Styler:\n        \"\"\"Applies conditional formatting to the summary table Styler object.\"\"\"\n        # Identify metric columns based on prefixes\n        metric_cols_mean = [c for c in styler.data.columns if c.startswith(\"Mean \")]\n        metric_cols_median = [c for c in styler.data.columns if c.startswith(\"Median \")]\n        metric_cols_std = [c for c in styler.data.columns if c.startswith(\"Std Dev \")]\n        all_metric_cols = metric_cols_mean + metric_cols_median + metric_cols_std\n\n        # General float formatting for metric columns\n        styler.format(\"{:.4f}\", subset=all_metric_cols, na_rep='N/A')\n        # Integer formatting for count\n        styler.format(\"{:,d}\", subset=[\"Sample Count\"])\n\n        # Apply background gradients column-wise\n        for col_name in all_metric_cols:\n            # Robustly extract the base metric abbreviation from the column name\n            metric_abbr = None\n            for abbr, full in self.FULL_METRIC_NAMES.items():\n                 if full in col_name:\n                     metric_abbr = abbr\n                     break\n\n            if metric_abbr: # Check if we successfully identified the metric\n                # Use Reds for higher-is-worse (MAE, RMSE, Std Dev)\n                # Use Blues for higher-is-better (SSIM, PSNR)\n                # Std Dev is always \"higher is worse/more spread\" -> use Reds\n                if metric_abbr in self.HIGHER_IS_BETTER_METRICS and not col_name.startswith(\"Std Dev\"):\n                    cmap = 'Blues'\n                else: # MAE, RMSE, and all Std Devs use Reds\n                    cmap = 'Reds_r' # Use reversed Reds: lower values are redder (better for errors)\n\n                # Apply gradient styling\n                try:\n                    styler.background_gradient(subset=[col_name], cmap=cmap, axis=0) # axis=0 for column-wise\n                except Exception as e:\n                    print(f\"Warning: Could not apply style to column '{col_name}': {e}\") # Handle potential errors\n            else:\n                print(f\"Warning: Could not determine metric abbreviation for styling column '{col_name}'\")\n\n\n        # Add table-level styles\n        styler.set_table_styles([\n            {'selector': 'th', 'props': [('text-align', 'center'), ('font-weight','bold')]},\n            {'selector': 'td', 'props': [('text-align', 'center')]}\n        ])\n        styler.set_caption(\"Summary Statistics by Dataset\").set_table_styles([{\n            'selector': 'caption',\n            'props': [('font-size', '16px'), ('font-weight', 'bold'), ('text-align','center')]\n        }])\n        styler.set_properties(**{'border': '1px solid black', 'margin': 'auto', 'width':'95%'})\n\n        return styler\n\n\n    def _display_summary_table(self):\n        \"\"\"Displays the styled metrics summary using itable. Uses sorted dataset order.\"\"\"\n        print(\"\\n--- Analysis Summary Table ---\")\n        if self.metrics_summary_table_df is not None and not self.metrics_summary_table_df.empty:\n             # Ensure the DataFrame being styled is correctly ordered (should be from _summarize_stats)\n             styled_df = self.metrics_summary_table_df.style.pipe(self._style_summary_table)\n             show(styled_df, classes=\"display compact cell-border\", paging=False, searching=False, info=False)\n        else:\n            print(\" Summary metrics table is empty or not generated.\")\n\n\n    def _create_metric_boxplot(self, metric_abbr: str):\n        \"\"\"\n        Creates and saves a boxplot using Seaborn.\n        If viz_engine is 'interactive', also displays a Plotly boxplot.\n        Uses sorted dataset order.\n        \"\"\"\n        full_metric_name = self._get_full_metric_name(metric_abbr)\n        if self.all_samples_df is None or self.all_samples_df.empty or full_metric_name not in self.all_samples_df.columns:\n            print(f\"Cannot create boxplot for '{full_metric_name}': Data not available.\")\n            return\n\n        # Prepare data: Handle NaN and Inf before plotting\n        df_plot = self.all_samples_df[['Dataset', full_metric_name]].copy()\n        # Replace Inf with NaN to handle potential issues in plotting/stats\n        df_plot[full_metric_name] = df_plot[full_metric_name].replace([np.inf, -np.inf], np.nan)\n        # PSNR specifically can be very high, maybe cap it for visualization if needed?\n        # if metric_abbr == 'psnr': df_plot[full_metric_name] = df_plot[full_metric_name].clip(upper=60) # Example cap\n        df_plot = df_plot.dropna(subset=[full_metric_name]) # Drop NaN\n\n        if df_plot.empty:\n             print(f\"Cannot create boxplot for '{full_metric_name}': No valid finite data points.\")\n             return\n\n        title = f\"{full_metric_name} Distribution by Class\"\n        filename_base = f\"boxplot_{metric_abbr}\"\n\n        # --- 1. Save Seaborn Plot ---\n        fig_seaborn = None\n        try:\n            # Use sorted_datasets for consistent plot order\n            plot_order = self.sorted_datasets\n            # Filter df_plot to only include datasets present in the order AND have data\n            df_plot_filtered = df_plot[df_plot['Dataset'].isin(plot_order)]\n\n            if df_plot_filtered.empty:\n                 print(f\"Cannot create seaborn boxplot for '{full_metric_name}': No valid data points for specified dataset order.\")\n                 return\n\n            # Further ensure plot_order only contains datasets actually present in the filtered df\n            final_plot_order = [ds for ds in plot_order if ds in df_plot_filtered['Dataset'].unique()]\n            if not final_plot_order:\n                 print(f\"Cannot create seaborn boxplot for '{full_metric_name}': No datasets left after filtering.\")\n                 return\n\n            fig_seaborn, ax = plt.subplots(figsize=(max(8, len(final_plot_order) * 1.2), 6)) # Adjust width based on number of datasets\n            sns.boxplot(data=df_plot_filtered, x=\"Dataset\", y=full_metric_name, order=final_plot_order,\n                        showfliers=True, ax=ax) # Use the filtered data and specified order\n            ax.set_title(title + \" (Seaborn)\")\n            ax.set_xlabel(\"Dataset Class\")\n            ax.set_ylabel(full_metric_name)\n            plt.xticks(rotation=45, ha='right')\n            plt.tight_layout()\n\n            self._save_figure(fig_seaborn, filename_base) # Save the Seaborn figure\n\n        except Exception as e:\n             print(f\"Error creating/saving seaborn boxplot for {full_metric_name}: {e}\")\n        finally:\n             if fig_seaborn is not None:\n                 plt.close(fig_seaborn) # Close the Seaborn figure after saving\n\n        # --- 2. Display Plotly Plot (if interactive) ---\n        if self.viz_engine == \"interactive\":\n            try:\n                # Use sorted_datasets for consistent plot order in Plotly\n                plot_order_plotly = self.sorted_datasets\n                df_plot_filtered_plotly = df_plot[df_plot['Dataset'].isin(plot_order_plotly)] # Ensure filtering here too\n\n                if df_plot_filtered_plotly.empty:\n                    print(f\"Cannot create plotly boxplot for '{full_metric_name}': No valid data points for specified dataset order.\")\n                    return\n\n                # Further ensure plot_order only contains datasets actually present in the filtered df\n                final_plot_order_plotly = [ds for ds in plot_order_plotly if ds in df_plot_filtered_plotly['Dataset'].unique()]\n                if not final_plot_order_plotly:\n                     print(f\"Cannot create plotly boxplot for '{full_metric_name}': No datasets left after filtering.\")\n                     return\n\n\n                fig_plotly = px.box(df_plot_filtered_plotly, x=\"Dataset\", y=full_metric_name,\n                                    title=title + \" (Plotly)\",\n                                    category_orders={\"Dataset\": final_plot_order_plotly}, # Enforce order of available data\n                                    points=\"outliers\", # Show outliers like seaborn's showfliers=True\n                                    labels={\"Dataset\": \"Dataset Class\", full_metric_name: full_metric_name})\n                fig_plotly.update_layout(xaxis={'categoryorder':'array', 'categoryarray':final_plot_order_plotly}) # Another way to enforce order\n                fig_plotly.update_xaxes(tickangle=45)\n                fig_plotly.show() # Display the interactive Plotly figure\n\n            except Exception as e:\n                print(f\"Error creating/showing plotly boxplot for {full_metric_name}: {e}\")\n        # else: print(\"Skipping interactive boxplot display.\") # Optional message\n\n\n    # 3. Refactor _plot_summary_samples to handle a single dataset\n    def _plot_summary_samples_for_dataset(self, ds_name: str, metric_sort_by_abbr: str):\n        \"\"\"Displays and saves Top/Mid/Worst samples for a specific class. CORRECTED Best/Worst logic.\"\"\"\n        full_metric_name = self._get_full_metric_name(metric_sort_by_abbr)\n\n        details_list = self.sample_details.get(ds_name, [])\n        if not details_list:\n            print(f\"  No samples found for {ds_name} to plot best/median/worst.\")\n            return # Skip if no samples for this dataset\n\n        # Filter NaN/Inf for sorting metric\n        valid_details = [d for d in details_list if metric_sort_by_abbr in d and np.isfinite(d[metric_sort_by_abbr])]\n        if not valid_details:\n            print(f\"  No valid finite metric values ({metric_sort_by_abbr}) found for {ds_name} to sort summary samples.\")\n            return # Skip if no sortable samples\n\n        # Determine sort direction:\n        # reverse=True (descending) if lower is better (MAE, RMSE) -> puts highest error first\n        # reverse=False (ascending) if higher is better (SSIM, PSNR) -> puts lowest score first\n        sort_descending = metric_sort_by_abbr not in self.HIGHER_IS_BETTER_METRICS\n        sorted_list = sorted(valid_details, key=lambda d: d[metric_sort_by_abbr], reverse=sort_descending)\n\n        num_samples = len(sorted_list)\n        if num_samples == 0: return # Should be caught above, but safety check\n\n        # --- CORRECTED SELECTION LOGIC ---\n        # After sorting with `reverse=sort_descending`:\n        # Index 0 always contains the sample that is \"worst\" according to the metric\n        # (highest error for MAE/RMSE, lowest score for SSIM/PSNR).\n        # Index -1 always contains the sample that is \"best\" according to the metric\n        # (lowest error for MAE/RMSE, highest score for SSIM/PSNR).\n        worst_sample = sorted_list[0]    # The first element is always the worst\n        best_sample = sorted_list[-1]    # The last element is always the best\n        # --- END CORRECTION ---\n\n        mid_idx = num_samples // 2\n        # Ensure mid_idx is valid (can happen if num_samples is 1 or 2)\n        mid_idx = max(0, min(mid_idx, num_samples - 1)) # Prevent index error for small lists\n        mid_sample = sorted_list[mid_idx]\n\n\n        # Prepare the trio for plotting with correct sample assignment\n        trio = []\n        # Check if best/median/worst samples are valid before adding\n        if best_sample:\n             trio.append((f\"BEST ({metric_sort_by_abbr.upper()}={best_sample[metric_sort_by_abbr]:.3f})\", best_sample))\n        if mid_sample:\n             trio.append((f\"MEDIAN ({metric_sort_by_abbr.upper()}={mid_sample[metric_sort_by_abbr]:.3f})\", mid_sample))\n        if worst_sample:\n             trio.append((f\"WORST ({metric_sort_by_abbr.upper()}={worst_sample[metric_sort_by_abbr]:.3f})\", worst_sample))\n\n        # Handle cases where best/median/worst might be the same sample if few samples exist\n        unique_trio = []\n        seen_indices = set()\n        for label, sample in trio:\n            # Use a unique identifier like the original index in file + path to check uniqueness\n            # sample_id = (sample['seis_path'], sample['vel_path'], sample['sample_idx_in_file']) # Too complex maybe\n            # Let's use the metric value AND index as a proxy, good enough usually\n            sample_id = (sample[metric_sort_by_abbr], sample['sample_idx_in_file'])\n            if sample_id not in seen_indices:\n                unique_trio.append((label, sample))\n                seen_indices.add(sample_id)\n            # If it's a duplicate but we have fewer than 3 samples, allow it?\n            # For simplicity, let's just use unique ones. If only 1 sample, only \"BEST\" (or \"WORST\") will show.\n            # If only 2 samples, BEST/WORST will show.\n\n        if not unique_trio:\n             print(f\"  Could not identify distinct best/median/worst samples for {ds_name}.\")\n             return\n\n        fig = None # Initialize fig\n        try:\n            num_rows_plot = len(unique_trio)\n            print(f\"  Generating { '/'.join([t[0].split(' ')[0] for t in unique_trio]) } plot for: {ds_name} (by {full_metric_name})\")\n            # Adjust figsize and rows based on how many unique samples we have\n            fig, axs = plt.subplots(num_rows_plot, 3, figsize=(13, 4 * num_rows_plot), squeeze=False) # Ensure axs is 2D\n            fig.suptitle(f\"{ds_name}: {' / '.join([t[0].split(' ')[0] for t in unique_trio])} samples (by {full_metric_name})\\n\"\n                         f\"(File: SampleIdx)\", fontsize=14)\n            filename_base = f\"{ds_name}_summary_samples_by_{metric_sort_by_abbr}\"\n\n            plot_successful = False\n            for row_idx, (label, sample) in enumerate(unique_trio):\n                pred_np, true_np = sample[\"pred\"], sample[\"true\"]\n                if pred_np.ndim != 2 or true_np.ndim != 2:\n                     # print(f\"    Skipping row {row_idx+1} for {ds_name} summary plot: Invalid dimensions.\")\n                     for col_idx in range(3): axs[row_idx, col_idx].axis('off'); axs[row_idx, col_idx].set_visible(False)\n                     continue # Skip this row if data is bad\n\n                diff_np = pred_np - true_np\n                # Use the actual metric value for the label\n                sort_metric_val = sample.get(metric_sort_by_abbr, np.nan)\n                sample_info = f\"{os.path.basename(sample['vel_path'])}: {sample['sample_idx_in_file']}\"\n\n                # Handle potential NaNs in min/max calculation for color limits\n                valid_pred = pred_np[np.isfinite(pred_np)]\n                valid_true = true_np[np.isfinite(true_np)]\n                if valid_pred.size == 0 and valid_true.size == 0:\n                    vmin, vmax = 0, 1 # Fallback if no finite values\n                else:\n                    vmin = min(np.min(valid_pred) if valid_pred.size > 0 else np.inf,\n                               np.min(valid_true) if valid_true.size > 0 else np.inf)\n                    vmax = max(np.max(valid_pred) if valid_pred.size > 0 else -np.inf,\n                               np.max(valid_true) if valid_true.size > 0 else -np.inf)\n                    # Ensure vmin < vmax and handle edge case where min/max might still be inf\n                    if not np.isfinite(vmin) or not np.isfinite(vmax) or vmin >= vmax:\n                         vmin, vmax = 0, 1 # Reset to default if calculation failed\n\n\n                # Handle difference color limits\n                valid_diff = diff_np[np.isfinite(diff_np)]\n                if valid_diff.size == 0:\n                    vmin_diff, vmax_diff = -1, 1 # Fallback\n                else:\n                    diff_abs_max = np.max(np.abs(valid_diff))\n                    vmin_diff = -max(diff_abs_max, 1e-6) # Ensure non-zero range\n                    vmax_diff = max(diff_abs_max, 1e-6)\n\n\n                # Plotting logic\n                try:\n                    im0 = axs[row_idx, 0].imshow(pred_np, aspect=\"auto\", cmap=\"viridis\", vmin=vmin, vmax=vmax)\n                    title_str = f\"{label}\\nPredicted ({metric_sort_by_abbr.upper()}={sort_metric_val:.3f})\" if np.isfinite(sort_metric_val) else f\"{label}\\nPredicted\"\n                    axs[row_idx, 0].set_title(title_str + f\"\\n{sample_info}\", fontsize=9)\n                    plt.colorbar(im0, ax=axs[row_idx, 0], fraction=0.046, pad=0.04)\n                except Exception as plot_err:\n                    print(f\"Error plotting predicted for {label}: {plot_err}\")\n                    axs[row_idx, 0].set_title(f\"{label}\\nPredicted (Plot Error)\")\n                    axs[row_idx, 0].axis('off')\n\n\n                try:\n                    im1 = axs[row_idx, 1].imshow(true_np, aspect=\"auto\", cmap=\"viridis\", vmin=vmin, vmax=vmax)\n                    axs[row_idx, 1].set_title(\"Ground Truth\", fontsize=9)\n                    plt.colorbar(im1, ax=axs[row_idx, 1], fraction=0.046, pad=0.04)\n                except Exception as plot_err:\n                    print(f\"Error plotting ground truth for {label}: {plot_err}\")\n                    axs[row_idx, 1].set_title(\"Ground Truth (Plot Error)\")\n                    axs[row_idx, 1].axis('off')\n\n\n                try:\n                    im2 = axs[row_idx, 2].imshow(diff_np, aspect=\"auto\", cmap=\"RdBu\", vmin=vmin_diff, vmax=vmax_diff)\n                    axs[row_idx, 2].set_title(\"Difference (Pred - True)\", fontsize=9)\n                    plt.colorbar(im2, ax=axs[row_idx, 2], fraction=0.046, pad=0.04)\n                except Exception as plot_err:\n                    print(f\"Error plotting difference for {label}: {plot_err}\")\n                    axs[row_idx, 2].set_title(\"Difference (Plot Error)\")\n                    axs[row_idx, 2].axis('off')\n\n                for col_idx in range(3):\n                    if axs[row_idx, col_idx].get_title() and \"Plot Error\" not in axs[row_idx, col_idx].get_title():\n                        axs[row_idx, col_idx].set_xticks([]); axs[row_idx, col_idx].set_yticks([])\n                plot_successful = True # Mark success if at least one row tried plotting\n\n            if plot_successful:\n                plt.tight_layout(rect=[0, 0.03, 1, 0.95]) # Adjust for suptitle\n                self._save_figure(fig, filename_base) # Save\n                plt.show() # Display\n            else:\n                # print(f\"  Summary sample plot for {ds_name} not generated as no valid rows were plotted.\")\n                if fig: plt.close(fig) # Close if nothing was plotted but figure was created\n\n        except Exception as e:\n            print(f\"Error creating/saving/showing summary sample plot for {ds_name}: {e}\")\n            if fig: plt.close(fig) # Ensure closure on error\n        finally:\n             # Ensure figure is closed AFTER showing/saving or if error occurs\n             if fig is not None and plt.fignum_exists(fig.number):\n                 plt.close(fig)\n\n\n    # --- Diagnostic Grid Plotting Functions (accept 'ax') ---\n    # These remain largely the same, just ensure they handle empty/invalid data gracefully\n    # and return True on success, False on failure.\n\n    def _plot_residual_histogram(self, ds_name: str, ax: plt.Axes):\n        \"\"\"Plots histogram/KDE of residuals on given axis 'ax'.\"\"\"\n        if ds_name not in self.sample_details or not self.sample_details[ds_name]: return False # Indicate failure\n        all_residuals = []\n        for d in self.sample_details[ds_name]:\n             pred, true = d.get(\"pred\"), d.get(\"true\") # Use .get for safety\n             if pred is not None and true is not None and pred.ndim == 2 and true.ndim == 2 and pred.shape == true.shape:\n                 # Calculate residuals only for finite pairs\n                 mask = np.isfinite(pred) & np.isfinite(true)\n                 if np.any(mask):\n                    all_residuals.append((pred[mask] - true[mask]).flatten())\n        if not all_residuals: return False\n        try:\n            resids = np.concatenate(all_residuals)\n        except ValueError: # Handle case where all_residuals might be empty after filtering\n            return False\n        # Ensure finite values for stats and plotting (should be already, but extra check)\n        resids = resids[np.isfinite(resids)]\n        if resids.size == 0: return False\n\n        mean_resid, std_resid = np.mean(resids), np.std(resids)\n        sns.histplot(resids, kde=True, bins=50, ax=ax)\n        ax.set_title(f\"Residual Distribution\\nMean={mean_resid:.3f}, Std={std_resid:.3f}\", fontsize=10)\n        ax.set_xlabel(\"Velocity Error (Pred - True) (m/s)\")\n        ax.grid(True, linestyle='--', alpha=0.6)\n        return True # Indicate success\n\n\n    def _plot_depth_profile(self, ds_name: str, ax: plt.Axes, error_metric_abbr: str = 'mae'):\n        \"\"\"Plots average error profile vs. depth on given axis 'ax'.\"\"\"\n        if ds_name not in self.sample_details or not self.sample_details[ds_name]: return False\n        profiles, valid_sample_count, profile_length = [], 0, None\n        for d in self.sample_details[ds_name]:\n            pred, true = d.get(\"pred\"), d.get(\"true\")\n            if pred is None or true is None or pred.ndim != 2 or true.ndim != 2 or pred.shape != true.shape: continue\n            # Determine profile length from first valid sample\n            if profile_length is None: profile_length = pred.shape[0]\n            # Skip if depth dimension doesn't match\n            if pred.shape[0] != profile_length: continue\n\n            # Calculate error map based on metric, only for finite pairs\n            mask = np.isfinite(pred) & np.isfinite(true)\n            error_map = np.full_like(pred, np.nan) # Initialize with NaN\n\n            if error_metric_abbr == 'mae':\n                error_map[mask] = np.abs(pred[mask] - true[mask])\n            elif error_metric_abbr == 'rmse': # Actually MSE profile here\n                error_map[mask] = (pred[mask] - true[mask])**2\n            else: continue # Skip unsupported metrics\n\n            # Calculate profile (mean along horizontal axis), handle rows with only NaNs/Infs\n            with warnings.catch_warnings(): # Suppress mean of empty slice warning\n                warnings.simplefilter(\"ignore\", category=RuntimeWarning)\n                # nanmean ignores NaNs introduced above or originally present\n                profile = np.nanmean(error_map, axis=1)\n\n            # Check if the resulting profile has valid numbers and matches expected length\n            if profile.shape != (profile_length,) or not np.any(np.isfinite(profile)): continue\n\n            # Replace any remaining NaNs/Infs in the profile (e.g., if a whole row was NaN) with NaN\n            profile[~np.isfinite(profile)] = np.nan\n            profiles.append(profile)\n            valid_sample_count += 1\n\n        if not profiles or valid_sample_count == 0: return False\n\n        # Average the profiles, ignoring NaNs column-wise (depth-wise)\n        with warnings.catch_warnings():\n            warnings.simplefilter(\"ignore\", category=RuntimeWarning)\n            avg_profile = np.nanmean(profiles, axis=0)\n\n        # Check avg_profile for inf/nan before sqrt or plotting\n        if not np.any(np.isfinite(avg_profile)): return False # Nothing to plot if all NaN\n\n        depth_indices = np.arange(len(avg_profile))\n\n        if error_metric_abbr == 'rmse': # Take sqrt of the averaged MSE profile\n            # Take sqrt only of non-negative finite values\n            valid_mse_mask = np.isfinite(avg_profile) & (avg_profile >= 0)\n            plot_profile = np.full_like(avg_profile, np.nan)\n            plot_profile[valid_mse_mask] = np.sqrt(avg_profile[valid_mse_mask])\n            error_label = \"Mean RMSE\"\n        else: # MAE case\n            plot_profile = avg_profile # Already calculated as nanmean of absolute errors\n            error_label = \"Mean MAE\"\n\n        # Plot only the finite parts\n        valid_idx = np.isfinite(plot_profile)\n        if not np.any(valid_idx): return False\n\n        ax.plot(plot_profile[valid_idx], depth_indices[valid_idx])\n        ax.invert_yaxis()\n        ax.set_title(f\"Depth-wise Error ({error_label})\\n({valid_sample_count} samples)\", fontsize=10)\n        ax.set_xlabel(f\"{error_label} (m/s)\")\n        ax.set_ylabel(\"Depth Index\")\n        ax.grid(True, linestyle='--', alpha=0.6)\n        # Add horizontal line at 0 for reference\n        ax.axvline(0, color='grey', linestyle='--', linewidth=0.8, alpha=0.7)\n        return True\n\n\n    def _spectral_amplitude(self, arr: np.ndarray) -> np.ndarray:\n        \"\"\"Computes the 2D FFT amplitude spectrum (shifted). Handles non-finite inputs.\"\"\"\n        if arr.ndim != 2: return np.array([])\n        # Replace non-finite values with the mean of finite values (or 0 if none exist)\n        finite_vals = arr[np.isfinite(arr)]\n        fill_val = np.mean(finite_vals) if finite_vals.size > 0 else 0.0\n        arr_finite = np.nan_to_num(arr, nan=fill_val, posinf=fill_val, neginf=fill_val)\n        try:\n            # Ensure the array is not empty after nan_to_num (edge case)\n            if arr_finite.size == 0: return np.array([])\n            fft_result = np.fft.fft2(arr_finite)\n            return np.abs(np.fft.fftshift(fft_result))\n        except Exception: return np.array([])\n\n\n    def _plot_frequency_error(self, ds_name: str, ax: plt.Axes):\n        \"\"\"Plots mean log-amplitude spectrum of error on given axis 'ax'.\"\"\"\n        if ds_name not in self.sample_details or not self.sample_details[ds_name]: return False\n        err_spectra, valid_sample_count, spec_shape = [], 0, None\n        for d in self.sample_details[ds_name]:\n            pred, true = d.get(\"pred\"), d.get(\"true\")\n            if pred is None or true is None or pred.ndim != 2 or true.ndim != 2 or pred.shape != true.shape: continue\n            # Calculate error only where both are finite\n            mask = np.isfinite(pred) & np.isfinite(true)\n            err = np.full_like(pred, np.nan) # Initialize with NaN\n            err[mask] = pred[mask] - true[mask]\n\n            spec = self._spectral_amplitude(err) # Handles non-finite in err\n            if spec.size == 0: continue # Skip if spectral amplitude failed\n            if spec_shape is None: spec_shape = spec.shape\n            if spec.shape != spec_shape: continue # Ensure consistent shapes\n            err_spectra.append(spec)\n            valid_sample_count += 1\n        if not err_spectra or valid_sample_count == 0: return False\n\n        # Average the valid spectra (spectra should be all finite now)\n        mean_spec = np.mean(err_spectra, axis=0)\n\n        # Check if mean result is valid (should be if inputs were valid)\n        if not np.any(np.isfinite(mean_spec)): return False\n\n        im = ax.imshow(np.log1p(mean_spec), aspect='auto', cmap='magma') # log1p handles potential zeros in mean_spec\n        ax.set_title(f\"Mean Log-Amp Residual Spectrum\\n({valid_sample_count} samples)\", fontsize=10)\n        ax.set_xlabel(\"Frequency (kx - Shifted)\")\n        ax.set_ylabel(\"Frequency (ky - Shifted)\")\n        ax.set_xticks([])\n        ax.set_yticks([])\n        # Optional: add a small colorbar plt.colorbar(im, ax=ax, fraction=0.046, pad=0.04)\n        return True\n\n\n    def _plot_error_heatmap(self, ds_name: str, ax: plt.Axes):\n        \"\"\"Plots heatmap of average absolute error on given axis 'ax'.\"\"\"\n        if ds_name not in self.sample_details or not self.sample_details[ds_name]: return False\n        abs_resids, valid_sample_count, heatmap_shape = [], 0, None\n        for d in self.sample_details[ds_name]:\n            pred, true = d.get(\"pred\"), d.get(\"true\")\n            if pred is None or true is None or pred.ndim != 2 or true.ndim != 2 or pred.shape != true.shape: continue\n            if heatmap_shape is None: heatmap_shape = pred.shape\n            if pred.shape != heatmap_shape: continue\n\n            # Calculate absolute error only where both are finite\n            mask = np.isfinite(pred) & np.isfinite(true)\n            abs_error = np.full_like(pred, np.nan) # Initialize with NaN\n            abs_error[mask] = np.abs(pred[mask] - true[mask])\n\n            # Replace non-finite errors with NaN (already handled above)\n            # abs_error[~np.isfinite(abs_error)] = np.nan\n            if np.all(np.isnan(abs_error)): continue # Skip if whole error map is NaN\n\n            abs_resids.append(abs_error)\n            valid_sample_count += 1\n        if not abs_resids or valid_sample_count == 0: return False\n\n        # Average using nanmean\n        with warnings.catch_warnings():\n            warnings.simplefilter(\"ignore\", category=RuntimeWarning)\n            mean_abs_resid = np.nanmean(abs_resids, axis=0)\n\n        if not np.any(np.isfinite(mean_abs_resid)): return False # Check if result is valid\n\n        im = ax.imshow(mean_abs_resid, cmap='hot', aspect='auto')\n        ax.set_title(f\"Mean Absolute Error Heatmap\\n({valid_sample_count} samples)\", fontsize=10)\n        ax.set_xlabel(\"Horizontal Index\")\n        ax.set_ylabel(\"Depth Index\")\n        # Optional: add a small colorbar plt.colorbar(im, ax=ax, label='Mean |Error|', fraction=0.046, pad=0.04)\n        return True\n\n\n    def _scatter_error_vs_truth(self, ds_name: str, ax: plt.Axes):\n        \"\"\"Scatter plot of pixel absolute error vs. true velocity on given axis 'ax'.\"\"\"\n        if ds_name not in self.sample_details or not self.sample_details[ds_name]: return False\n        truth_vals, err_vals, valid_sample_count = [], [], 0\n        for d in self.sample_details[ds_name]:\n            pred, true = d.get(\"pred\"), d.get(\"true\")\n            if pred is None or true is None or pred.ndim != 2 or true.ndim != 2 or pred.shape != true.shape: continue\n            # Only include finite pairs for scatter plot\n            mask = np.isfinite(pred) & np.isfinite(true)\n            if not np.any(mask): continue # Skip if no finite pairs\n\n            abs_error = np.abs(pred[mask] - true[mask])\n            # Ensure errors are also finite (should be if pred/true are finite)\n            err_mask = np.isfinite(abs_error)\n            if not np.any(err_mask): continue\n\n            # Append only the finite error values and corresponding true values\n            err_vals.append(abs_error[err_mask])\n            truth_vals.append(true[mask][err_mask]) # Apply err_mask to the masked true values\n\n            valid_sample_count += 1\n\n        if not truth_vals or not err_vals or valid_sample_count == 0: return False\n\n        # Use try-except for concatenation as large arrays might cause memory issues\n        try:\n            truth = np.concatenate(truth_vals)\n            err = np.concatenate(err_vals)\n        except MemoryError:\n            print(f\"  MemoryError concatenating points for scatter plot ({ds_name}). Skipping.\")\n            return False\n        except ValueError: # Handle case where arrays might be empty after filtering\n             print(f\"  ValueError (likely empty arrays) concatenating points for scatter plot ({ds_name}). Skipping.\")\n             return False\n\n\n        # Double check for inf/nan after concatenation (unlikely but possible)\n        final_mask = np.isfinite(truth) & np.isfinite(err)\n        truth, err = truth[final_mask], err[final_mask]\n        if truth.size == 0: return False\n\n        max_points = 50000 # Keep subsampling\n        if len(truth) > max_points:\n             indices = np.random.choice(len(truth), max_points, replace=False)\n             truth, err = truth[indices], err[indices]\n\n        # Use rasterized=True for potentially dense plots to keep output sizes smaller\n        sns.scatterplot(x=truth, y=err, s=3, alpha=0.1, edgecolor=None, ax=ax, rasterized=True)\n        ax.set_title(f\"Error vs. Ground Truth\\n({len(truth)} points plotted from {valid_sample_count} samples)\", fontsize=10)\n        ax.set_xlabel(\"True Velocity (m/s)\")\n        ax.set_ylabel(f\"Absolute Error (m/s)\")\n        ax.grid(True, linestyle='--', alpha=0.6)\n        ax.axhline(0, color='grey', lw=0.5)\n        # Consider y-axis limits if errors are huge? ax.set_ylim(bottom=0, top=...)\n        return True\n\n\n    def _plot_mae_boxplot_single(self, ds_name: str, ax: plt.Axes):\n        \"\"\"Plots a single boxplot for MAE for a specific dataset on 'ax'.\"\"\"\n        full_mae_name = self._get_full_metric_name(\"mae\")\n        if self.all_samples_df is None or self.all_samples_df.empty or full_mae_name not in self.all_samples_df.columns: return False\n\n        # Prepare data: Filter, handle NaN/Inf\n        ds_data = self.all_samples_df[self.all_samples_df['Dataset'] == ds_name][[full_mae_name]].copy()\n        ds_data[full_mae_name] = ds_data[full_mae_name].replace([np.inf, -np.inf], np.nan)\n        ds_data = ds_data.dropna()\n        if ds_data.empty: return False\n\n        sns.boxplot(y=ds_data[full_mae_name], ax=ax, showfliers=True, width=0.4) # Make box narrower\n        ax.set_title(f\"{full_mae_name}\", fontsize=10)\n        ax.set_ylabel(\"(m/s)\") # Keep ylabel concise\n        ax.set_xlabel(ds_name) # Label x-axis with dataset name\n        ax.set_xticks([]) # Remove x-ticks as the label serves the purpose\n        ax.grid(True, axis='y', linestyle='--', alpha=0.6)\n        return True\n\n\n    def _plot_diagnostic_grid(self, ds_name: str):\n        \"\"\"Creates, saves, and displays the 2x3 diagnostic grid for a single dataset.\"\"\"\n        print(f\"  Generating diagnostic grid for: {ds_name}\")\n        fig = None # Initialize fig\n        try:\n            fig, axs = plt.subplots(2, 3, figsize=(16, 9)) # Slightly adjusted figsize\n            fig.suptitle(f\"Diagnostic Plots - {ds_name}\", fontsize=16, y=1.02) # Adjust title position\n            filename_base = f\"{ds_name}_diagnostic_grid\"\n\n            plot_funcs = [\n                self._plot_residual_histogram,\n                lambda ds, ax: self._plot_depth_profile(ds, ax, error_metric_abbr='mae'), # Pass MAE explicitly\n                self._plot_frequency_error,\n                self._plot_error_heatmap,\n                self._scatter_error_vs_truth,\n                self._plot_mae_boxplot_single # Use MAE boxplot here\n            ]\n            plot_names = [ # For error messages\n                \"Residual Histogram\", \"Depth Profile (MAE)\", \"Frequency Error\",\n                \"Error Heatmap\", \"Error vs Truth Scatter\", \"MAE Boxplot\"\n            ]\n\n\n            any_plot_successful = False\n            for i, (plot_func, plot_name) in enumerate(zip(plot_funcs, plot_names)):\n                row, col = divmod(i, 3)\n                ax = axs[row, col]\n                try:\n                    # Pass ds_name and ax to the plot function\n                    success = plot_func(ds_name, ax=ax) # Call the specific plot function\n                    if success:\n                        any_plot_successful = True\n                    else: # If function indicated failure (e.g., no data), hide the axis\n                        # print(f\"    Skipping plot '{plot_name}' for {ds_name}: No valid data or plot failed.\")\n                        ax.set_visible(False)\n                except Exception as e:\n                    print(f\"    Error generating plot '{plot_name}' for {ds_name}: {e}\")\n                    # Display error on the plot itself and hide axes\n                    ax.text(0.5, 0.5, 'Plot Error', horizontalalignment='center', verticalalignment='center', transform=ax.transAxes, color='red', fontsize=10)\n                    ax.set_xticks([])\n                    ax.set_yticks([])\n                    # Optionally make frame invisible too: ax.spines[:].set_visible(False)\n\n            if any_plot_successful:\n                plt.tight_layout(rect=[0, 0, 1, 0.97]) # Adjust layout for suptitle\n                self._save_figure(fig, filename_base) # Save\n                plt.show() # Display\n            else:\n                print(f\"  Skipping diagnostic grid display/save for {ds_name} - no plots could be generated.\")\n                # Close the potentially empty figure if it was created\n                if fig is not None and plt.fignum_exists(fig.number):\n                    plt.close(fig)\n\n        except Exception as e:\n            print(f\"Error creating diagnostic grid layout for {ds_name}: {e}\")\n            # Ensure figure is closed if layout creation failed mid-way\n            if fig is not None and plt.fignum_exists(fig.number): plt.close(fig)\n        finally:\n             # Close the figure if it exists and was potentially shown/saved\n             if fig is not None and plt.fignum_exists(fig.number):\n                 plt.close(fig)\n\n    # --- 4. Helper for Cross-Family Grid Plots ---\n    def _create_cross_family_grid(\n        self,\n        plot_function: Callable[[str, plt.Axes], bool],\n        plot_name: str,\n        filename_base: str,\n        fig_title: str,\n        cols: int = 5 # Default to 5 columns\n    ):\n        \"\"\"\n        Generates a grid of plots, one for each dataset family, using the provided plot function.\n\n        Args:\n            plot_function: A plotting function (like _plot_residual_histogram) that accepts\n                           (ds_name, ax) and returns True on success, False on failure.\n            plot_name: A descriptive name for the plot type (for messages).\n            filename_base: Base name for the saved PNG file.\n            fig_title: The main title for the entire figure grid.\n            cols: Number of columns in the grid layout.\n        \"\"\"\n        datasets_to_plot = [ds for ds in self.sorted_datasets if ds in self.sample_details and self.sample_details[ds]]\n        if not datasets_to_plot:\n            print(f\"Skipping cross-family grid '{plot_name}': No datasets with data.\")\n            return\n\n        n_datasets = len(datasets_to_plot)\n        cols = min(cols, n_datasets) # Don't have more columns than datasets\n        if cols <= 0: return # Skip if no datasets or cols=0\n        rows = math.ceil(n_datasets / cols)\n        fig_height = rows * 4 # Adjust multiplier as needed for aspect ratio\n        fig_width = cols * 4 # Adjust multiplier as needed\n\n        fig = None\n        try:\n            print(f\"\\n--- Generating Cross-Family Grid: {plot_name} ---\")\n            fig, axs = plt.subplots(rows, cols, figsize=(fig_width, fig_height), squeeze=False) # Ensure axs is always 2D array\n            # Adjust title y position based on rows/layout\n            title_y_pos = 0.98 if rows <= 2 else 1.0\n            fig.suptitle(fig_title, fontsize=16, y=title_y_pos)\n\n            any_plot_successful = False\n            for i, ds_name in enumerate(datasets_to_plot):\n                row, col = divmod(i, cols)\n                ax = axs[row, col]\n                try:\n                    # Temporarily set title during function call for context\n                    ax.set_title(ds_name, fontsize=12) # Set subplot title to dataset name BEFORE plotting\n                    success = plot_function(ds_name, ax=ax)\n                    if success:\n                        any_plot_successful = True\n                        # Keep title if successful\n                    else:\n                        ax.set_visible(False) # Hide axes if plot failed for this dataset\n                        ax.set_title(\"\") # Clear title if failed\n                except Exception as e:\n                    print(f\"    Error in cross-family grid for {ds_name} ({plot_name}): {e}\")\n                    ax.clear() # Clear potential partial plot\n                    ax.text(0.5, 0.5, f'{ds_name}\\nPlot Error', horizontalalignment='center', verticalalignment='center', transform=ax.transAxes, color='red', fontsize=10)\n                    ax.set_xticks([])\n                    ax.set_yticks([])\n                    ax.set_title(\"\") # Clear title on error\n\n            # Hide unused axes\n            for i in range(n_datasets, rows * cols):\n                row, col = divmod(i, cols)\n                axs[row, col].set_visible(False)\n\n            if any_plot_successful:\n                plt.tight_layout(rect=[0, 0.01, 1, 0.95 if title_y_pos > 0.98 else 0.97]) # Adjust layout rect based on title pos\n                self._save_figure(fig, filename_base)\n                plt.show()\n            else:\n                print(f\"Cross-family grid '{plot_name}' not generated as no subplots were successful.\")\n                if fig is not None and plt.fignum_exists(fig.number): plt.close(fig)\n\n        except Exception as e:\n            print(f\"Error creating cross-family grid layout for '{plot_name}': {e}\")\n            if fig is not None and plt.fignum_exists(fig.number): plt.close(fig)\n        finally:\n            if fig is not None and plt.fignum_exists(fig.number): plt.close(fig)\n\n\n    # --- 3. Restructure generate_report ---\n    def generate_report(self):\n        \"\"\"\n        Generates all summaries and plots in the specified order:\n        1. Summary Stats Table\n        2. Overall Metric Boxplots\n        3. For EACH dataset:\n            a. Diagnostic Grid\n            b. Best/Median/Worst Samples Plot\n        4. Cross-Family Comparison Grids (at the end)\n        \"\"\"\n        # 1. Calculate stats, create DataFrames, save CSVs\n        self._summarize_stats()\n\n        # 2. Display interactive styled summary table\n        self._display_summary_table()\n\n        # Check if analysis yielded any results before proceeding with plots\n        if not self.sample_details or self.all_samples_df is None or self.all_samples_df.empty:\n             print(\"\\nNo samples analyzed successfully or DataFrame empty. Skipping visualization generation.\")\n             # Call the function to identify other potential errors\n             self.comment_on_potential_errors()\n             return\n\n        print(\"\\n--- Generating Visualizations (saving and displaying) ---\")\n\n        # 3. Overall Boxplots for key metrics (across all datasets)\n        print(\"\\nGenerating Overall Metric Boxplots...\")\n        metrics_to_plot_abbr = [m for m in sorted(list(self.metrics_to_compute)) # Plot in consistent order\n                                if self._get_full_metric_name(m) in self.all_samples_df.columns]\n        valid_metrics_found_for_boxplot = False\n        for metric_abbr in metrics_to_plot_abbr:\n            # Check if data is valid before attempting plot\n            full_name = self._get_full_metric_name(metric_abbr)\n            # Check for any finite values after handling inf/nan\n            temp_series = self.all_samples_df[full_name].replace([np.inf, -np.inf], np.nan).dropna()\n            if not temp_series.empty:\n                 self._create_metric_boxplot(metric_abbr)\n                 valid_metrics_found_for_boxplot = True\n        if not valid_metrics_found_for_boxplot:\n             print(\" No metrics with valid finite data found for overall boxplots.\")\n\n        # 4. Per-dataset detailed plots (Grid THEN Best/Median/Worst)\n        print(\"\\nGenerating Per-Dataset Diagnostic Grid and Summary Samples...\")\n        datasets_with_data = [ds for ds in self.sorted_datasets if ds in self.sample_details and self.sample_details[ds]]\n\n        if not datasets_with_data:\n             print(\" No datasets have analyzed samples for detailed diagnostic plots.\")\n        else:\n             # Determine the sort metric for Best/Median/Worst ONCE\n             # Prioritize MAE if available and valid\n             sort_metric_abbr = \"mae\"\n             full_sort_metric_name = self._get_full_metric_name(sort_metric_abbr)\n             can_sort = False\n\n             if sort_metric_abbr in self.metrics_to_compute and full_sort_metric_name in self.all_samples_df.columns:\n                 # Check for finite values in the sorting column across all valid datasets\n                 metric_series = self.all_samples_df[self.all_samples_df['Dataset'].isin(datasets_with_data)][full_sort_metric_name]\n                 if metric_series.replace([np.inf, -np.inf], np.nan).notna().any():\n                      can_sort = True\n                      print(f\" Sorting Best/Median/Worst samples by {full_sort_metric_name}.\")\n\n             # Fallback if primary sort metric (MAE) failed or not computed\n             if not can_sort:\n                 print(f\" Primary sort metric {full_sort_metric_name} not computed or has no valid finite values.\")\n                 # Look for *any* valid metric to sort by, preferring lower-is-better, then higher-is-better\n                 fallback_candidates = []\n                 # Prefer lower-is-better metrics first (RMSE)\n                 for m_abbr in sorted([m for m in self.metrics_to_compute if m not in self.HIGHER_IS_BETTER_METRICS and m != 'mae']):\n                      m_full = self._get_full_metric_name(m_abbr)\n                      if m_full in self.all_samples_df.columns:\n                         metric_series = self.all_samples_df[self.all_samples_df['Dataset'].isin(datasets_with_data)][m_full]\n                         if metric_series.replace([np.inf, -np.inf], np.nan).notna().any():\n                            fallback_candidates.append(m_abbr)\n                 # Then check higher-is-better metrics (SSIM, PSNR)\n                 for m_abbr in sorted([m for m in self.metrics_to_compute if m in self.HIGHER_IS_BETTER_METRICS]):\n                      m_full = self._get_full_metric_name(m_abbr)\n                      if m_full in self.all_samples_df.columns:\n                         metric_series = self.all_samples_df[self.all_samples_df['Dataset'].isin(datasets_with_data)][m_full]\n                         if metric_series.replace([np.inf, -np.inf], np.nan).notna().any():\n                            fallback_candidates.append(m_abbr)\n\n                 if fallback_candidates:\n                      sort_metric_abbr = fallback_candidates[0] # Pick the first valid fallback\n                      full_sort_metric_name = self._get_full_metric_name(sort_metric_abbr)\n                      can_sort = True\n                      print(f\" Falling back to sorting Best/Median/Worst samples by {full_sort_metric_name}.\")\n                 else:\n                      print(\" No valid metrics available to sort Best/Median/Worst samples. Skipping these plots.\")\n\n             # Now loop through datasets and plot\n             for ds_name in datasets_with_data: # Iterate in sorted order\n                 print(f\"\\n--- Processing Dataset: {ds_name} ---\")\n                 # a. Diagnostic Grid\n                 self._plot_diagnostic_grid(ds_name)\n                 # b. Best/Median/Worst samples plot (if sorting is possible)\n                 if can_sort:\n                     self._plot_summary_samples_for_dataset(ds_name, metric_sort_by_abbr=sort_metric_abbr)\n                 else:\n                     # Message printed once above if cannot sort\n                     pass\n                     # print(f\"  Skipping Best/Median/Worst for {ds_name} due to lack of valid sort metric.\")\n\n        # 5. Cross-Family Comparison Grids (at the end)\n        print(\"\\n--- Generating Cross-Family Comparison Grids ---\")\n        # Determine number of columns based on datasets actually plotted\n        n_datasets_plotted = len(datasets_with_data)\n        n_cols_grid = min(5, n_datasets_plotted) if n_datasets_plotted > 0 else 1\n\n        # Grid of Residual Distributions\n        self._create_cross_family_grid(\n            plot_function=self._plot_residual_histogram,\n            plot_name=\"Residual Distributions\",\n            filename_base=\"cross_family_residual_histograms\",\n            fig_title=\"Residual Distributions Across Families\",\n            cols=n_cols_grid\n        )\n\n        # Grid of Depth-wise Errors (MAE)\n        self._create_cross_family_grid(\n            plot_function=lambda ds, ax: self._plot_depth_profile(ds, ax, error_metric_abbr='mae'),\n            plot_name=\"Depth-wise MAE\",\n            filename_base=\"cross_family_depth_mae\",\n            fig_title=\"Mean Absolute Error vs. Depth Across Families\",\n            cols=n_cols_grid\n        )\n\n        # Grid of Mean Log-Amp Residual Spectrums\n        self._create_cross_family_grid(\n            plot_function=self._plot_frequency_error,\n            plot_name=\"Residual Spectrums\",\n            filename_base=\"cross_family_residual_spectrums\",\n            fig_title=\"Mean Log-Amplitude Residual Spectrum Across Families\",\n            cols=n_cols_grid\n        )\n\n        # Grid of Mean Absolute Error Heatmaps\n        self._create_cross_family_grid(\n            plot_function=self._plot_error_heatmap,\n            plot_name=\"MAE Heatmaps\",\n            filename_base=\"cross_family_mae_heatmaps\",\n            fig_title=\"Mean Absolute Error Heatmap Across Families\",\n            cols=n_cols_grid\n        )\n\n        # Grid of Error vs. Ground Truth Scatter Plots\n        self._create_cross_family_grid(\n            plot_function=self._scatter_error_vs_truth,\n            plot_name=\"Error vs. Ground Truth Scatter\",\n            filename_base=\"cross_family_error_vs_truth_scatter\",\n            fig_title=\"Absolute Error vs. True Velocity Across Families (Sampled)\",\n            cols=n_cols_grid\n        )\n\n        # --- End of Report Generation ---\n        print(\"\\n--- All Report Generation Steps Complete ---\")\n        # Call the function to identify other potential errors\n        self.comment_on_potential_errors()\n\n\n\n    def get_worst_cases(self, k: int = 20, metric_abbr: str = \"mae\") -> List[Dict[str, Any]]:\n        \"\"\"Returns the top 'k' samples with the 'worst' score based on the specified metric. Uses sorted dataset order.\"\"\"\n        full_metric_name = self._get_full_metric_name(metric_abbr)\n        if metric_abbr not in self.metrics_to_compute:\n            print(f\"Warning: Metric '{full_metric_name}' not computed.\")\n            return []\n        if not self.sample_details:\n             print(\"Warning: No analysis details available.\")\n             return []\n\n        # Flatten list respecting sorted dataset order\n        all_samples_flat = [sample for ds_name in self.sorted_datasets\n                            for sample in self.sample_details.get(ds_name, [])]\n        # Filter for finite values of the metric\n        valid_samples = [s for s in all_samples_flat if metric_abbr in s and np.isfinite(s[metric_abbr])]\n        if not valid_samples:\n             print(f\"No valid finite samples found for metric '{full_metric_name}'.\")\n             return []\n\n        # Determine sort direction for finding the \"worst\"\n        # We want the highest error (MAE/RMSE) or lowest score (SSIM/PSNR) first.\n        # This is exactly what `reverse=sort_descending` does in the _plot_summary function.\n        sort_descending = metric_abbr not in self.HIGHER_IS_BETTER_METRICS\n        valid_samples.sort(key=lambda x: x[metric_abbr], reverse=sort_descending)\n\n        num_to_show = min(k, len(valid_samples))\n        # Label clarifies if \"worst\" means high error or low score\n        worst_label = \"Highest Error\" if sort_descending else \"Lowest Score\"\n        print(f\"\\n--- Top {num_to_show} Worst Cases (by {full_metric_name} - {worst_label}) ---\")\n        result_k = valid_samples[:num_to_show]\n        for i, sample in enumerate(result_k):\n            print(f\" #{i+1}: Dataset={sample['dataset']}, File={os.path.basename(sample['vel_path'])}, \"\n                  f\"SampleIdx={sample['sample_idx_in_file']}, {full_metric_name}={sample[metric_abbr]:.4f}\")\n        return result_k\n\n\n    def plot_sample_prediction(\n        self,\n        dataset_name: str,\n        sample_index: int = 0,\n    ):\n        \"\"\"Visualizes and saves a specific analyzed sample's prediction vs. ground truth.\"\"\"\n        # Ensure dataset_name exists in sample_details (which respects sorted order)\n        if dataset_name not in self.sample_details:\n            # Check if it was discovered but had no samples processed\n            if dataset_name in self._pairs:\n                 print(f\"Warning: Dataset '{dataset_name}' exists but no samples were analyzed (check max_samples_per_class or file errors). Cannot plot.\")\n                 return\n            else:\n                 # Check if it's a typo vs the available datasets\n                 print(f\"Error: Dataset '{dataset_name}' not found. Available datasets: {self.sorted_datasets}\")\n                 return # More informative than raising ValueError\n\n        details_list = self.sample_details[dataset_name]\n        if not details_list:\n             print(f\"No samples were analyzed for dataset '{dataset_name}'. Cannot plot index {sample_index}.\")\n             return\n        if not 0 <= sample_index < len(details_list):\n             print(f\"Error: Index {sample_index} out of range for {len(details_list)} analyzed samples in {dataset_name}.\")\n             return # More informative than raising IndexError\n\n        sample = details_list[sample_index]\n        pred_np, true_np = sample.get(\"pred\"), sample.get(\"true\") # Use get for safety\n        if pred_np is None or true_np is None:\n             print(f\"Error plotting sample {sample_index} ({dataset_name}): Stored 'pred' or 'true' data is missing.\")\n             return\n        if pred_np.ndim != 2 or true_np.ndim != 2:\n             print(f\"Error plotting sample {sample_index} ({dataset_name}): Stored 'pred' ({pred_np.shape}) or 'true' ({true_np.shape}) data is not 2D.\")\n             return\n\n        file_info = f\"{os.path.basename(sample['vel_path'])} (Orig Idx {sample['sample_idx_in_file']})\"\n        filename_base = f\"{dataset_name}_sample_{sample_index}_pred_vs_true\"\n\n        fig = None # Initialize fig\n        try:\n            fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n\n            # --- Robust Color Limit Calculation ---\n            valid_pred = pred_np[np.isfinite(pred_np)]\n            valid_true = true_np[np.isfinite(true_np)]\n            if valid_pred.size == 0 and valid_true.size == 0: vmin, vmax = 0, 1\n            else:\n                vmin = min(np.min(valid_pred) if valid_pred.size > 0 else np.inf,\n                           np.min(valid_true) if valid_true.size > 0 else np.inf)\n                vmax = max(np.max(valid_pred) if valid_pred.size > 0 else -np.inf,\n                           np.max(valid_true) if valid_true.size > 0 else -np.inf)\n                if not np.isfinite(vmin) or not np.isfinite(vmax) or vmin >= vmax: vmin, vmax = 0, 1\n\n            # --- Robust Difference Calculation and Limits ---\n            diff = np.full_like(pred_np, np.nan)\n            mask = np.isfinite(pred_np) & np.isfinite(true_np)\n            diff[mask] = pred_np[mask] - true_np[mask]\n            valid_diff = diff[np.isfinite(diff)]\n            if valid_diff.size == 0: vmin_diff, vmax_diff = -1, 1\n            else:\n                diff_abs_max = np.max(np.abs(valid_diff))\n                vmin_diff = -max(diff_abs_max, 1e-6)\n                vmax_diff = max(diff_abs_max, 1e-6)\n\n            # Plotting with error handling per subplot\n            try:\n                im0 = axes[0].imshow(pred_np, aspect='auto', cmap='viridis', vmin=vmin, vmax=vmax)\n                axes[0].set_title(\"Predicted Velocity\")\n                plt.colorbar(im0, ax=axes[0], fraction=0.046, pad=0.04)\n            except Exception as e: axes[0].set_title(\"Predicted (Plot Error)\"); axes[0].axis('off')\n\n            try:\n                im1 = axes[1].imshow(true_np, aspect='auto', cmap='viridis', vmin=vmin, vmax=vmax)\n                axes[1].set_title(\"Ground Truth Velocity\")\n                plt.colorbar(im1, ax=axes[1], fraction=0.046, pad=0.04)\n            except Exception as e: axes[1].set_title(\"Ground Truth (Plot Error)\"); axes[1].axis('off')\n\n            try:\n                im2 = axes[2].imshow(diff, aspect='auto', cmap='RdBu', vmin=vmin_diff, vmax=vmax_diff)\n                axes[2].set_title(f\"Difference (Pred - True)\")\n                plt.colorbar(im2, ax=axes[2], fraction=0.046, pad=0.04)\n            except Exception as e: axes[2].set_title(\"Difference (Plot Error)\"); axes[2].axis('off')\n\n\n            for ax in axes:\n                if \"Plot Error\" not in (ax.get_title() or \"\"):\n                    ax.set_xticks([]); ax.set_yticks([])\n\n            # Title\n            title = f\"{dataset_name}, Analyzed Sample #{sample_index} ({file_info})\\n\"\n            metrics_str = []\n            for m_abbr in sorted(list(self.metrics_to_compute)): # Consistent metric order\n                 val = sample.get(m_abbr, np.nan)\n                 if np.isfinite(val):\n                     # More sophisticated formatting based on metric?\n                     format_str = \".2f\" if m_abbr == 'psnr' else \".4f\" # Use more precision for others\n                     metrics_str.append(f\"{self._get_full_metric_name(m_abbr)}={val:{format_str}}\")\n            title += \", \".join(metrics_str)\n            plt.suptitle(title, fontsize=12)\n            plt.tight_layout(rect=[0, 0.03, 1, 0.92]) # Adjust rect for suptitle\n\n            # Save and Show\n            self._save_figure(fig, filename_base)\n            plt.show()\n\n        except Exception as e:\n            print(f\"Error plotting sample prediction for {dataset_name} index {sample_index}: {e}\")\n        finally:\n            # Close figure AFTER showing or if error\n            if fig is not None and plt.fignum_exists(fig.number):\n                plt.close(fig)\n\n    def comment_on_potential_errors(self):\n        return","metadata":{"_uuid":"7006da94-abc3-4d52-9579-918fce02d298","_cell_guid":"3532edd1-dbb0-44b5-99d5-b86754bd1ed9","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-05-13T02:08:24.041793Z","iopub.execute_input":"2025-05-13T02:08:24.041985Z","iopub.status.idle":"2025-05-13T02:08:36.197930Z","shell.execute_reply.started":"2025-05-13T02:08:24.041969Z","shell.execute_reply":"2025-05-13T02:08:36.197353Z"},"jupyter":{"outputs_hidden":false},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Usage Cell - Extremely Simple.","metadata":{}},{"cell_type":"code","source":"analyzer = FWIPredictionAnalyzer(\n    root_dir=\"/kaggle/input/waveform-inversion/train_samples\",\n    model=model,\n    device=\"cuda\" if torch.cuda.is_available() else \"cpu\",\n    max_samples_per_class=500,\n    viz_engine=\"seaborn\"  # or \"interactive\"\n)\nsummary_stats = analyzer.analyze()","metadata":{"_uuid":"053f7f35-bd83-4213-bb6a-eedf1dfa83cf","_cell_guid":"25934d5d-5530-4caf-b7e2-0836db09f066","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-05-13T02:08:36.199422Z","iopub.execute_input":"2025-05-13T02:08:36.199799Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Analyzing the Residual Histograms","metadata":{"_uuid":"064489bf-3ddb-458b-b000-36d7fcfcd676","_cell_guid":"d893ce8b-a57d-4862-b952-6c0aac82311e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"from IPython.display import Image, display\ndisplay(Image(filename=\"/kaggle/working/fwi_analysis_output/cross_family_residual_histograms.png\"))","metadata":{"_uuid":"e97a8002-6178-41a5-b2c2-6819a82979f5","_cell_guid":"cc6dc42c-4d56-4ff0-ad72-fe4dc646223f","trusted":true,"collapsed":false,"_kg_hide-input":true,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n| Term on the plot                 | Plain-English meaning                                                                                                                                                      |\n| -------------------------------- | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |\n| **Velocity Error (Pred − True)** | *Residual* = (model prediction) − (actual value) for every grid-cell, measured in metres per second (m/s). A residual of **0 m/s** means the cell was predicted perfectly. |\n| **Count**                        | How many grid-cells fell into each error bin.                                                                                                                              |\n| **One panel per folder**         | Each subplot pools every residual coming from **one geological sub-family** (e.g. `FlatVel_A`, `CurveFault_B`, …).                                                         |\n\n> **Concrete example**\n> Imagine the true velocity at some point underground is **3 200 m/s** but the model predicts **3 050 m/s**.\n> Residual = ▶ 3 050 − 3 200 = **−150 m/s** (an under-prediction).\n> That single number lands in the “−150 m/s” bar of its family’s histogram.\n\n---\n\n### 2.  How to Read a Residual Distribution (Beginner-Friendly)\n\n1. **Peak position (bias)**\n   *Does the spike sit at 0 m/s?*\n\n   * **Centered at 0** → on average the model is not systematically high or low.\n   * **Shifted left** (negative) → tends to *under-predict* velocities.\n   * **Shifted right** (positive) → tends to *over-predict*.\n\n2. **Width (variance)**\n   *How fat are the tails?*\n\n   * **Narrow spike** → most errors are small → contributes little to MAE.\n   * **Wide spread / heavy tails** → large occasional errors → inflates MAE.\n\n3. **Skewness / asymmetry**\n   One tail longer than the other? That tail is where costly mistakes live.\n\n4. **Outliers**\n   Low-density bars far from 0 indicate rare but severe misses.\n\n---\n\n### 3.  What the Ten Panels Are Telling Us\n\n| Family                            | Visual cues                                    | Interpretation                                                                              |\n| --------------------------------- | ---------------------------------------------- | ------------------------------------------------------------------------------------------- |\n| **FlatVel\\_A & FlatVel\\_B**       | Razor-thin peaks hugging 0                     | Excellent fit on the simplest layered cases.                                                |\n| **CurveVel\\_A & CurveVel\\_B**     | Slightly wider, peak \\~ −50 m/s                | The model *slightly under-estimates* velocities where layers curve; still quite accurate.   |\n| **FlatFault\\_A & FlatFault\\_B**   | Narrow but peak sits ≈ −100 m/s                | Consistent under-prediction around faults; probably misses the velocity jump across faults. |\n| **CurveFault\\_A & CurveFault\\_B** | Wider than “Flat”, visible left tail           | Curvature *plus* faults add difficulty; errors larger and more negative.                    |\n| **Style\\_A & Style\\_B**           | Tails stretching ±400 m/s, Style\\_B the widest | The “style-transfer” textures are much harder; these two buckets dominate your MAE.         |\n\n*(Exact numbers are approximate; the qualitative shapes carry the message.)*\n\n---\n\n### 4.  Why This Matters for **MAE**\n\n$$\n\\text{MAE} \\;=\\; \\frac{1}{N}\\,\\sum_{i=1}^{N} \\bigl| \\text{Residual}_i \\bigr|\n$$\n\n**Translation:**\n“Take every residual, drop the minus sign, add them all up, and divide by how many there are.”\n\n*Implication:* wide or biased residuals in any family **directly raise** the MAE. From the histograms:\n\n* Biggest contributors → **Style\\_B, Style\\_A, CurveFault\\_B** (widest spreads).\n* Systematic bias → negative shift in most *Fault* buckets (every cell adds ≈ 50-100 m/s to MAE even if the spread is narrow).\n\n---\n\n### 5.  Practical Tweaks to Drive MAE Down\n\n| Idea                                                                                                                                               | How it helps                                                                       | Pros                                 | Cons / Trade-offs                                                             |\n| -------------------------------------------------------------------------------------------------------------------------------------------------- | ---------------------------------------------------------------------------------- | ------------------------------------ | ----------------------------------------------------------------------------- |\n| **Family-specific fine-tuning**<br>Train a small head or specialist model per geological bucket.                                                   | Removes one-size-fits-all bias; lets each head learn its quirks.                   | Simple to bolt on; quick bias win.   | Extra parameters; risk of over-fitting if some buckets are small.             |\n| **Bias-correction layer**<br>Add a learnable scalar or shallow network whose sole job is to shift predictions so the mean residual ≈ 0.            | Eliminates systematic under-prediction seen in Fault families.                     | Cheap; no extra data needed.         | Only fixes *mean* error, not spread.                                          |\n| **Loss re-weighting / focal MAE**<br>Give higher weight to Style & CurveFault samples during training.                                             | Forces the optimiser to pay more attention to the worst cases.                     | Targets the “tail” that drives MAE.  | Might hurt performance on easy families; needs careful tuning.                |\n| **Multi-scale architecture**<br>Add wavelet or UNet-style skip paths so the net “sees” both fine textures (Style) and large blocks (Fault planes). | Captures high-frequency patterns lurking in Style\\_B.                              | Physics-motivated; generally robust. | Heavier model; longer training time.                                          |\n| **Edge-aware or gradient loss**<br>Add a term that penalises velocity jumps if the model blurs faults.                                             | Encourages sharper fault boundaries → reduces under-prediction right at the fault. | Directly addresses bias origin.      | Choosing weight for extra loss needs validation.                              |\n| **Data augmentation (texture synthesis)**                                                                                                          | Increases exposure to Style-like randomness, shrinking the tails.                  | Cheap once implemented.              | Synthetic textures must resemble Style family or risk hurting generalisation. |\n| **Switch to a *Huber* or *Charbonnier* loss during pre-training**                                                                                  | Less sensitive to narrow easy peaks, more pressure on large residuals.             | Naturally pulls in heavy tails.      | Huber needs δ hyper-parameter; may converge slower.                           |\n| **Ensemble of physics-guided + data-driven models**                                                                                                | Different error modes cancel, lowering MAE.                                        | Proven trick in competitions.        | High compute; ensemble diversity must be real, not cosmetic.                  |\n\n---\n\n### 6.  Putting the Pieces Together—A Suggested Roadmap\n\n1. **Quick win (days):**\n\n   * Compute the mean residual **per family** on a validation split and subtract that bias at inference.\n   * Re-train with a *bias-correction layer*; verify MAE drop.\n\n2. **Medium term (weeks):**\n\n   * Duplicate the last few layers and fine-tune separately for **Style** scenarios (use a family ID flag).\n   * Introduce a *Huber* loss with δ chosen on the validation set; track tail shrinkage.\n\n3. **Longer term (competition life-cycle):**\n\n   * Redesign the backbone to be *multi-scale* (e.g. UNet or ConvNeXt with dilated blocks).\n   * Augment the training set with procedurally generated “Style-like” textures; oversample CurveFault\\_B.\n   * Ensemble a physics-based adjoint inversion with your neural network predictions.\n\n---\n\n### 7.  Key Take-aways for Beginners\n\n* **Histogram’s job:** show where and how often the model is wrong.\n* **Perfect model:** spike at 0 m/s with no tails.\n* **Your model:** great on simple layers, biased low near faults, wobbly on artistic “Style” geology.\n* **Fixes:** first kill the bias, then squeeze the tails—especially in the Style and CurveFault families.\n\nEach change above chips away at MAE; combining two or three complementary tactics often wins Kaggle contests. Good luck refining your inversion model!","metadata":{"_uuid":"dcccc6c7-32e1-44bf-a4ae-8086ab3b0fa9","_cell_guid":"59a1735a-5acc-4d80-a92e-6ab1234b93e3","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"## Depth Wise Error Curve Analysis","metadata":{"_uuid":"60f993c0-31c2-4363-b379-84aaabc378a8","_cell_guid":"1d35272d-99f1-45e7-846f-e0b4f15a84b5","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"image_path = \"/kaggle/working/fwi_analysis_output/cross_family_depth_mae.png\"\ndisplay(Image(filename=image_path))","metadata":{"_uuid":"35223747-a681-452c-bb4c-17eedf38efe1","_cell_guid":"6fdf5590-b2f4-46be-9b84-504faac64274","trusted":true,"collapsed":false,"_kg_hide-input":true,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n\n| Plot element                                           | Plain-English meaning                                                                                                        | Concrete example                                                                                                                                                                                       |\n| ------------------------------------------------------ | ---------------------------------------------------------------------------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ |\n| **Mean MAE (m/s)** on the *x-axis*                     | The *average* absolute error at one fixed depth level across *every* grid-cell and *every* sample in that geological folder. | If the blue line sits at **15 m/s** when the *Depth Index* is **25**, it means that, on average, the model is off by 15 m/s for all cells that lie 25 grid-steps below the surface in that sub-family. |\n| **Depth Index** on the *y-axis* (top ≈ 0, bottom ≈ 70) | Think of it as vertical pixel rows in the 70 × 70 velocity map: 0 = surface, 70 = deepest layer.                             | A spike near index 60 hints at problems in the “basement” layer of that synthetic geology.                                                                                                             |\n| **One panel per folder**                               | Same ten folders you saw earlier (`CurveVel_A`, `Style_B`, …).                                                               | Lets you spot which depth ranges are error-prone *and* in which geological setting.                                                                                                                    |\n\n---\n\n## 2  Quick-Scan Cheat Sheet\n\n| Family                            | Overall pattern                                                           | Simple reading                                                                            | Likely reason                                                      |\n| --------------------------------- | ------------------------------------------------------------------------- | ----------------------------------------------------------------------------------------- | ------------------------------------------------------------------ |\n| **FlatVel\\_A / FlatVel\\_B**       | Flat-ish curve, low (2–12 m/s) except small bump at depth ≈ 65            | Model nails simple horizontal layers; struggles a bit on the deepest sand/shale contrast. | Few reflections arrive from great depth → weaker training signal.  |\n| **CurveVel\\_A / CurveVel\\_B**     | Errors rise almost linearly with depth (10 → 60 m/s)                      | Curved layers become harder the deeper they get.                                          | Travel paths bend; later arrivals are weaker and harder to invert. |\n| **FlatFault\\_A / FlatFault\\_B**   | Two plateaus: low near surface, sharp jump at mid-depth (index ≈ 50 – 60) | Fault throw located there; model blurs the discontinuity.                                 | Neural net smoothes velocities across the slip plane.              |\n| **CurveFault\\_A / CurveFault\\_B** | Low at top, staircase jumps at depths where faults + curvature co-exist   | Compounded complexity → high local MAE.                                                   | Interference of dipping layers and fault shadow.                   |\n| **Style\\_A / Style\\_B**           | High everywhere (40 – 100 m/s) but still trending upward                  | Artistic “blob” texture difficult regardless of depth; deeper even worse.                 | High-frequency “noise-like” texture + weaker signal.               |\n\n---\n\n## 3  Why Depth Matters for Overall MAE\n\nThe competition metric is\n\n$$\n\\text{MAE} \\;=\\; \\frac{1}{N}\\sum\\limits_{i=1}^{N}|v_{\\text{pred},i}-v_{\\text{true},i}|\n$$\n\n**In words:** add up the *absolute* velocity errors for every single cell and divide by how many cells you have.\n\nBecause each horizontal layer contains the same number of cells, **deep layers contribute just as much** to the leaderboard as shallow ones. That means the *sloping* right-hand side of many panels is a silent MAE killer: you might look great at the surface but bleed points down below.\n\n---\n\n## 4  Actionable Ideas Focused on Depth-Related Errors\n\n| Idea                                                                                                       | How it targets these curves                                                                             | Pros                                                          | Cons / Trade-offs                                                           |\n| ---------------------------------------------------------------------------------------------------------- | ------------------------------------------------------------------------------------------------------- | ------------------------------------------------------------- | --------------------------------------------------------------------------- |\n| **Depth-weighted loss**<br>(*e.g.* multiply errors by a weight that grows with depth index)                | Forces the optimiser to pay more attention to bottom layers where MAE is largest.                       | Easy to implement; directly flattens the right-hand slope.    | Might *increase* shallow-layer error; need weight schedule tuning.          |\n| **Curriculum training by depth**<br>(start training with top 30 layers, gradually un-freeze deeper layers) | Lets the net learn easy reflections first, then specialise on deeper arrivals.                          | Mimics how humans interpret seismic; can improve convergence. | Adds training stages; hyper-parameters decide when to “unlock” depth bands. |\n| **Positional encoding / depth embedding**                                                                  | Gives the model an explicit notion of “how deep am I?” so it can vary receptive fields.                 | Lightweight architectural change.                             | Only helps if the net *uses* the embedding; needs validation.               |\n| **Multi-scale time-windowing**                                                                             | Feed earlier time slices for shallow focus and later time slices for deep focus into separate branches. | Tailors the receptive field to depth-specific signal.         | Network complexity grows; must blend branches well.                         |\n| **Augment with extra low-frequency shots**                                                                 | Low frequencies penetrate deeper; simulated augmentation can strengthen deep information.               | Pure data trick; no code changes.                             | Must keep augmentation realistic; risk of domain shift.                     |\n| **Edge-aware penalty near fault depths**                                                                   | Add a term that penalises blurring across depth-indexed fault zones (e.g. from a fault mask).           | Specifically attacks the jump at index ≈ 50 in fault plots.   | Needs a fault mask (easy in synthetic data, harder elsewhere).              |\n| **Hybrid inversion (physics + CNN) only at deep layers**                                                   | Use a lightweight adjoint or ray-based update after the CNN to refine bottom half.                      | Physics good at depth where data are sparse.                  | Adds compute; model hand-off must be seamless.                              |\n\n---\n\n### Example: Simple Depth-Weighted MAE\n\n$$\n\\mathcal{L} \\;=\\; \\frac{1}{N}\\sum_{i=1}^{N}w(z_i)\\,|v_{\\text{pred},i}-v_{\\text{true},i}|\n$$\n\nWith\n$w(z) = 1 + \\alpha \\frac{z}{70}$\n\n*Plain English translation:*\n\n> “Keep the usual MAE but multiply each cell’s error by a weight that grows linearly from **1 ×** at the surface to **1 + α** at the bottom. If α = 1, the deepest layer counts double.”\n\nTry α ≈ 0.5 first; monitor whether deep-layer spikes shrink without blowing up surface error.\n\n---\n\n## 5  Putting It All Together—A Minimal Roadmap\n\n1. **Sanity check**\n\n   * Compute per-depth MAE on a validation split to ensure the figure you saw is not a fluke.\n2. **Quick prototype** (1 – 2 days)\n\n   * Implement depth-weighted loss with a modest α.\n   * Re-train for a few epochs; re-plot depth curves.\n3. **Iterative refinement** (1 week)\n\n   * Add depth positional encoding *and/or* curriculum unlocking.\n   * If using UNet, increase dilation or use larger kernels in deeper decoder blocks.\n4. **Advanced stage** (competition duration)\n\n   * Combine data augmentation (low-frequency shots) with a physics-guided deep refinement module.\n   * Ensembling: average a “shallow-tuned” model with a “deep-tuned” one.\n\n---\n\n## 6  Key Take-aways\n\n* **Steeper curves to the right tell you exactly where your leaderboard points are leaking.**\n* **Depth-awareness—either in the loss, the architecture, or the data—tends to pay off.**\n* **Always re-plot depth-wise MAE after each experiment**; it is the fastest visual feedback loop for this problem.\n\nWith targeted depth fixes layered on top of the family-specific ideas we discussed earlier, you can attack both vertical (*depth*) and horizontal (*geology-style*) dimensions of the error landscape—and that’s usually the recipe for a podium finish. Good luck diving deeper!","metadata":{"_uuid":"ee828945-f33c-4add-b66a-0b449586ab95","_cell_guid":"159fa99f-f2a9-4e13-a8d8-36a68ca960b9","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"## Mean Log Amplitude Residual Spectrum Analysis","metadata":{"_uuid":"c80c07c3-bd08-46c4-adfd-c3fb3fa88795","_cell_guid":"0dd7a578-da6c-491e-b7e3-9581e3dedb01","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"image_path = \"/kaggle/working/fwi_analysis_output/cross_family_residual_spectrums.png\"\ndisplay(Image(filename=image_path))","metadata":{"_uuid":"aa16fd9c-550e-41e3-b212-ed9c6d86c352","_cell_guid":"6a4b537a-a4c5-484b-89a3-ce2d761d0deb","trusted":true,"collapsed":false,"_kg_hide-input":true,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n\n| Plot element                           | Plain-English meaning                                                                                                                                                             | Concrete example                                                                                                                                                                                     |\n| -------------------------------------- | --------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |\n| **Heat-map pixel**                     | Average strength (log-amplitude) of the *residual’s* 2-D Fourier transform at a particular spatial frequency.  Bright ≈ large error energy, dark ≈ little error.                  | A bright spot at the centre means the model is systematically off at the **overall background** velocity (a very low-frequency error).                                                               |\n| **Axes kx and ky (after “fft-shift”)** | Horizontal (kx) and vertical (ky) spatial frequencies measured in “cycles per model width/height”.  The centre (0, 0) is **low frequency**; edges are **high frequency** wiggles. | A streak straight up the vertical axis (kx ≈ 0) means errors that vary rapidly **with depth** but are nearly constant **horizontally**—exactly the kind you get from mis-estimating layer thickness. |\n| **One panel per folder**               | Same ten geological sub-families you’ve seen before.                                                                                                                              | Lets you see which kinds of structures (low vs high frequency, horizontal vs vertical) stump the model in each scenario.                                                                             |\n\n> **Concrete mental picture**\n> Think of the residual map as a black-and-white photo.\n> Running a 2-D Fourier transform turns that photo into a “recipe” of sine-wave patterns.\n> The plot shows *which sine-waves still need fixing* after your model’s best shot.\n\n---\n\n### 2 Crash-Course on Reading 2-D Spectra (Beginner Edition)\n\n1. **Brightness at the centre** → model has trouble with the *overall* trend or very large-scale shapes (low frequencies).\n2. **Rings or fuzzy blobs far from centre** → it misses fine details (high frequencies).\n3. **Vertical or horizontal bars** → errors are *anisotropic*—strong only in one direction (e.g. layers).\n4. **Diagonal streaks** → errors follow dipping structures (fault planes or curved beds).\n\n*(Imagine shining a flashlight from the centre outwards; whatever areas light up first mark the frequencies costing you MAE.)*\n\n---\n\n### 3 What the Ten Panels Say—Folder by Folder\n\n| Family                            | Dominant pattern                                           | Interpretation                                                                                          | How it links to geology                                                                     |\n| --------------------------------- | ---------------------------------------------------------- | ------------------------------------------------------------------------------------------------------- | ------------------------------------------------------------------------------------------- |\n| **FlatVel\\_A / FlatVel\\_B**       | Bright vertical stripe (kx ≈ 0) + weaker horizontal stripe | Misses mostly *vertical* variations (depth layering) and a bit of lateral smoothing.                    | Layers are horizontal; network slightly blurs layer boundaries and baseline velocity.       |\n| **CurveVel\\_A / CurveVel\\_B**     | Elongated diamond / X-shape                                | Errors concentrate along directions of folded layers—your model under-captures curvature.               | Needs better handling of dipping reflectors and bending wavefronts.                         |\n| **FlatFault\\_A / FlatFault\\_B**   | Cross (+) with extra blobs near diagonal                   | Vertical stripe = layer bias; diagonal wings = fault throw errors.                                      | Blurring across the fault plane leaves directional, mid-frequency artefacts.                |\n| **CurveFault\\_A / CurveFault\\_B** | Wide, faint X with dark wedge along ky ≈ 0                 | Both curvature and faulting hurt; model okay for pure horizontal variation (ky ≈ 0) but shaky for dips. | Compounded complexity; needs multi-directional receptive fields.                            |\n| **Style\\_A / Style\\_B**           | Round, bright central blob fading isotropically            | Errors spread *uniformly* in all directions—model lacks both coarse and fine texture fidelity.          | Style images inject random, isotropic patterns like “blobs”; net never learned those bases. |\n\n---\n\n### 4 Why Frequency-Domain Errors Feed the Leaderboard MAE\n\n$$\nE(f_x,f_y) = \\bigl|\\mathcal{F}\\{v_\\text{pred}-v_\\text{true}\\}(f_x,f_y)\\bigr|\n$$\n\n*Translation:*\n“Take the residual map, run a Fourier transform **F**, and look at the absolute size of each sine-wave component with horizontal frequency $f_x$ and vertical frequency $f_y$.”\n\nThe **MAE** you care about in pixel space equals the “energy” in *all* these frequency bins added up (Parseval’s theorem).\nSo a bright streak—even if narrow—acts like a leaking pipe: every pixel row or column along that frequency keeps adding to your MAE score.\n\n---\n\n### 5 Practical Tweaks—Now in Frequency Language\n\n| Tactic                                                                     | Frequency-space angle                                                                                        | Pros                                                 | Cons / Trade-offs                                                  |\n| -------------------------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------ | ---------------------------------------------------- | ------------------------------------------------------------------ |\n| **Add a *spectral loss*** (e.g. L1 or L2 on FFT of prediction vs truth)    | Makes optimiser explicitly minimise those bright bins.                                                       | Simple extra term; directly shrinks bad frequencies. | Need to balance with spatial MAE; may slow training.               |\n| **High-frequency emphasis (Laplacian or gradient loss)**                   | Penalises edges being too smooth → kills vertical stripes in FlatVel & Fault families.                       | Cheap; already common in imaging.                    | Doesn’t cure low-freq bias.                                        |\n| **Low-frequency bias correction layer**                                    | Learns a global offset or very low-order polynomial across the map.                                          | Darkens the bright centre blob.                      | Ignores fine detail; only fixes the dot at (0,0).                  |\n| **Directional convolutions / group CNNs**                                  | Provide orientation-aware kernels to grab diagonal dips visible in Curve/Fault spectra.                      | Targets X-shaped artefacts.                          | More parameters; harder to train.                                  |\n| **Fourier Feature Embeddings**                                             | Encode spatial coords with sin/cos at multiple frequencies; helps Style families capture isotropic textures. | Proven in NeRF-style models.                         | Minor overhead; still needs enough capacity to use those features. |\n| **Multi-resolution supervision** (predict and compare at ¼, ½, full scale) | Ensures both low and high frequencies receive gradient signal.                                               | Balances coarse and fine; stable.                    | More bookkeeping in loss code.                                     |\n| **Frequency-aware data augmentation**                                      | Inject band-limited noise or mixup in Fourier space to widen training bandwidth.                             | Network experiences missing bands before test time.  | Must avoid unrealistic artefacts; monitor validation carefully.    |\n\n---\n\n### 6 Example: Simple Spectral Loss Add-On\n\n$$\n\\mathcal{L}_\\text{total} \\;=\\; \\mathcal{L}_\\text{MAE} \\;+\\; \\lambda \\frac{1}{N_{f}}\\!\\sum_{f_x,f_y} \\bigl|\\,\\log E_\\text{pred}(f_x,f_y)-\\log E_\\text{true}(f_x,f_y)\\bigr|\n$$\n\n**Plain English translation:**\n\n> “Keep the usual MAE, **plus** a penalty that says ‘make the log-amplitude of every sine-wave in your prediction look like the true one’. The *λ* knob controls how strongly you care about frequency mistakes.”\n\nTypical starting point: λ ≈ 0.1. Monitor the spectrum: those bright crosses should dim after a few epochs.\n\n---\n\n### 7 Suggested Roadmap Focused on Spectral Fixes\n\n1. **Re-plot after every tweak** to see which bins darken—fast visual feedback.\n2. **Stage 1 (days):** Add a low-frequency bias layer + spectral loss with small λ.\n3. **Stage 2 (week):** Switch final convolution block to *directional* kernels or add Fourier features; keep spectral loss.\n4. **Stage 3 (competition life-cycle):** Ensemble a *texture-specialist* subnet (trained with high λ) with your main model; average predictions in pixel space.\n\n---\n\n### 8 Key Take-aways for Beginners\n\n* **Bright centres** → global bias; **vertical stripes** → blurred layers; **diagonal wings** → missed dips / faults; **round blobs** → universal texture miss.\n* A small **spectral loss** is often the cheapest, most direct way to dim those artefacts and cut MAE.\n* Never trust a single metric—always look at **histograms (error sign), depth curves (vertical location), and spectra (frequency content)** together.  Fixing all three views usually lands you on the Kaggle medal board.\n\nNow you have a frequency-domain compass to guide the next round of model surgery—happy tuning!","metadata":{"_uuid":"607dbb8a-7c0e-445c-a3fc-5dc1e8eca309","_cell_guid":"c3430f95-2f17-498d-b7b9-97e4dd4aa52f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"## Pixel-wise MAE Heat-maps","metadata":{"_uuid":"856f8a06-017c-4b64-8705-d5ff7af897e9","_cell_guid":"93f694bc-97bd-404a-b846-ac504c7c4d3e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"image_path = \"/kaggle/working/fwi_analysis_output/cross_family_mae_heatmaps.png\"\ndisplay(Image(filename=image_path))","metadata":{"_uuid":"db7a9caf-7d7d-403d-90bd-13345bcb063c","_cell_guid":"aba32d52-4f53-497b-8d22-6ef302231514","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n\nEach little picture is a **70 × 70 grid** (depth ↓, horizontal →) whose colour encodes the **mean absolute error** (MAE, in m/s) of the model at that exact cell, averaged over every sample that lives in the folder named above the panel.\n\n| Axis                          | Plain-English meaning                                                        | Concrete example                                                                                   |\n| ----------------------------- | ---------------------------------------------------------------------------- | -------------------------------------------------------------------------------------------------- |\n| **Horizontal Index (x-axis)** | Left-to-right position in the seismic section. 0 = far left, 69 = far right. | Column 34 is the middle of the model.                                                              |\n| **Depth Index (y-axis)**      | Vertical position; 0 at the surface, 69 at the deepest layer.                | Row 60 is deep basement.                                                                           |\n| **Colour**                    | Dark red ≈ small error, yellow/white ≈ large error.                          | A bright horizontal stripe near depth 65 means the model is badly wrong for that whole deep layer. |\n\n> **Concrete example**\n> In **FlatVel\\_A**, the bottom five rows glow bright yellow, so the model is off by, say, **120 m/s** there, while the top rows are dark—perhaps only **5 m/s** wrong.\n\n---\n\n## 2 Quick How-To Read a Heat-map\n\n1. **Horizontal stripes** → mistakes aligned with layers (model blurs layer boundaries).\n2. **Vertical bands** → errors tied to specific x-positions (often faults or edges).\n3. **Patches or blobs** → localised trouble spots (curved folds, salt bodies, “Style” blobs).\n4. **Trend with depth** → if the lower half is brighter, deeper information is missing.\n\n---\n\n## 3 Folder-by-Folder Observations\n\n| Family                            | Visual cues                                                                                                           | What it says                                                                      | Likely cause                                                                 |\n| --------------------------------- | --------------------------------------------------------------------------------------------------------------------- | --------------------------------------------------------------------------------- | ---------------------------------------------------------------------------- |\n| **FlatVel\\_A / FlatVel\\_B**       | Clean horizontal stripes that brighten with depth; very uniform left → right.                                         | Model great at shallow flat layers, steadily worse in the low-frequency basement. | Weak seismic energy at late times; network not depth-aware.                  |\n| **CurveVel\\_A / CurveVel\\_B**     | Fan-shaped bright zone widening with depth; mottled texture.                                                          | Errors largest beneath the *steepest* curved portions.                            | Curvature bends wave-paths; single-scale kernels struggle to align features. |\n| **FlatFault\\_A / FlatFault\\_B**   | Bright horizontal band right where the fault plane meets bottom layers; slight vertical band at the fault x-position. | Mis-estimates velocity jump across the fault and “smears” it vertically.          | Loss doesn’t penalise sharp discontinuities strongly enough.                 |\n| **CurveFault\\_A / CurveFault\\_B** | Combination of fan (curvature) and horizontal bright band (fault).                                                    | Two mechanisms interact, hence larger patchy errors.                              | Need orientation-aware filters and fault-sharp penalties simultaneously.     |\n| **Style\\_A / Style\\_B**           | Whole middle depth band bright and noisy; top shallow zone darker.                                                    | Model okay near surface but loses the random “Style” texture deeper down.         | Training data scarce for those isotropic blobs; network over-smooths.        |\n\n*(Exact µs vary, but the qualitative patterns are robust.)*\n\n---\n\n## 4 Why These Patterns Hurt the Leaderboard MAE\n\n$$\n\\text{MAE}=\\frac{1}{70\\times70}\\sum_{z=0}^{69}\\sum_{x=0}^{69}\\bigl|\\text{error}(z,x)\\bigr|\n$$\n\n> **Plain-English translation**\n> “Add up the absolute error of **every pixel**; divide by 4 900.”\n> A bright swath that spans even 10 % of the grid can dominate the total, so these heat-maps tell you *where* to spend modelling effort.\n\n---\n\n## 5 Practical Tweaks that Attack Spatial Bias Directly\n\n| Idea                                                                                     | Attacks which panels?                                | Pros                                            | Cons / Trade-offs                                  |\n| ---------------------------------------------------------------------------------------- | ---------------------------------------------------- | ----------------------------------------------- | -------------------------------------------------- |\n| **Depth-weighted or layer-aware loss**                                                   | FlatVel, FlatFault (deep stripes)                    | Targets uniform bottom-layer errors.            | Might raise shallow MAE if weights too aggressive. |\n| **Coordinate Convolutions or Positional Embeddings**<br>(give x and z as extra channels) | Curve & CurveFault (fan-shaped)                      | Lets kernels adapt to “where” they are.         | Small extra parameter cost.                        |\n| **Fault-mask edge loss**<br>(penalise smoothing across a provided fault label)           | FlatFault & CurveFault (horizontal + vertical bands) | Sharpens velocity jump.                         | Needs synthetic fault mask (easy here).            |\n| **Direction-selective filters / deformable convs**                                       | Curve families (dips), Style (random)                | Captures orientation-specific textures.         | Heavier model, tuning complexity.                  |\n| **Multi-scale supervision**<br>(predict at ¼, ½, full resolution)                        | All panels (layer & blob texture)                    | Gives gradients at both coarse and fine levels. | More bookkeeping in code.                          |\n| **Style-focused data augmentation**<br>(procedurally mix blob patterns)                  | Style\\_A / B                                         | Cheap and often effective.                      | Must preserve physics realism.                     |\n\n---\n\n### Worked Example — Coordinate Convs\n\nAdd two constant channels:\n\n```text\nX(i,j) = i / 69          # horizontal coordinate, 0 → 1\nZ(i,j) = j / 69          # depth coordinate, 0 → 1\n```\n\nConcatenate them to the feature tensor before the first conv layer.\n**Intuition:** a kernel can now learn “if Z > 0.8 tighten weights because deep layers are tricky”.\nThis usually flattens the bottom bright band in **FlatVel\\_B** within a few epochs.\n\n---\n\n## 6 Minimal Roadmap\n\n1. **Baseline re-plot:** Verify these heat-maps on your validation split.\n2. **1–2 day quick win:** Add X/Z coordinate channels **and** a modest depth-weighted MAE (e.g. weight = 1 + 0.5·depth/69).\n3. **1 week iteration:**\n\n   * Add a fault-edge penalty when folder contains *Fault*.\n   * Train a small *Style-specialist* head with extra high-frequency loss.\n4. **Final stretch:** Ensemble the “spatial-bias-fixed” net with your earlier “frequency-fixed” net from the previous step.\n\n---\n\n## 7 Key Take-aways for Beginners\n\n* **Heat-maps spotlight *where* the model is wrong.**\n* **Bottom layers** and **fault zones** leak the most MAE.\n* Give your network *knowledge of position* (coordinates, depth weights) and *special treatment of discontinuities* (edge losses) to plug those leaks.\n* Re-generate this graphic after every experiment; when the yellow fades, your leaderboard score will follow.","metadata":{"_uuid":"6797484a-fd81-4bfa-83d1-27f55c80bde9","_cell_guid":"16940b2c-b6f9-4b83-a0ac-defdc024ed34","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"## Inference","metadata":{"_uuid":"bd70aaea-489e-4cd4-aa5d-9d8711295a5a","_cell_guid":"b7139b1f-d2b4-482b-8cf3-1adcc683fc22","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"DO_INFERENCE = True\n\nif DO_INFERENCE:\n\n    print(f\"Starting inference on {len(test_files)} test files...\", flush=True)\n    t0 = time.time()\n    \n    x_cols = [f\"x_{i}\" for i in range(1, 70, 2)]  # as in the example\n    fieldnames = [\"oid_ypos\"] + x_cols\n\n    # Use torch.inference_mode for slight performance & memory improvements\n    with torch.inference_mode():\n        with open(\"submission.csv\", \"w\", newline=\"\") as csvfile:\n            writer = csv.DictWriter(csvfile, fieldnames=fieldnames)\n            writer.writeheader()\n    \n            # Provide a total for tqdm so it starts immediately\n            for inputs_batch, oids_test in tqdm(\n                dl_test,\n                desc=\"Batches\",\n                total=len(dl_test),  # show total count\n                miniters=1           # forces refresh after each batch\n            ):\n                inputs_batch = inputs_batch.to(device, non_blocking=True)\n                \n                # Forward pass\n                outputs = model(inputs_batch)  # shape: [B, 1, 70, 70]\n                y_preds = outputs[:, 0].cpu().numpy()  # shape: (B, 70, 70)\n    \n                # Write each y-row to a line in the CSV\n                for y_pred, oid_test in zip(y_preds, oids_test):\n                    for y_pos in range(70):\n                        row_vals = [y_pred[y_pos, x_pos] for x_pos in range(1, 70, 2)]\n                        row_dict = dict(zip(x_cols, row_vals))\n                        row_dict[\"oid_ypos\"] = f\"{oid_test}_y_{y_pos}\"\n                        writer.writerow(row_dict)\n    \n    elapsed = format_time(time.time() - t0)\n    print(f\"Inference complete. Time: {elapsed}\")\n    print(\"Wrote results to 'submission.csv'.\")","metadata":{"_uuid":"96fed46f-986f-4173-b066-9f39e15a5257","_cell_guid":"a49f0ddf-d2d3-4c6b-8bfd-d82637237323","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}