{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Example Use of Deepwave Library\n\nThis is a PINN competition after all.\nYet the top public notebooks barely have any physics in them.\nIn order to instill physics in your NN, you'll need to implement the forward wave propagation equations yourself in order to backprop through the physical forward pass.\nFortunately, this is not necessary, as there exists a few open source FWI libraries.\nOne of them is Deepwave, which I will demonstrate its usage with this notebook.\nThe following is based on the first half of the official tutorial [Full-Waveform Inversion (FWI)](https://ausargeo.com/deepwave/example_fwi).\nAlongside that, you get a bunch of helper functions as well.\n\nIf you wish to skip all the clutter, the main training loop is entered in cell 16.","metadata":{}},{"cell_type":"code","source":"!pip install deepwave pytorch_msssim -q","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:47.661638Z","iopub.execute_input":"2025-06-04T21:54:47.661898Z","iopub.status.idle":"2025-06-04T21:54:50.779203Z","shell.execute_reply.started":"2025-06-04T21:54:47.661876Z","shell.execute_reply":"2025-06-04T21:54:50.778227Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport random\nfrom dataclasses import dataclass, field\nfrom enum import Enum\nfrom functools import partial\nfrom pathlib import Path\nfrom typing import Callable, Generator\n\nimport deepwave\nimport matplotlib.lines as mlines\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom deepwave import scalar\nfrom deepwave.wavelets import ricker\nfrom IPython.display import HTML\nfrom matplotlib import gridspec\nimport matplotlib.animation as animation\nfrom matplotlib.legend_handler import HandlerTuple\nfrom pytorch_msssim import ssim\nfrom scipy.ndimage import gaussian_filter\nfrom scipy.fft import fft, fftfreq, fftshift\nfrom torch.optim import Optimizer\nfrom torch.optim.lr_scheduler import _LRScheduler\nfrom tqdm.notebook import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:50.780241Z","iopub.execute_input":"2025-06-04T21:54:50.780535Z","iopub.status.idle":"2025-06-04T21:54:53.057429Z","shell.execute_reply.started":"2025-06-04T21:54:50.780504Z","shell.execute_reply":"2025-06-04T21:54:53.056640Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"path_data = Path(\"/kaggle/input/waveform-inversion\")\npath_train = path_data / \"train_samples\"\npath_test = path_data / \"test\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.058246Z","iopub.execute_input":"2025-06-04T21:54:53.058630Z","iopub.status.idle":"2025-06-04T21:54:53.062335Z","shell.execute_reply.started":"2025-06-04T21:54:53.058605Z","shell.execute_reply":"2025-06-04T21:54:53.061561Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Classes and functions","metadata":{}},{"cell_type":"code","source":"class FamilyName(str, Enum):\n    CurveFault_A = \"CurveFault_A\"\n    CurveFault_B = \"CurveFault_B\"\n    CurveVel_A   = \"CurveVel_A\"\n    CurveVel_B   = \"CurveVel_B\"\n    FlatFault_A  = \"FlatFault_A\"\n    FlatFault_B  = \"FlatFault_B\"\n    FlatVel_A    = \"FlatVel_A\"\n    FlatVel_B    = \"FlatVel_B\"\n    Style_A      = \"Style_A\"\n    Style_B      = \"Style_B\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.064317Z","iopub.execute_input":"2025-06-04T21:54:53.064517Z","iopub.status.idle":"2025-06-04T21:54:53.075912Z","shell.execute_reply.started":"2025-06-04T21:54:53.064502Z","shell.execute_reply":"2025-06-04T21:54:53.075227Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def iter_sample_pairs(\n    base_path: Path, family: FamilyName, random_sample: bool = False\n) -> Generator[tuple[Path, Path], None, None]:\n    \"\"\"\n    Yield (seismic, velocity) .npy file pairs from a given family.\n\n    Parameters\n    ----------\n    base_path : Path\n        Root path containing all family directories.\n    family : FamilyName\n        The dataset family to sample from.\n    random_sample : bool, optional\n        Whether to yield pairs in random order (default is sequential).\n\n    Yields\n    ------\n    seismic_path : Path\n        Path to seismic data .npy file.\n    model_path : Path\n        Path to velocity model .npy file.\n    \"\"\"\n    family_data_path = base_path / family.value\n\n    # Fault family of data has different directory structures, which requires different handling.\n    if \"fault\" in family_name.name.lower():\n        yield from _iter_fault_pairs(family_data_path, random_sample)\n    else:\n        yield from _iter_nonfault_pairs(family_data_path, random_sample)\n\n\ndef _iter_fault_pairs(family_data_path: Path, random_sample: bool) -> Generator[tuple[Path, Path], None, None]:\n    seis_files = sorted(f for f in family_data_path.glob(\"seis*.npy\"))\n    vel_files = sorted(f for f in family_data_path.glob(\"vel*.npy\"))\n\n    if not seis_files or not vel_files:\n        warnings.warn(f\"No seis/vel .npy files found in {family_data_path}!\")\n        return\n\n    pairs = []\n    for s_file in seis_files:\n        expected_v = s_file.name.replace(\"seis\", \"vel\", 1)\n        v_file = family_data_path / expected_v\n        if v_file.exists():\n            pairs.append((s_file, v_file))\n\n    if not pairs:\n        warnings.warn(f\"No directly matched seis/vel pairs in {family_data_path}!\")\n        return\n\n    if random_sample:\n        while True:\n            yield random.choice(pairs)\n    else:\n        for pair in pairs:\n            yield pair\n\n\ndef _iter_nonfault_pairs(family_data_path: Path, random_sample: bool) -> Generator[tuple[Path, Path], None, None]:\n    seis_path = family_data_path / \"data\"\n    vel_path = family_data_path / \"model\"\n\n    seis_files = sorted(f for f in seis_path.glob(\"*.npy\"))\n    vel_files = sorted(f for f in vel_path.glob(\"*.npy\"))\n\n    if not seis_files or not vel_files:\n        warnings.warn(f\"No data/model files found in {family_data_path}!\")\n        return\n\n    pairs = []\n    for s_file in seis_files:\n        expected_v = s_file.name.replace(\"data\", \"model\", 1)\n        v_file = vel_path / expected_v\n        if v_file.exists():\n            pairs.append((s_file, v_file))\n\n    if not pairs:\n        warnings.warn(f\"No matched data/model pairs in {family_data_path}!\")\n        return\n\n    if random_sample:\n        while True:\n            yield random.choice(pairs)\n    else:\n        for pair in pairs:\n            yield pair\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.076679Z","iopub.execute_input":"2025-06-04T21:54:53.076859Z","iopub.status.idle":"2025-06-04T21:54:53.091341Z","shell.execute_reply.started":"2025-06-04T21:54:53.076844Z","shell.execute_reply":"2025-06-04T21:54:53.090650Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def total_variation_loss(v: torch.Tensor) -> torch.Tensor:\n    \"\"\"\n    Compute isotropic total variation (TV) loss to encourage piecewise smoothness and discourage noise/artifacts.\n\n    $$ \\mathcal{L}_{\\text{TV}} = \\sum_{i,j} \\sqrt{(\\partial_x v_{i,j})^2 + (\\partial_y v_{i,j})^2} $$\n\n    Parameters\n    ----------\n    v : torch.Tensor\n        Velocity field tensor of shape (B, C, H, W). C is often == 1.\n\n    Returns\n    -------\n    torch.Tensor\n        Scalar TV loss.\n    \"\"\"\n    dx = v[:, :, 1:, :] - v[:, :, :-1, :]\n    dy = v[:, :, :, 1:] - v[:, :, :, :-1]\n    return torch.mean(torch.sqrt(dx**2 + 1e-8)) + torch.mean(torch.sqrt(dy**2 + 1e-8))\n\n\ndef laplacian_loss(v: torch.Tensor) -> torch.Tensor:\n    \"\"\"\n    Compute Laplacian loss to penalize curvature and promote smooth gradients.\n\n    This discourages sharp, unrealistic jumps while preserving structure better than TV.\n\n    $$ \\mathcal{L}_{\\text{Laplacian}} = \\|\\Delta v\\|_1 \\quad \\text{or} \\quad \\|\\nabla^2 v\\|_2^2 $$\n\n    Parameters\n    ----------\n    v : torch.Tensor\n        Velocity field tensor of shape (B, C, H, W). C is often == 1.\n\n    Returns\n    -------\n    torch.Tensor\n        Scalar Laplacian loss.\n    \"\"\"\n    kernel = (\n        torch.tensor(\n            [\n                [0, 1, 0],\n                [1, -4, 1],\n                [0, 1, 0],\n            ],\n            dtype=v.dtype,\n            device=v.device,\n        )\n        .unsqueeze(0)\n        .unsqueeze(0)\n    )\n    laplace = F.conv2d(v, kernel, padding=1)\n    return torch.mean(torch.abs(laplace))\n\n\ndef monotonic_depth_loss(v: torch.Tensor) -> torch.Tensor:\n    \"\"\"\n    Penalize decreases in velocity with depth to enforce monotonicity along vertical axis.\n\n    Velocity typically increases with depth. Encourage ∂v/∂y ≥ 0.\n\n    Parameters\n    ----------\n    v : torch.Tensor\n        Velocity field tensor of shape (B, C, H, W). C is often == 1.\n\n    Returns\n    -------\n    torch.Tensor\n        Scalar monotonicity loss.\n    \"\"\"\n    dy = v[:, :, 1:, :] - v[:, :, :-1, :]\n    return torch.mean(F.relu(-dy))\n\n\ndef ssim_loss(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n    \"\"\"\n    Structural Similarity Index (SSIM) loss.\n\n    If ground truth velocity maps have consistent structure, SSIM loss can help preserve patterns.\n\n    Parameters\n    ----------\n    pred : torch.Tensor\n        Predicted tensor of shape (B, C, H, W).\n    target : torch.Tensor\n        Ground truth tensor of same shape as pred.\n\n    Returns\n    -------\n    torch.Tensor\n        Scalar SSIM loss.\n    \"\"\"\n    return 1 - ssim(pred, target, data_range=target.max() - target.min(), size_average=True)\n\n\ndef fourier_domain_loss(pred: torch.Tensor, target: torch.Tensor, p: int = 1) -> torch.Tensor:\n    \"\"\"\n    Compute L_p loss between Fourier magnitudes of prediction and target.\n\n    Parameters\n    ----------\n    pred : torch.Tensor\n        Predicted tensor of shape (B, 1, H, W).\n    target : torch.Tensor\n        Ground truth tensor of shape (B, 1, H, W).\n    p : int, optional\n        Norm degree (1 for L1, 2 for L2), by default 1.\n\n    Returns\n    -------\n    torch.Tensor\n        Scalar loss in frequency domain.\n    \"\"\"\n    # Apply FFT2 to spatial dimensions\n    pred_fft = torch.fft.fft2(pred.to(dtype=torch.float32).squeeze(1), norm=\"ortho\")\n    target_fft = torch.fft.fft2(target.to(dtype=torch.float32).squeeze(1), norm=\"ortho\")\n\n    # Take magnitude\n    pred_mag = torch.abs(pred_fft)\n    target_mag = torch.abs(target_fft)\n\n    loss = F.l1_loss(pred_mag, target_mag) if p == 1 else F.mse_loss(pred_mag, target_mag)\n    return loss\n\n\ndef edge_loss(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n    \"\"\"\n    Compute L1 loss between Sobel edge magnitudes of prediction and target.\n\n    Parameters\n    ----------\n    pred : torch.Tensor\n        Predicted tensor of shape (B, C, H, W).\n    target : torch.Tensor\n        Ground truth tensor of shape (B, C, H, W).\n\n    Returns\n    -------\n    torch.Tensor\n        Scalar edge loss.\n    \"\"\"\n    B, C, H, W = pred.shape\n    N = B * C\n    pred = pred.view(N, 1, H, W)\n    target = target.view(N, 1, H, W)\n\n    sobel_x = (\n        torch.tensor(\n            [\n                [1, 0, -1],\n                [2, 0, -2],\n                [1, 0, -1],\n            ],\n            dtype=pred.dtype,\n            device=pred.device,\n        ).view(1, 1, 3, 3)\n        / 8.0\n    )\n    sobel_y = sobel_x.transpose(2, 3)\n\n    pred_gx = F.conv2d(pred, sobel_x, padding=1)\n    pred_gy = F.conv2d(pred, sobel_y, padding=1)\n\n    sobel_x = sobel_x.to(target.device, dtype=target.dtype)\n    sobel_y = sobel_x.transpose(2, 3)\n\n    target_gx = F.conv2d(target, sobel_x, padding=1)\n    target_gy = F.conv2d(target, sobel_y, padding=1)\n\n    pred_grad = torch.sqrt(pred_gx**2 + pred_gy**2 + 1e-8).view(B, C, H, W)\n    target_grad = torch.sqrt(target_gx**2 + target_gy**2 + 1e-8).view(B, C, H, W)\n\n    return F.l1_loss(pred_grad, target_grad)\n\n\nclass LossAggregator(nn.Module):\n    \"\"\"\n    A weighted loss aggregator for combining multiple scalar loss components.\n\n    Methods\n    -------\n    register(name, fn, weight):\n        Add a new named loss function with specified weight.\n\n    forward(pred, target, return_components):\n        Compute total loss and optionally return per-component losses.\n    \"\"\"\n\n    def __init__(self):\n        super().__init__()\n        self.loss_fns: Dict[str, Callable] = {}\n        self.weights: Dict[str, float] = {}\n\n    def register(self, name: str, fn: Callable, weight: float = 1.0) -> None:\n        \"\"\"\n        Register a named loss function and its weight.\n\n        Parameters\n        ----------\n        name : str\n            Unique name of the loss.\n        fn : Callable\n            Loss function taking (pred, target) or (pred) as arguments.\n        weight : float, optional\n            Weight for this loss component, by default 1.0.\n        \"\"\"\n        self.loss_fns[name] = fn\n        self.weights[name] = weight\n\n    def forward(\n        self,\n        pred: torch.Tensor,\n        target: torch.Tensor,\n        return_components: bool = False,\n        components_with_weight: bool = False,\n    ) -> tuple[torch.Tensor, dict[str, float]] | torch.Tensor:\n        \"\"\"\n        Compute total weighted loss from all registered loss functions.\n\n        Parameters\n        ----------\n        pred : torch.Tensor\n            Predicted output tensor.\n        target : torch.Tensor\n            Ground truth tensor.\n        return_components : bool, optional\n            If True, return dictionary of individual component losses.\n\n        Returns\n        -------\n        torch.Tensor or (torch.Tensor, Dict[str, float])\n            Total scalar loss or (loss, component_dict).\n        \"\"\"\n        loss_total = torch.tensor(0.0, device=pred.device)\n        components = {}\n\n        for name, fn in self.loss_fns.items():\n            if name in {\"monotonic\", \"tv\", \"laplacian\"}:\n                # Apply to pred only\n                value = fn(pred)\n            else:\n                value = fn(pred, target)\n            weighted_value = self.weights[name] * value\n            loss_total += weighted_value\n            if components_with_weight:\n                components[name] = weighted_value.item()\n            else:\n                components[name] = value.item()\n\n        return (loss_total, components) if return_components else loss_total\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.092057Z","iopub.execute_input":"2025-06-04T21:54:53.092244Z","iopub.status.idle":"2025-06-04T21:54:53.109876Z","shell.execute_reply.started":"2025-06-04T21:54:53.092226Z","shell.execute_reply":"2025-06-04T21:54:53.109166Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class WarmupPlateauCooldownLR(_LRScheduler):\n    \"\"\"\n    Scheduler with linear warm-up, plateau, and linear cool-down learning rates.\n\n    Parameters\n    ----------\n    optimizer : Optimizer\n        Optimizer whose LR will be scheduled.\n    lr_init : float\n        Initial learning rate before warm-up.\n    lr_peak : float\n        Learning rate after warm-up.\n    lr_final : float\n        Learning rate after cool-down.\n    steps_warmup : int\n        Number of steps for linear warm-up.\n    steps_plateau : int\n        Number of steps to hold at lr_peak after warm-up.\n    steps_cooldown : int\n        Number of steps to linearly decay to lr_final.\n    last_epoch : int, optional\n        Index of last epoch. Default: -1.\n    \"\"\"\n    def __init__(\n        self,\n        optimizer: Optimizer,\n        lr_init: float,\n        lr_peak: float,\n        lr_final: float,\n        steps_warmup: int,\n        steps_plateau: int,\n        steps_cooldown: int,\n        last_epoch: int = -1\n    ):\n        self.lr_init = lr_init\n        self.lr_peak = lr_peak\n        self.lr_final = lr_final\n        self.steps_warmup = steps_warmup\n        self.steps_plateau = steps_plateau\n        self.steps_cooldown = steps_cooldown\n        self.total_steps = steps_warmup + steps_plateau + steps_cooldown\n        super().__init__(optimizer, last_epoch)\n\n    def get_lr(self):\n        step = self.last_epoch + 1\n\n        if step < self.steps_warmup:\n            # Linear warm-up: lr_init -> lr_peak\n            scale = step / self.steps_warmup\n            lr = self.lr_init + scale * (self.lr_peak - self.lr_init)\n        elif step < self.steps_warmup + self.steps_plateau:\n            # Hold at lr_peak\n            lr = self.lr_peak\n        elif step < self.total_steps:\n            # Linear cooldown: lr_peak -> lr_final\n            t = step - self.steps_warmup - self.steps_plateau\n            scale = t / self.steps_cooldown\n            lr = self.lr_peak + scale * (self.lr_final - self.lr_peak)\n        else:\n            # After schedule ends, hold lr_final\n            lr = self.lr_final\n\n        return [lr for _ in self.optimizer.param_groups]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.110759Z","iopub.execute_input":"2025-06-04T21:54:53.111021Z","iopub.status.idle":"2025-06-04T21:54:53.126450Z","shell.execute_reply.started":"2025-06-04T21:54:53.110998Z","shell.execute_reply":"2025-06-04T21:54:53.125735Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_ricker_wavelet(freq: float, nt: float, dt: float, peak_time: float, n_sources: int) -> torch.Tensor:\n    \"\"\"\n    Generate a Ricker wavelet tensor for all sources.\n\n    Parameters\n    ----------\n    freq : float\n        Ricker wavelet central frequency (Hz).\n    nt : float\n        Number of time steps.\n    dt : float\n        Temporal step size (seconds).\n    peak_time : float\n        Time in seconds with offset where the wavelet is at its peak.\n    n_sources : int\n        Number of independent sources.\n\n    Returns\n    -------\n    torch.Tensor\n        Tensor of shape (n_sources, 1, nt) with repeated wavelet for all sources.\n    \"\"\"\n    return -ricker(freq, nt, dt, peak_time).repeat(n_sources, 1).view(n_sources, 1, -1)\n\n\ndef make_src_locations(src_indices: list[int]) -> torch.Tensor:\n    \"\"\"\n    Construct source coordinates at surface depth.\n\n    Parameters\n    ----------\n    src_indices : list[int]\n        Indices of horizontal source locations.\n\n    Returns\n    -------\n    torch.Tensor\n        Tensor of shape (n_sources, 1, 2) containing (z, x) coordinates.\n    \"\"\"\n    n_sources = len(src_indices)\n    src_locations = torch.zeros(n_sources, 2, dtype=torch.float32)\n    src_locations[:, 1] = torch.tensor(src_indices, dtype=torch.long)\n    return src_locations.view(n_sources, 1, 2)\n\n\ndef make_rec_locations(n_receivers: int, n_sources: int) -> torch.Tensor:\n    \"\"\"\n    Construct receiver coordinates for all sources.\n\n    Parameters\n    ----------\n    n_receivers : int\n        Number of receivers.\n    n_sources : int\n        Number of independent sources.\n\n    Returns\n    -------\n    torch.Tensor\n        Tensor of shape (n_sources, n_receivers, 2) containing (z, x) coordinates.\n    \"\"\"\n    rec_locations = torch.zeros(n_receivers, 2, dtype=torch.float32)\n    rec_locations[:, 1] = torch.arange(n_receivers, dtype=torch.float32)\n    return rec_locations.unsqueeze(0).repeat(n_sources, 1, 1)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.127385Z","iopub.execute_input":"2025-06-04T21:54:53.127608Z","iopub.status.idle":"2025-06-04T21:54:53.146089Z","shell.execute_reply.started":"2025-06-04T21:54:53.127589Z","shell.execute_reply":"2025-06-04T21:54:53.145372Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@dataclass(frozen=True)\nclass OpenFWIConfig:\n    \"\"\"\n    Configuration of physical parameters for seismic forward modeling and inversion.\n\n    These configurations should be the same throughout our training dataset.\n\n    Attributes\n    ----------\n    dx : float\n        Spatial grid spacing in x-direction (meters).\n    dz : float\n        Spatial grid spacing in z-direction (meters).\n    dt : float\n        Temporal step size (seconds).\n    nt : int\n        Number of time steps.\n    nx : int\n        Number of horizontal grid points.\n    nz : int\n        Number of vertical grid points.\n    freq : float\n        Ricker wavelet central frequency (Hz).\n    peak_index : int\n        Time step where the wavelet is at its peak.\n    unexplained_offset : int\n        Offset used to determine peak time in wavelet generation relative to `peak_time` (index 76).\n    src_indices : list[int]\n        Indices of horizontal source locations.\n    \"\"\"\n\n    dx: float = 10.0\n    dz: float = 10.0\n    dt: float = 0.001\n    nt: int = 1000\n    nx: int = 70\n    nz: int = 70\n    freq: float = 15.0\n    peak_index: int = 76\n    unexplained_offset: int = -4\n    src_indices: list[int] = field(default_factory=lambda: [0, 17, 34, 52, 69])","metadata":{"trusted":true,"_kg_hide-input":false,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.146887Z","iopub.execute_input":"2025-06-04T21:54:53.147119Z","iopub.status.idle":"2025-06-04T21:54:53.162061Z","shell.execute_reply.started":"2025-06-04T21:54:53.147101Z","shell.execute_reply":"2025-06-04T21:54:53.161403Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VelocityModel(torch.nn.Module):\n    \"\"\"\n    Put velocity field through a sigmoid to resolve extreme velocity values, by mapping it within given bounds [v_min, v_max].\n\n    This is how one can (potentially) include neural networks in the physics model.\n    Here it is just a sigmoid function. But you can put any MLP/CNN etc here.\n\n    Parameters\n    ----------\n    initial : torch.Tensor\n        Initial guess for the velocity field. Must be in the range [v_min, v_max].\n    v_min : float\n        Minimum allowed velocity value.\n    v_max : float\n        Maximum allowed velocity value.\n    \"\"\"\n    def __init__(self, initial: torch.Tensor, v_min: float, v_max: float):\n        super().__init__()\n        self.v_min = v_min\n        self.v_max = v_max\n        normalized = (initial - v_min) / (v_max - v_min)\n        self.model = torch.nn.Parameter(torch.logit(normalized.clamp(1e-6, 1 - 1e-6)))\n\n    def forward(self) -> torch.Tensor:\n        return torch.sigmoid(self.model) * (self.v_max - self.v_min) + self.v_min","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.162791Z","iopub.execute_input":"2025-06-04T21:54:53.163025Z","iopub.status.idle":"2025-06-04T21:54:53.180481Z","shell.execute_reply.started":"2025-06-04T21:54:53.163009Z","shell.execute_reply":"2025-06-04T21:54:53.179915Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data loading","metadata":{}},{"cell_type":"code","source":"# A for-loop to demonstrate how to iterate through the entire dataset.\n#  As this is just a code demo, I added `break` statements to prevent it from running.\nfor family_name in FamilyName:\n    break  # TODO: Remove to use this loop.\n\n    # Sequential iteration\n    for i, (path_seis, path_vel) in enumerate(iter_sample_pairs(path_train, family_name, random_sample=False)):\n        print(f\"{i}-th sample of {family_name.name} -> Seismic: {path_seis.name}, Velocity: {path_vel.name}\")\n        if i == 2:\n            break  # TODO: Remove to use this loop.\n\n    # Random sampling\n    gen = iter_sample_pairs(path_train, family_name, random_sample=True)\n    path_seis, path_vel = next(gen)\n    print(f\"Random sample {family_name.name} -> Seismic: {path_seis.name}, Velocity: {path_vel.name}\")\n\n    # Batch of 500 input data per file pair!\n    batch_seismic_data = np.load(path_seis)\n    batch_velocity_data = np.load(path_vel)\n    print(f\"Shape of a batch of seismic data  : {batch_seismic_data.shape}\")\n    print(f\"Shape of a batch of velocity data : {batch_velocity_data.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.181379Z","iopub.execute_input":"2025-06-04T21:54:53.181660Z","iopub.status.idle":"2025-06-04T21:54:53.196780Z","shell.execute_reply.started":"2025-06-04T21:54:53.181637Z","shell.execute_reply":"2025-06-04T21:54:53.196240Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load random sample instead.\nfamily_name = random.choice(list(FamilyName))\ngen = iter_sample_pairs(path_train, family_name, random_sample=True)\npath_seis, path_vel = next(gen)\nprint(f\"Random sample {family_name.name} -> Seismic: {path_seis.name}, Velocity: {path_vel.name}\")\n\nbatch_seismic_data = np.load(path_seis)\nbatch_velocity_data = np.load(path_vel)\nprint(f\"Shape of a batch of seismic data  : {batch_seismic_data.shape}\")\nprint(f\"Shape of a batch of velocity data : {batch_velocity_data.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.197438Z","iopub.execute_input":"2025-06-04T21:54:53.197640Z","iopub.status.idle":"2025-06-04T21:54:53.499734Z","shell.execute_reply.started":"2025-06-04T21:54:53.197621Z","shell.execute_reply":"2025-06-04T21:54:53.498910Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"random_index = random.choice(range(batch_seismic_data.shape[0]))\nprint(f\"Randomly select index {random_index}.\")\nsample_seis = torch.Tensor(batch_seismic_data[random_index])\nsample_vel = torch.Tensor(batch_velocity_data[random_index].squeeze())\nprint(f\"sample_seis: {sample_seis.shape}\")\nprint(f\"sample_vel: {sample_vel.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.502275Z","iopub.execute_input":"2025-06-04T21:54:53.502539Z","iopub.status.idle":"2025-06-04T21:54:53.507777Z","shell.execute_reply.started":"2025-06-04T21:54:53.502524Z","shell.execute_reply":"2025-06-04T21:54:53.507056Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## User configuration","metadata":{}},{"cell_type":"code","source":"@dataclass\nclass TrainConfig:\n    \"\"\"\n    Configuration for training loop.\n\n    Attributes\n    ----------\n    init_with_true_vel : bool\n        Whether to use a blurred version of ground truth to init the velocities. Demo only. Default: False.\n    epochs : int\n        Number of training epochs.\n    lr_init : float\n        Adam learning rate, initial.\n    lr_peak : float\n        Adam learning rate, peak.\n        Use a crazy learning rate here in case the initial velocity guess is way off the target,\n        if we use direct physics optimization..\n    lr_final : float\n        Adam learning rate, final.\n    downscale_lr : float\n        Should you choose to use the above VelocityModel, the learning rate should be way lower, like in normal ML.\n    steps_warmup : int\n        Steps for scheduler to warm up learning rate.\n    steps_plateau : int\n        Steps for scheduler to hold the learning rate constant.\n    steps_cooldown : int\n        Steps for scheduler to cool down learning rate to a final base rate.\n    \"\"\"\n\n    init_with_true_vel: bool = False\n    epochs: int = 1000\n    lr_init: float = 1e-2\n    lr_peak: float = 1e0\n    lr_final: float = 1e-2\n    downscale_lr = 1e3\n    steps_warmup: int = 100\n    steps_plateau: int = 500\n    steps_cooldown: int = 400","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.508524Z","iopub.execute_input":"2025-06-04T21:54:53.508745Z","iopub.status.idle":"2025-06-04T21:54:53.522869Z","shell.execute_reply.started":"2025-06-04T21:54:53.508723Z","shell.execute_reply":"2025-06-04T21:54:53.522236Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training loop","metadata":{}},{"cell_type":"code","source":"def fit_one_sample(\n    config_openfwi: OpenFWIConfig,\n    config_train: TrainConfig,\n    target_seis: torch.Tensor,\n    target_vel: torch.Tensor | None = None,\n    use_nn: bool = False,\n) -> tuple[torch.Tensor, torch.Tensor | torch.nn.Module, torch.Tensor, torch.Tensor, dict[str : list[float]]]:\n    \"\"\"\n    Run full training loop for a single sample in the batch.\n\n    Parameters\n    ----------\n    config_openfwi : OpenFWIConfig\n        Configuration parameters for physical simulation grid.\n    config_train : TrainConfig\n        Configuration parameters for training loop.\n    target_seis : torch.Tensor\n        One sample from a batch of seismic data, shape (S, T, R).\n    target_vel (optional) : torch.Tensor\n        One sample from a batch of velocity models, shape (H, W). Only used for eval.\n    use_nn : bool\n        Whether to use the example `VelocityModel(nn.Module)` class.\n\n    Returns\n    -------\n    pred_seis : torch.Tensor\n        Predicted seismogram, shape (S, T, R).\n    pred_vel : torch.Tensor | torch.nn.Module\n        Final optimized velocity model, shape (H, W).\n    init_vel : torch.Tensor\n        Initial velocity model, shape (H, W), for debug.\n    wavelet : torch.Tensor\n        Generated wavelet, shape (S, 1, T), for debug.\n    losses : dict[str: list[float]]]\n        A dictionary of named losses over epochs.\n    \"\"\"\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    target_seis = target_seis.to(device)\n    if target_vel is not None:\n        target_vel = target_vel.to(device)\n\n    wavelet = make_ricker_wavelet(\n        freq=config_openfwi.freq,\n        nt=config_openfwi.nt,\n        dt=config_openfwi.dt,\n        peak_time=(config_openfwi.peak_index + config_openfwi.unexplained_offset) / config_openfwi.nt,\n        n_sources=len(config_openfwi.src_indices),\n    ).to(device)\n    src_locations = make_src_locations(src_indices=config_openfwi.src_indices).to(device)\n    rec_locations = make_rec_locations(n_receivers=config_openfwi.nx, n_sources=len(config_openfwi.src_indices)).to(\n        device\n    )\n\n    if target_vel is not None and config_train.init_with_true_vel:\n        # If `target_vel` is provided, this means we're not in test-time, and it's okay to cheat a little bit.\n        #  We initialize the model with a blurred target in order to help the model converge easier.\n        init_vel = torch.tensor(1 / gaussian_filter(1 / target_vel.cpu().numpy(), 5))  # .requires_grad_()\n    else:\n        # init_vel = torch.ones(config_openfwi.nz, config_openfwi.nx, dtype=torch.float32) * 3500\n        init_vel = torch.linspace(3000, 4000, steps=config_openfwi.nz, dtype=torch.float32)\n        init_vel = init_vel.unsqueeze(1).repeat(1, config_openfwi.nx)\n\n    # You can choose to either optimize the velocity field directly as a trainable Tensor,\n    #  or use a PyTorch nn.Module class.\n    # Option 1) Wrap a model to help convergence. Naturally a model is trainable.\n    if use_nn:\n        init_vel = init_vel.requires_grad_()\n        velocity_model = VelocityModel(init_vel, 1000, 8000).to(device)\n        downscale_lr = config_train.downscale_lr\n    # Option 2) `nn.Parameter` is still a Tensor and can be manipulated as such. But now it is trainable.\n    else:\n        velocity_model = nn.Parameter(init_vel.clone().detach().to(device))\n        downscale_lr = 1.0\n\n    if isinstance(velocity_model, nn.Module):\n        optimizer = torch.optim.Adam(velocity_model.parameters(), lr=config_train.lr_init / downscale_lr)\n    else:\n        optimizer = torch.optim.Adam([velocity_model], lr=config_train.lr_init)\n\n    scheduler = WarmupPlateauCooldownLR(\n        optimizer=optimizer,\n        lr_init=config_train.lr_init / downscale_lr,\n        lr_peak=config_train.lr_peak / downscale_lr,\n        lr_final=config_train.lr_final / downscale_lr,\n        steps_warmup=config_train.steps_warmup,\n        steps_plateau=config_train.steps_plateau,\n        steps_cooldown=config_train.steps_cooldown,\n    )\n\n    criterion_seis = LossAggregator()\n    criterion_seis.register(\"ssim\", ssim_loss, weight=50.0)\n    criterion_seis.register(\"l1\", nn.L1Loss(), weight=10.0)\n    criterion_seis.register(\"l2\", nn.MSELoss(), weight=5.0)\n    criterion_seis.register(\"fourier\", fourier_domain_loss, weight=50.0)\n    criterion_seis.register(\"edge\", edge_loss, weight=10.0)\n\n    criterion_vel = LossAggregator()\n    criterion_vel.register(\"monotonic\", monotonic_depth_loss, weight=0.5)\n    criterion_vel.register(\"tv\", total_variation_loss, weight=0.01)\n    criterion_vel.register(\"laplacian\", laplacian_loss, weight=0.005)\n\n    # TODO: Upsample `wavelet` here when ground truth breaks CFL conditions.\n\n    losses: dict[str, list[float]] = {}\n\n    for step in (pbar := tqdm(range(config_train.epochs))):  # or desired number of iterations\n        optimizer.zero_grad()\n\n        # My own convention: model is normal, field is transposed.\n        #  But this is probably me using Deepwave wrong.\n        if isinstance(velocity_model, nn.Module):\n            velocity_field = velocity_model()\n        else:\n            velocity_field = velocity_model\n\n        # Forward modeling.\n        #  It returns a tuple of 7 elements, only the last element is the relevant `receiver_amplitudes`.\n        out = scalar(\n            v=velocity_field,\n            grid_spacing=config_openfwi.dx,  # Union[int, float, List[float], Tensor]\n            dt=config_openfwi.dt,  # float\n            source_amplitudes=wavelet,\n            source_locations=src_locations,\n            receiver_locations=rec_locations,\n            # nt=nt,  # You cannot specify both the source amplitudes and `nt`.\n            accuracy=8,  # Default: 4. Max: 8.\n            # pml_width=20,  # Default: 20.\n            pml_freq=config_openfwi.freq,\n            freq_taper_frac=0.2,\n            time_pad_frac=0.2,\n            time_taper=True,\n        )\n        # Need to swap output dimensions because time step comes before spatial width in our target receiver amplitudes.\n        pred_seis = out[-1]  # (5, 70, 1000)\n        pred_seis = pred_seis.movedim(-2, -1)  # (5, 1000, 70)\n        # wavefield_nt = out[0]\n        # assert wavefield_nt.shape == (5, 110, 110)\n        # wavefield_nt = wavefield_nt[..., 20:90, 20:90]\n\n        seis_loss, seis_loss_components = criterion_seis(\n            pred=pred_seis.unsqueeze(0),\n            target=target_seis.unsqueeze(0),\n            return_components=True,\n            components_with_weight=True,\n        )\n        vel_loss, vel_loss_components = criterion_vel(\n            pred=velocity_field.unsqueeze(0).unsqueeze(0),\n            target=None,\n            return_components=True,\n            components_with_weight=True,\n        )\n        # Example scales of loss terms:\n        # {'ssim': 0.03497380018234253, 'l1': 0.20045273005962372, 'fourier': 0.00616673706099391, 'edge': 0.049981553107500076}\n        # {'monotonic': 1.7209473848342896, 'tv': 6.263477802276611, 'laplacian': 237.21453857421875}\n\n        # We won't have `target_vel` during test-time.\n        if target_vel is not None:\n            vel_mae_loss = nn.L1Loss()(velocity_field, target_vel)\n        else:\n            vel_mae_loss = None\n\n        loss = seis_loss + vel_loss\n        loss.backward()\n\n        # Clip gradients by quantile.\n        if isinstance(velocity_model, nn.Module):\n            all_grads = torch.cat([p.grad.detach().abs().flatten() for p in velocity_model.parameters() if p.grad is not None])\n            global_clip_value = torch.quantile(all_grads, 0.98)\n            torch.nn.utils.clip_grad_value_(velocity_model.parameters(), global_clip_value.item())\n        else:\n            torch.nn.utils.clip_grad_value_(velocity_field, torch.quantile(velocity_field.grad.detach().abs(), 0.98))\n\n        optimizer.step()\n        scheduler.step()\n\n        pbar.set_description(f\"Step {step:03d}\")\n        pbar.set_postfix(\n            lr=f\"{scheduler.get_last_lr()[0]:.1e}\",\n            seis_loss=f\"{seis_loss.item():.6f}\",\n            vel_loss=f\"{vel_loss.item():.6f}\",\n            vel_mae_loss=f\"{vel_mae_loss.item():.6f}\" if vel_mae_loss is not None else None,\n        )\n\n        for key, val in seis_loss_components.items():\n            key = f\"seis_{key}\"\n            if key not in losses:\n                losses[key] = []\n            losses[key].append(val)\n        for key, val in vel_loss_components.items():\n            key = f\"vel_{key}\"\n            if key not in losses:\n                losses[key] = []\n            losses[key].append(val)\n        if vel_mae_loss is not None:\n            key = \"vel_eval_mae\"\n            if key not in losses:\n                losses[key] = []\n            losses[key].append(vel_mae_loss.item())\n\n    return (\n        pred_seis.detach().cpu(),\n        velocity_model.detach().cpu() if isinstance(velocity_model, torch.Tensor) else velocity_model.cpu(),\n        init_vel.detach().cpu(),\n        wavelet.detach().cpu(),\n        losses,\n    )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:54:53.523738Z","iopub.execute_input":"2025-06-04T21:54:53.524244Z","iopub.status.idle":"2025-06-04T21:54:53.545115Z","shell.execute_reply.started":"2025-06-04T21:54:53.524220Z","shell.execute_reply":"2025-06-04T21:54:53.544422Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Main simulation code here!\n#  `sample_vel` is only used for reference (vel_mae_loss). Not used in training unless you force it in `TrainConfig`.\npred_seis, pred_vel, init_vel, wavelet, losses = fit_one_sample(\n    config_openfwi=OpenFWIConfig(),\n    config_train=TrainConfig(),\n    target_seis=sample_seis,\n    target_vel=sample_vel,\n    use_nn=False,  # Set `False` to demonstrate direct physics optimization.\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:58:10.271841Z","iopub.execute_input":"2025-06-04T21:58:10.272176Z","iopub.status.idle":"2025-06-04T21:59:15.960445Z","shell.execute_reply.started":"2025-06-04T21:58:10.272153Z","shell.execute_reply":"2025-06-04T21:59:15.959585Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_named_losses(losses: dict[str, list[float]], skip_first: int = 0) -> plt.Figure:\n    \"\"\"\n    Plot named loss curves over epochs.\n\n    Parameters\n    ----------\n    losses : dict of str to list of float\n        A dictionary where each key is the name of a loss component\n        and each value is a list of loss values per epoch.\n    skip_first : int\n        Number of initial epochs to skip.\n\n    Returns\n    -------\n    fig : matplotlib.figure.Figure\n        The matplotlib Figure object containing the loss plots.\n    \"\"\"\n    fig = plt.figure(figsize=(12, 8))\n    ax = fig.add_subplot(1, 1, 1)\n\n    for idx, (name, values) in enumerate(losses.items()):\n        epochs = range(1, len(values) + 1)\n        if name[:9] == \"vel_eval_\":\n            ls = \"-.\"\n            values = [v / 100.0 for v in values]\n            name = f\"{name} / 100\"\n        elif name[:4] == \"vel_\":\n            ls = \"--\"\n        else:\n            ls = \"-\"\n        ax.plot(epochs[skip_first:], values[skip_first:], label=name, ls=ls, lw=3)\n\n    ax.set_title(\"Loss Curves Over Epochs\")\n    ax.set_xlabel(\"Epoch\")\n    ax.set_ylabel(\"Loss\")\n    ax.legend()\n    ax.grid(True)\n\n    fig.tight_layout()\n    return fig\n\n\nfig = plot_named_losses(losses, skip_first=200)\nfig.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:59:15.962003Z","iopub.execute_input":"2025-06-04T21:59:15.962252Z","iopub.status.idle":"2025-06-04T21:59:16.346091Z","shell.execute_reply.started":"2025-06-04T21:59:15.962232Z","shell.execute_reply":"2025-06-04T21:59:16.345270Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visualizations\n\nWe visualize the fitted receivers' signal and the optimized velocity model.","metadata":{}},{"cell_type":"code","source":"def compare_seismogram_traces(\n    pred_seis: torch.Tensor,\n    true_seis: torch.Tensor,\n    family_name: str,\n    source_idx_to_plot: int = 0,\n    receiver_idx_to_plot: int = 0,\n    num_traces_to_plot: int = 20,\n    offset_trace: float = 10.0,\n) -> plt.Figure:\n    \"\"\"\n    Plot predicted and ground truth seismogram traces for a single source across multiple receivers.\n\n    Parameters\n    ----------\n    pred_seis : torch.Tensor\n        Predicted seismogram tensor of shape (num_sources, time_steps, num_receivers).\n    true_seis : torch.Tensor\n        Ground truth seismogram tensor with the same shape as pred_seis.\n    family_name : str\n        Identifier for the seismic shot family, used in the plot title.\n    source_idx_to_plot : int, optional\n        Index of the source to visualize (default is 0, max is 4).\n    receiver_idx_to_plot : int, optional\n        Starting index of receiver traces to visualize (default is 0).\n    num_traces_to_plot : int, optional\n        Number of receiver traces to visualize (default is 20).\n    offset_trace : float, optional\n        Vertical offset between traces for visual separation (default is 10.0).\n\n    Returns\n    -------\n    matplotlib.figure.Figure\n        The Matplotlib Figure object containing the seismogram plot.\n    \"\"\"\n\n    num_sources, time_steps, num_receivers = pred_seis.shape\n\n    assert source_idx_to_plot < num_sources, \"Invalid source index\"\n    assert receiver_idx_to_plot + num_traces_to_plot <= num_receivers, \"Receiver range exceeds bounds\"\n\n    fig, ax = plt.subplots(figsize=(16, max(4, num_traces_to_plot / 2)), constrained_layout=True)\n    prop_cycle = plt.rcParams[\"axes.prop_cycle\"]\n    colors = prop_cycle.by_key()[\"color\"]\n\n    legend_handles = []\n    legend_labels = []\n\n    for i in range(num_traces_to_plot):\n        color = colors[i % len(colors)]\n\n        trace_pred = pred_seis[source_idx_to_plot, :, receiver_idx_to_plot + i] - i * offset_trace\n        trace_true = true_seis[source_idx_to_plot, :, receiver_idx_to_plot + i] - i * offset_trace\n        line_pred = ax.plot(trace_pred, color=color, alpha=0.5)\n        line_true = ax.plot(trace_true, color=color, alpha=0.7, ls=\"--\")\n\n        proxy_pred = mlines.Line2D([], [], color=color, alpha=0.5)\n        proxy_true = mlines.Line2D([], [], color=color, alpha=0.7, linestyle=\"--\")\n        legend_handles.append((proxy_pred, proxy_true))\n        legend_labels.append(f\"Receiver {receiver_idx_to_plot + i:02d} (Pred vs GT)\")\n\n    ax.set_title(f\"Seismogram from {family_name}: Source #{source_idx_to_plot}\")\n    ax.set_xlabel(\"Time Step (ms)\")\n    ax.set_ylabel(\"Amplitude (Offset)\")\n    ax.set_yticks([])\n    ax.grid(True)\n    ax.legend(\n        legend_handles,\n        legend_labels,\n        handler_map={tuple: HandlerTuple(ndivide=None)},\n        handlelength=4,\n        loc=\"upper right\",\n    )\n    return fig\n\n\nfig = compare_seismogram_traces(\n    pred_seis=pred_seis.detach().cpu(),\n    true_seis=sample_seis.detach().cpu(),\n    family_name=family_name.name,\n    source_idx_to_plot=0,\n    receiver_idx_to_plot=0,\n    num_traces_to_plot=20,\n    offset_trace=10.0,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:59:16.347114Z","iopub.execute_input":"2025-06-04T21:59:16.347690Z","iopub.status.idle":"2025-06-04T21:59:17.536790Z","shell.execute_reply.started":"2025-06-04T21:59:16.347660Z","shell.execute_reply":"2025-06-04T21:59:17.536092Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compare_seismograms(\n    pred_seis: torch.Tensor,\n    true_seis: torch.Tensor,\n) -> plt.Figure:\n    \"\"\"\n    Plot predicted vs ground truth seismograms for 5 shots, including absolute difference.\n\n    Parameters\n    ----------\n    pred_seis : torch.Tensor\n        Predicted seismogram tensor of shape (5, 1000, 70), where 5 is the number of shots,\n        1000 is the number of time steps, and 70 is the number of receivers.\n    true_seis : torch.Tensor\n        Ground truth seismogram tensor with the same shape as pred_seis.\n\n    Returns\n    -------\n    matplotlib.figure.Figure\n        The Matplotlib Figure object with all subplots.\n    \"\"\"\n    data_dict = {\n        \"Prediction\": pred_seis,\n        \"Ground Truth\": true_seis,\n        \"Abs Diff\": torch.abs(pred_seis - true_seis),\n    }\n    n_shots, n_timesteps, n_receivers = next(iter(data_dict.values())).shape\n    vmin = min(arr.min() for arr in data_dict.values())\n    vmax = max(arr.max() for arr in data_dict.values())\n    cmap = \"seismic\"\n\n    fig = plt.figure(figsize=(24, 5 * len(data_dict)), constrained_layout=True)\n    gs = gridspec.GridSpec(\n        len(data_dict),\n        n_shots + 1,\n        figure=fig,\n        width_ratios=[1] * n_shots + [0.05],\n        height_ratios=[1] * len(data_dict),\n        hspace=0.0,\n        wspace=0.1,\n    )\n\n    for row_idx, (label, data) in enumerate(data_dict.items()):\n        for col_idx in range(n_shots):\n            ax = fig.add_subplot(gs[row_idx, col_idx])\n            im = ax.imshow(data[col_idx, :, :], aspect=\"auto\", cmap=cmap, vmin=vmin, vmax=vmax)\n            if row_idx == 0:\n                ax.set_title(f\"Source #{col_idx + 1}\")\n            ax.set_xlabel(\"70x Receivers\")\n            if col_idx == 0:\n                ax.set_ylabel(f\"{label}\\n1000x Time Steps\")\n            else:\n                ax.set_yticks([])\n\n    cax = fig.add_subplot(gs[:, -1])\n    cbar = fig.colorbar(im, cax=cax)\n    cbar.set_label(\"Amplitude\", fontsize=12)\n\n    fig.suptitle(\n        \"Ground Truth vs Prediction of 5 Shots (observed from 70 receivers over 1000 time steps)\", fontsize=18, y=1.02\n    )\n\n    return fig\n\n\nfig = compare_seismograms(pred_seis=pred_seis.detach().cpu(), true_seis=sample_seis.detach().cpu())\nfig.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:59:17.538455Z","iopub.execute_input":"2025-06-04T21:59:17.539230Z","iopub.status.idle":"2025-06-04T21:59:20.537857Z","shell.execute_reply.started":"2025-06-04T21:59:17.539177Z","shell.execute_reply":"2025-06-04T21:59:20.537088Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compare_velocity_models(\n    init_vel: torch.Tensor,\n    pred_vel: torch.Tensor,\n    true_vel: torch.Tensor,\n    family_name: str,\n) -> plt.Figure:\n    \"\"\"\n    Compare initial guess, prediction, and ground truth velocity models side-by-side.\n\n    Parameters\n    ----------\n    init_vel : torch.Tensor\n        Tensor containing the initial guess velocity model, shape (1, H, W) or (H, W).\n    pred_vel : torch.Tensor\n        Tensor containing the predicted velocity model, shape (1, H, W) or (H, W).\n    true_vel : torch.Tensor\n        Tensor containing the ground truth velocity model, shape (1, H, W) or (H, W).\n    family_name : str\n        Identifier for the seismic shot family, used in the plot title.\n\n    Returns\n    -------\n    matplotlib.figure.Figure\n        The Matplotlib Figure object containing the velocity model plots.\n    \"\"\"\n    data_dict = {\n        \"Initial Guess\": init_vel.squeeze(),\n        \"Prediction\": pred_vel.squeeze(),\n        \"Ground Truth\": true_vel.squeeze(),\n    }\n    dx = 10  # 10 meters per index\n    height, width = next(iter(data_dict.values())).shape\n    x_ticks = np.linspace(0, width - 1, 6)  # 6 ticks across width\n    y_ticks = np.linspace(0, height - 1, 6)  # 6 ticks down height\n    x_tick_labels = [f\"{i * dx:.0f}\" for i in x_ticks]\n    y_tick_labels = [f\"{-i * dx:.0f}\" for i in y_ticks]\n    vmin = min(arr.min() for arr in data_dict.values())\n    vmax = max(arr.max() for arr in data_dict.values())\n\n    fig = plt.figure(figsize=(18, 6), constrained_layout=True)\n    gs = gridspec.GridSpec(1, 4, figure=fig, width_ratios=[1, 1, 1, 0.05], wspace=0.1)\n\n    for idx, (title, data) in enumerate(data_dict.items()):\n        ax = fig.add_subplot(gs[0, idx])\n        im = ax.imshow(data, cmap=\"seismic\", aspect=\"equal\", vmin=vmin, vmax=vmax)\n        ax.set_title(title)\n        ax.set_xlabel(\"Width (m)\")\n        ax.set_xticks(x_ticks)\n        ax.set_xticklabels(x_tick_labels)\n        if idx == 0:\n            ax.set_ylabel(\"Depth (m)\")\n            ax.set_yticks(y_ticks)\n            ax.set_yticklabels(y_tick_labels)\n        else:\n            ax.set_yticks([])\n\n    cax = fig.add_subplot(gs[:, 3])\n    cbar = fig.colorbar(im, cax=cax)\n    cbar.set_label(\"Velocity (m/s)\")\n\n    fig.suptitle(f\"Compare Velocity Models ({family_name})\", fontsize=18, y=1.0)\n\n    return fig\n\n\nfig = compare_velocity_models(\n    init_vel=init_vel.detach().cpu(),\n    pred_vel=pred_vel.detach().cpu() if isinstance(pred_vel, torch.Tensor) else pred_vel().detach().cpu(),\n    true_vel=sample_vel.detach().cpu(),\n    family_name=family_name.name,\n)\nfig.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:59:20.538839Z","iopub.execute_input":"2025-06-04T21:59:20.539656Z","iopub.status.idle":"2025-06-04T21:59:21.125491Z","shell.execute_reply.started":"2025-06-04T21:59:20.539623Z","shell.execute_reply":"2025-06-04T21:59:21.124717Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_wavelet(wavelet_: np.ndarray) -> plt.Figure:\n    \"\"\"\n    Plot a Ricker wavelet in time and frequency domains.\n\n    Parameters\n    ----------\n    wavelet_ : np.ndarray\n        1D array representing the Ricker wavelet amplitude over time.\n        Assumes sampling rate of 1000 Hz and unit duration.\n\n    Returns\n    -------\n    matplotlib.figure.Figure\n        The Matplotlib Figure object containing the two subplots.\n    \"\"\"\n    sampling_rate = 1000  # Hz\n    dt = 1 / sampling_rate\n\n    # Frequency spectrum via FFT\n    N = len(wavelet_)\n    t = np.arange(0, N * dt, dt)\n    freq = fftshift(fftfreq(N, d=dt))\n    spectrum = np.abs(fftshift(fft(wavelet_)))\n    peak_freq = freq[spectrum.argmax()]\n\n    fig = plt.figure(figsize=(12, 6), constrained_layout=True)\n    gs = fig.add_gridspec(1, 2, figure=fig, width_ratios=[1, 1])\n\n    # Time domain plot\n    ax0 = fig.add_subplot(gs[0])\n    ax0.plot(t, wavelet_, color=\"C0\")\n    ax0.set_title(f\"Ricker Wavelet (Peak @ {abs(peak_freq)} Hz)\")\n    ax0.set_xlim(0, 0.2)  # restrict to relevant frequencies\n    ax0.set_xlabel(\"Time [s]\")\n    ax0.set_ylabel(\"Amplitude\")\n    ax0.grid(True)\n\n    # Frequency domain plot\n    ax1 = fig.add_subplot(gs[1])\n    ax1.plot(freq, spectrum, color=\"C1\")\n    ax1.set_xlim(0, 100)  # restrict to relevant frequencies\n    ax1.set_title(\"Frequency Spectrum\")\n    ax1.set_xlabel(\"Frequency [Hz]\")\n    ax1.set_ylabel(\"Magnitude\")\n    ax1.grid(True)\n\n    return fig\n\n\nfig = plot_wavelet(wavelet.numpy().squeeze()[0])\nfig.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:59:21.126508Z","iopub.execute_input":"2025-06-04T21:59:21.126818Z","iopub.status.idle":"2025-06-04T21:59:21.631696Z","shell.execute_reply.started":"2025-06-04T21:59:21.126793Z","shell.execute_reply":"2025-06-04T21:59:21.630949Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Animate wavefield\n\nAfter we've solved for a velocity model, we can visualize how the acoustic wave propagated from the source through a forward simulation.","metadata":{}},{"cell_type":"code","source":"def forward_propagation(\n    config_openfwi: OpenFWIConfig,\n    config_train: TrainConfig,\n    target_seis: torch.Tensor,\n    velocity_model: torch.Tensor | torch.nn.Module,\n) -> torch.Tensor:\n    \"\"\"\n    Forward propagate a fixed velocity model to obtain the wavefield history.\n\n    Note how there is no backprop needed here.\n\n    Parameters\n    ----------\n    config_openfwi : OpenFWIConfig\n        Configuration parameters for physical simulation grid.\n    config_train : TrainConfig\n        Configuration parameters for training loop.\n    target_seis : torch.Tensor\n        One sample from a batch of seismic data, shape (S, T, R).\n    velocity_model : torch.Tensor | torch.nn.Module\n        One sample from a batch of ground truth velocity models, shape (H, W), or a fitted one.\n\n    Returns\n    -------\n    wavefields : torch.Tensor\n        Predicted wavefields given velocity model and receiver observations over T, shape (T, S, P + H + P, P + W + P).\n        Here P is padding from `pml_width=20`.\n    \"\"\"\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    target_seis = target_seis.to(device)\n    velocity_model = velocity_model.clone().detach().requires_grad_(False).to(device)\n\n    if isinstance(velocity_model, nn.Module):\n        velocity_field = velocity_model()\n    else:\n        velocity_field = velocity_model\n\n    wavelet = make_ricker_wavelet(\n        freq=config_openfwi.freq,\n        nt=config_openfwi.nt,\n        dt=config_openfwi.dt,\n        peak_time=(config_openfwi.peak_index + config_openfwi.unexplained_offset) / config_openfwi.nt,\n        n_sources=len(config_openfwi.src_indices),\n    ).to(device)\n    src_locations = make_src_locations(src_indices=config_openfwi.src_indices).to(device)\n    rec_locations = make_rec_locations(n_receivers=config_openfwi.nx, n_sources=len(config_openfwi.src_indices)).to(\n        device\n    )\n\n    criterion_seis = LossAggregator()\n    criterion_seis.register(\"l1\", nn.L1Loss(), weight=1.0)\n\n    # We'll need to upsample the wavelet (source_amplitudes) because\n    #  ground truth max-velocity can be > 4242m/s, which breaks CFL condition.\n    cfl_dt, step_ratio = deepwave.common.cfl_condition(\n        config_openfwi.dx,\n        config_openfwi.dx,\n        config_openfwi.dt,\n        max_vel=10000,\n    )\n    wavelet = deepwave.common.upsample(wavelet, step_ratio)\n\n    losses: dict[str, list[float]] = {}\n\n    # Initial states\n    wavefield_nt, wavefield_ntm1 = None, None\n    psiy_ntm1, psix_ntm1, zetay_ntm1, zetax_ntm1 = None, None, None, None\n    wavefields = []\n\n    for step in (pbar := tqdm(range(config_openfwi.nt))):\n        # Forward propagation with previous time-step state (enables continuation).\n        #  It returns a tuple of 7 elements, only the last element is the relevant `receiver_amplitudes`.\n        step_ratio = 1\n        wavelet_chunk = wavelet[..., step*step_ratio:(step+1)*step_ratio]\n        out = scalar(\n            v=velocity_field,\n            grid_spacing=config_openfwi.dx,  # Union[int, float, List[float], Tensor]\n            dt=cfl_dt,  # float\n            source_amplitudes=wavelet_chunk,\n            source_locations=src_locations,\n            receiver_locations=rec_locations,\n            # nt=nt,  # You cannot specify both the source amplitudes and `nt`.\n            accuracy=8,  # Default: 4. Max: 8.\n            # pml_width=20,  # Default: 20.\n            pml_freq=config_openfwi.freq,\n            # We have examples where v=2000\n            freq_taper_frac=0.2,\n            time_pad_frac=0.2,\n            time_taper=True,\n            # Here you reinitialize the simulation step from the previous step's wavefields.\n            wavefield_0=wavefield_nt,\n            wavefield_m1=wavefield_ntm1,\n            psiy_m1=psiy_ntm1,\n            psix_m1=psix_ntm1,\n            zetay_m1=zetay_ntm1,\n            zetax_m1=zetax_ntm1,\n        )\n\n        # Extract receiver predictions and updated state for next loop.\n        wavefield_nt, wavefield_ntm1 = out[0].detach(), out[1].detach()\n        psiy_ntm1, psix_ntm1, zetay_ntm1, zetax_ntm1 = out[2].detach(), out[3].detach(), out[4].detach(), out[5].detach()\n        # Need to swap output dimensions because time step comes before spatial width in our target receiver amplitudes.\n        pred_seis = out[-1]  # (5, 70, 1000)\n        pred_seis = pred_seis.movedim(-2, -1)  # (5, 1000, 70)\n        wavefields.append(wavefield_nt.detach())\n\n        # Compute waveform misfit and backpropagate.\n        loss = criterion_seis(\n            pred=pred_seis.unsqueeze(0),\n            target=target_seis[..., step*step_ratio:(step+1)*step_ratio, :].unsqueeze(0),\n        )\n\n        pbar.set_description(f\"Step {step:03d}\")\n        pbar.set_postfix(\n            loss=f\"{loss.item():.6f}\",\n        )\n\n    return torch.stack(wavefields, dim=0)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:59:21.632515Z","iopub.execute_input":"2025-06-04T21:59:21.632763Z","iopub.status.idle":"2025-06-04T21:59:21.643449Z","shell.execute_reply.started":"2025-06-04T21:59:21.632737Z","shell.execute_reply":"2025-06-04T21:59:21.642618Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if isinstance(pred_vel, torch.Tensor):\n    pred_vel_ = pred_vel\nelse:\n    pred_vel_ = pred_vel().detach().cpu()\nprint(\"L1 loss:\", nn.L1Loss()(pred_vel_, sample_vel).item())\nprint(\"Type:\", type(sample_vel), type(pred_vel_))\nprint(\"Shape:\", sample_vel.shape, pred_vel_.shape)\nprint(\"Grad:\", sample_vel.requires_grad, pred_vel_.requires_grad)\nprint(\"Min:\", sample_vel.min().item(), pred_vel_.min().item())\nprint(\"Max:\", sample_vel.max().item(), pred_vel_.max().item())\nprint(\"Std:\", sample_vel.std().item(), pred_vel_.std().item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:59:21.644159Z","iopub.execute_input":"2025-06-04T21:59:21.644441Z","iopub.status.idle":"2025-06-04T21:59:21.661168Z","shell.execute_reply.started":"2025-06-04T21:59:21.644419Z","shell.execute_reply":"2025-06-04T21:59:21.660473Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"true_wavefields = forward_propagation(\n    config_openfwi=OpenFWIConfig(),\n    config_train=TrainConfig(),\n    target_seis=sample_seis,\n    velocity_model=sample_vel,\n)  # True-ish. It's still a simulation, but based on ground-truth velocity fields.\npred_wavefields = forward_propagation(\n    config_openfwi=OpenFWIConfig(),\n    config_train=TrainConfig(),\n    target_seis=sample_seis,\n    velocity_model=pred_vel if isinstance(pred_vel, torch.Tensor) else pred_vel().detach().cpu(),\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:59:21.661818Z","iopub.execute_input":"2025-06-04T21:59:21.662002Z","iopub.status.idle":"2025-06-04T21:59:33.484490Z","shell.execute_reply.started":"2025-06-04T21:59:21.661988Z","shell.execute_reply":"2025-06-04T21:59:33.483675Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"L1 loss:\", nn.L1Loss()(pred_wavefields, true_wavefields).item())\nprint(\"Type:\", type(true_wavefields), type(pred_wavefields))\nprint(\"Shape:\", true_wavefields.shape, pred_wavefields.shape)\nprint(\"Grad:\", true_wavefields.requires_grad, pred_wavefields.requires_grad)\nprint(\"Min:\", true_wavefields.min().item(), pred_wavefields.min().item())\nprint(\"Max:\", true_wavefields.max().item(), pred_wavefields.max().item())\nprint(\"Std:\", true_wavefields.std().item(), pred_wavefields.std().item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:59:33.486744Z","iopub.execute_input":"2025-06-04T21:59:33.486963Z","iopub.status.idle":"2025-06-04T21:59:33.498492Z","shell.execute_reply.started":"2025-06-04T21:59:33.486947Z","shell.execute_reply":"2025-06-04T21:59:33.497631Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def animate_wavefields(\n    pred_wave: torch.Tensor,\n    true_wave: torch.Tensor,\n    fps: int = 50,\n) -> animation.FuncAnimation:\n    \"\"\"\n    Animate predicted vs ground truth wavefields over time for multiple shots.\n\n    Parameters\n    ----------\n    pred_wave : torch.Tensor\n        Predicted wavefields of shape (1000, 5, 110, 110).\n    true_wave : torch.Tensor\n        Ground truth wavefields with the same shape.\n    fps : int\n        Frames per second for animation.\n\n    Returns\n    -------\n    matplotlib.animation.FuncAnimation\n        The animated figure object.\n    \"\"\"\n    assert pred_wave.shape == true_wave.shape\n    n_timesteps, n_shots, H, W = pred_wave.shape\n\n    data_dict = {\n        \"Prediction\": pred_wave[..., 20:90, 20:90],\n        \"Ground Truth\": true_wave[..., 20:90, 20:90],\n        \"Abs Diff\": torch.abs(pred_wave - true_wave)[..., 20:90, 20:90],\n    }\n\n    vmin = min(arr.min().item() for arr in data_dict.values())\n    vmax = max(arr.max().item() for arr in data_dict.values())\n    # cmap = \"seismic\" gray_r, gray, bone\n    cmap = \"bone\"\n\n    fig = plt.figure(figsize=(19.2, 10.8), dpi=100, constrained_layout=True)\n    gs = gridspec.GridSpec(\n        len(data_dict),\n        n_shots + 1,\n        figure=fig,\n        width_ratios=[1] * n_shots + [0.05],\n        height_ratios=[1] * len(data_dict),\n    )\n\n    axes = []\n    ims = []\n\n    for row_idx, (label, data) in enumerate(data_dict.items()):\n        row_axes = []\n        row_ims = []\n        for col_idx in range(n_shots):\n            ax = fig.add_subplot(gs[row_idx, col_idx])\n            im = ax.imshow(data[0, col_idx], vmin=vmin, vmax=vmax, cmap=cmap, animated=True)\n            if row_idx == 0:\n                ax.set_title(f\"Shot #{col_idx + 1}\")\n            if col_idx == 0:\n                ax.set_ylabel(f\"{label}\")\n            ax.set_xticks([])\n            ax.set_yticks([])\n            row_axes.append(ax)\n            row_ims.append(im)\n        axes.append(row_axes)\n        ims.append(row_ims)\n\n    cax = fig.add_subplot(gs[:, -1])\n    cbar = fig.colorbar(ims[0][0], cax=cax)\n    cbar.set_label(\"Amplitude\", fontsize=12)\n\n    fig.suptitle(\"Wavefield Evolution: Prediction vs Ground Truth\", fontsize=16)\n\n    pbar = tqdm(total=n_timesteps, desc=\"Animating frames\", leave=False, dynamic_ncols=True)\n\n    def update(frame_idx, pbar):\n        for row_idx, data in enumerate(data_dict.values()):\n            for col_idx in range(n_shots):\n                ims[row_idx][col_idx].set_array(data[frame_idx, col_idx])\n        pbar.n = frame_idx + 1\n        pbar.refresh()\n        return [im for row in ims for im in row]\n\n    anim = animation.FuncAnimation(\n        fig,\n        partial(update, pbar=pbar),\n        frames=n_timesteps,\n        interval=1000 / fps,\n        blit=True,\n    )\n\n    return anim\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:59:33.499354Z","iopub.execute_input":"2025-06-04T21:59:33.499587Z","iopub.status.idle":"2025-06-04T21:59:33.511110Z","shell.execute_reply.started":"2025-06-04T21:59:33.499569Z","shell.execute_reply":"2025-06-04T21:59:33.510301Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"anim = animate_wavefields(pred_wave=pred_wavefields[::50].detach().cpu(), true_wave=true_wavefields[::50].detach().cpu(), fps=50)\n_ = anim.save(\"wavefield_animation.mp4\", writer=\"ffmpeg\", fps=10, dpi=100)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:59:33.511899Z","iopub.execute_input":"2025-06-04T21:59:33.512161Z","iopub.status.idle":"2025-06-04T21:59:48.621591Z","shell.execute_reply.started":"2025-06-04T21:59:33.512141Z","shell.execute_reply":"2025-06-04T21:59:48.620778Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"HTML(\"\"\"\n<video width=\"1920\" height=\"1080\" controls>\n  <source src=\"wavefield_animation.mp4\" type=\"video/mp4\">\n  Your browser does not support the video tag.\n</video>\n\"\"\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-04T21:59:48.622567Z","iopub.execute_input":"2025-06-04T21:59:48.623518Z","iopub.status.idle":"2025-06-04T21:59:48.628849Z","shell.execute_reply.started":"2025-06-04T21:59:48.623496Z","shell.execute_reply":"2025-06-04T21:59:48.628014Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null}]}