{"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":"nvidiaTeslaT4","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":"Improved version of [Jaewook Kim’s vel_to_seis implementation](https://www.kaggle.com/code/jaewook704/waveform-inversion-vel-to-seis), and PyTorch version.","metadata":{"execution":{"iopub.status.busy":"2025-06-22T07:50:27.043700Z","iopub.execute_input":"2025-06-22T07:50:27.043931Z","iopub.status.idle":"2025-06-22T07:50:27.051857Z","shell.execute_reply.started":"2025-06-22T07:50:27.043913Z","shell.execute_reply":"2025-06-22T07:50:27.050910Z"}}},{"cell_type":"code","source":"import numpy as np\nfrom tqdm import tqdm\nimport torch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:38:52.891598Z","iopub.execute_input":"2025-06-23T06:38:52.891816Z","iopub.status.idle":"2025-06-23T06:38:57.869205Z","shell.execute_reply.started":"2025-06-23T06:38:52.891793Z","shell.execute_reply":"2025-06-23T06:38:57.868364Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seis_files = [\n    \"/kaggle/input/waveform-inversion/train_samples/FlatVel_A/data/data1.npy\",\n    \"/kaggle/input/waveform-inversion/train_samples/FlatVel_B/data/data1.npy\",\n    \n    \"/kaggle/input/waveform-inversion/train_samples/CurveVel_A/data/data1.npy\",\n    \"/kaggle/input/waveform-inversion/train_samples/CurveVel_B/data/data1.npy\",\n\n    \"/kaggle/input/waveform-inversion/train_samples/FlatFault_A/seis2_1_0.npy\",\n    \"/kaggle/input/waveform-inversion/train_samples/FlatFault_B/seis6_1_0.npy\",\n    \n    \"/kaggle/input/waveform-inversion/train_samples/CurveFault_A/seis2_1_0.npy\",\n    \"/kaggle/input/waveform-inversion/train_samples/CurveFault_B/seis6_1_0.npy\",\n\n    \"/kaggle/input/waveform-inversion/train_samples/Style_A/data/data1.npy\",\n    \"/kaggle/input/waveform-inversion/train_samples/Style_B/data/data1.npy\",\n    \n]\nvel_files = [\n    \"/kaggle/input/waveform-inversion/train_samples/FlatVel_A/model/model1.npy\",\n    \"/kaggle/input/waveform-inversion/train_samples/FlatVel_B/model/model1.npy\",\n    \n    \"/kaggle/input/waveform-inversion/train_samples/CurveVel_A/model/model1.npy\",\n    \"/kaggle/input/waveform-inversion/train_samples/CurveVel_B/model/model1.npy\",\n\n    \"/kaggle/input/waveform-inversion/train_samples/FlatFault_A/vel2_1_0.npy\",\n    \"/kaggle/input/waveform-inversion/train_samples/FlatFault_B/vel6_1_0.npy\",\n    \n    \"/kaggle/input/waveform-inversion/train_samples/CurveFault_A/vel2_1_0.npy\",\n    \"/kaggle/input/waveform-inversion/train_samples/CurveFault_B/vel6_1_0.npy\",\n\n    \"/kaggle/input/waveform-inversion/train_samples/Style_A/model/model1.npy\",\n    \"/kaggle/input/waveform-inversion/train_samples/Style_B/model/model1.npy\",\n    \n]\ntypes = [\n    'FlatVel_A', 'FlatVel_B', 'CurveVel_A', 'CurveVel_B', 'FlatFault_A', 'FlatFault_B', 'CurveFault_A', 'CurveFault_B', 'Style_A', 'Style_B'\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:38:57.870533Z","iopub.execute_input":"2025-06-23T06:38:57.870903Z","iopub.status.idle":"2025-06-23T06:38:57.876166Z","shell.execute_reply.started":"2025-06-23T06:38:57.870876Z","shell.execute_reply":"2025-06-23T06:38:57.875392Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:40:53.064216Z","iopub.execute_input":"2025-06-23T06:40:53.064900Z","iopub.status.idle":"2025-06-23T06:40:53.152011Z","shell.execute_reply.started":"2025-06-23T06:40:53.064875Z","shell.execute_reply":"2025-06-23T06:40:53.151020Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_vels = []\nall_seis = []\n\nfor Z in tqdm(range(len(vel_files))):\n    vel = torch.from_numpy(np.load(vel_files[Z])).float().to(device)\n    seis = torch.from_numpy(np.load(seis_files[Z])).float().to(device)\n    all_vels.append(vel)\n    all_seis.append(seis)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:40:54.631201Z","iopub.execute_input":"2025-06-23T06:40:54.632059Z","iopub.status.idle":"2025-06-23T06:41:35.193395Z","shell.execute_reply.started":"2025-06-23T06:40:54.632030Z","shell.execute_reply":"2025-06-23T06:41:35.192812Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# No torch (for benchmark)\n\nCode base from [Jaewook Kim’s notebook](https://www.kaggle.com/code/jaewook704/waveform-inversion-vel-to-seis), corrected for 0-based index.","metadata":{}},{"cell_type":"code","source":"### CODE 0-BASED OK !!! <MAE> ~ 0.000012 !!!\n\n# Based on:\n# https://arxiv.org/pdf/2111.02926\n# https://csim.kaust.edu.sa/files/SeismicInversion/Chapter.FD/lab.FD2.8/lab.html\n# https://www.kaggle.com/code/jaewook704/waveform-inversion-vel-to-seis\n\ndef ricker(f, dt, nt):\n    \"\"\"\n    Generate a Ricker wavelet of central frequency `f`.\n\n    Args:\n        f (float): Dominant frequency of the wavelet (Hz).\n        dt (float): Time step (s).\n        nt (int): Total number of time samples.\n\n    Returns:\n        w (np.ndarray): The Ricker wavelet padded to length `nt`.\n        tw (np.ndarray): Corresponding time axis.\n    \"\"\"\n    nw = int(2.2 / f / dt)\n    nw = 2 * (nw // 2) + 1  # Ensure odd length for centering\n    nc = nw // 2  # Center index\n\n    k = np.arange(nw)\n    alpha = (nc - k) * f * dt * np.pi\n    beta = alpha ** 2\n    w0 = (1.0 - 2.0 * beta) * np.exp(-beta)  # Ricker formula\n\n    if nt < len(w0):\n        raise ValueError(\"nt is too small to hold the Ricker wavelet.\")\n    \n    w = np.zeros(nt)\n    w[:len(w0)] = w0  # Pad wavelet to length nt\n    tw = np.arange(len(w)) * dt\n\n    return w, tw\n\ndef padvel(v0, nbc):\n    \"\"\"\n    Pad velocity model on all sides using edge values.\n\n    Args:\n        v0 (np.ndarray): Original velocity model.\n        nbc (int): Number of boundary cells.\n\n    Returns:\n        np.ndarray: Padded velocity model.\n    \"\"\"\n    return np.pad(v0, ((nbc, nbc), (nbc, nbc)), mode='edge')\n\ndef expand_source(s0, nt):\n    \"\"\"\n    Expand source wavelet to total time length by zero-padding.\n\n    Args:\n        s0 (array): Input source wavelet.\n        nt (int): Total number of time steps.\n\n    Returns:\n        np.ndarray: Padded source wavelet.\n    \"\"\"\n    s0 = np.asarray(s0).flatten()\n    s = np.zeros(nt)\n    s[:len(s0)] = s0\n    return s\n\ndef adjust_sr(coord, dx, nbc):\n    \"\"\"\n    Convert real-world source/receiver coordinates to grid indices.\n\n    Args:\n        coord (dict): Contains keys 'sx', 'sz', 'gx', 'gz' (in meters).\n        dx (float): Grid spacing.\n        nbc (int): Number of absorbing boundary cells.\n\n    Returns:\n        Tuple of ints/arrays: isx, isz, igx, igz (grid indices).\n    \"\"\"\n    isx = int(round(coord['sx'] / dx)) + nbc\n    isz = int(round(coord['sz'] / dx)) + nbc\n    igx = (np.round(np.array(coord['gx']) / dx) + nbc).astype(int)\n    igz = (np.round(np.array(coord['gz']) / dx) + nbc).astype(int)\n\n    return isx, isz, igx, igz\n\ndef AbcCoef2D(vel, nbc, dx):\n    \"\"\"\n    Compute 2D absorbing boundary damping coefficients.\n\n    Args:\n        vel (np.ndarray): Padded velocity model.\n        nbc (int): Number of absorbing cells.\n        dx (float): Grid spacing.\n\n    Returns:\n        np.ndarray: Damping coefficient matrix.\n    \"\"\"\n    nzbc, nxbc = vel.shape\n    velmin = np.min(vel)\n    nz = nzbc - 2 * nbc\n    nx = nxbc - 2 * nbc\n\n    a = (nbc - 1) * dx\n    kappa = 3.0 * velmin * np.log(1e7) / (2.0 * a)\n\n    damp1d = kappa * ((np.arange(nbc) * dx / a) ** 2)\n    damp = np.zeros((nzbc, nxbc))\n\n    # Left/right boundaries\n    for iz in range(nzbc):\n        damp[iz, :nbc] = damp1d[::-1]\n        damp[iz, nx + nbc:] = damp1d\n\n    # Top/bottom boundaries\n    for ix in range(nbc, nbc + nx):\n        damp[:nbc, ix] = damp1d[::-1]\n        damp[nz + nbc:, ix] = damp1d\n\n    return damp\n\ndef a2d_mod_abc24(v, nbc, dx, nt, dt, s, coord, isFS):\n    \"\"\"\n    Simulate 2D acoustic wave propagation using a 2-4 finite difference scheme with ABC.\n\n    Args:\n        v (np.ndarray): Velocity model (nz, nx).\n        nbc (int): Number of absorbing boundary cells.\n        dx (float): Grid spacing (m).\n        nt (int): Number of time steps.\n        dt (float): Time step (s).\n        s (np.ndarray): Source time function.\n        coord (dict): Source and receiver positions.\n        isFS (bool): Free surface flag (unused here).\n\n    Returns:\n        np.ndarray: Seismograms at receiver positions (nt, ng).\n    \"\"\"\n    ng = len(coord['gx'])\n    seis = np.zeros((nt, ng))\n\n    # Coefficients for 2nd–4th order finite difference\n    c1, c2, c3 = -2.5, 4.0 / 3.0, -1.0 / 12.0\n\n    v = padvel(v, nbc)\n    abc = AbcCoef2D(v, nbc, dx)\n\n    alpha = (v * dt / dx) ** 2\n    kappa = abc * dt\n    temp1 = 2 + 2 * c1 * alpha - kappa\n    temp2 = 1 - kappa\n    beta_dt = (v * dt) ** 2\n    s = expand_source(s, nt)\n    isx, isz, igx, igz = adjust_sr(coord, dx, nbc)\n\n    # Time stepping variables\n    p0 = np.zeros_like(v)  # u^{n-1}\n    p1 = np.zeros_like(v)  # u^{n}\n\n    for it in range(nt):\n        # 2D finite difference update (2nd–4th order Laplacian)\n        p = (temp1 * p1 - temp2 * p0 +\n             alpha * (\n                 c2 * (np.roll(p1, 1, axis=1) + np.roll(p1, -1, axis=1) +\n                       np.roll(p1, 1, axis=0) + np.roll(p1, -1, axis=0)) +\n                 c3 * (np.roll(p1, 2, axis=1) + np.roll(p1, -2, axis=1) +\n                       np.roll(p1, 2, axis=0) + np.roll(p1, -2, axis=0))\n             ))\n\n        # Inject source\n        p[isz, isx] += beta_dt[isz, isx] * s[it]\n\n        # Record receiver data\n        for ig in range(ng):\n            seis[it, ig] = p[igz[ig], igx[ig]]\n\n        p0, p1 = p1, p  # Shift time steps\n\n    return seis\n\ndef vel_to_seis(vel):\n    \"\"\"\n    Generate seismic data from a 2D velocity model using 5 source locations.\n\n    Args:\n        vel (np.ndarray): Velocity model (70, 70).\n\n    Returns:\n        np.ndarray: Seismic data (5, 1001, 70).\n    \"\"\"\n    nz, nx = 70, 70\n    dx = 10\n    nbc = 120\n    nt = 1001\n    dt = 1e-3\n    freq = 15\n    isFS = False\n\n    s, _ = ricker(freq, dt, nt)\n\n    # Common source depth and receiver line\n    coord = {\n        'sz': 1 * dx,\n        'gx': np.arange(nx) * dx,\n        'gz': np.ones(nx) * dx\n    }\n\n    seis_data = []\n    for source_x in [0, 17, 34, 52, 69]:\n        coord['sx'] = source_x * dx\n        seis = a2d_mod_abc24(vel, nbc, dx, nt, dt, s, coord, isFS)\n        seis_data.append(seis)\n\n    return np.stack(seis_data, axis=0)\n\ndef plot_seis(seis):\n    \"\"\"\n    Plot seismograms for 5 sources.\n\n    Args:\n        seis (np.ndarray): Seismic data (5, 1000, 70).\n    \"\"\"\n    import matplotlib.pyplot as plt\n\n    fig, ax = plt.subplots(1, 5, figsize=(20, 5))\n    for i in range(5):\n        ax[i].imshow(seis[i, :, :], extent=[0, 70, 1000, 0], aspect='auto',\n                     cmap='gray', vmin=-0.5, vmax=0.5)\n        ax[i].set_xticks(range(0, 70, 10))\n        ax[i].set_xticklabels(range(0, 700, 100))\n        ax[i].set_yticks(range(0, 2000, 1000))\n        ax[i].set_yticklabels(range(0, 2, 1))\n        ax[i].set_xlabel('Offset (m)', fontsize=12)\n        ax[i].set_ylabel('Time (s)', fontsize=12)\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:43:42.802457Z","iopub.execute_input":"2025-06-23T06:43:42.803212Z","iopub.status.idle":"2025-06-23T06:43:42.820772Z","shell.execute_reply.started":"2025-06-23T06:43:42.803181Z","shell.execute_reply":"2025-06-23T06:43:42.820026Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"Z = 0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:43:48.728633Z","iopub.execute_input":"2025-06-23T06:43:48.728915Z","iopub.status.idle":"2025-06-23T06:43:48.732281Z","shell.execute_reply.started":"2025-06-23T06:43:48.728890Z","shell.execute_reply":"2025-06-23T06:43:48.731796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_vels[0].dtype","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:47:01.451523Z","iopub.execute_input":"2025-06-23T06:47:01.452055Z","iopub.status.idle":"2025-06-23T06:47:01.456327Z","shell.execute_reply.started":"2025-06-23T06:47:01.452034Z","shell.execute_reply":"2025-06-23T06:47:01.455666Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\nseis_notorch = vel_to_seis(all_vels[Z][0,0].cpu())[:,:-1,:]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:48:11.807001Z","iopub.execute_input":"2025-06-23T06:48:11.807712Z","iopub.status.idle":"2025-06-23T06:48:20.071920Z","shell.execute_reply.started":"2025-06-23T06:48:11.807687Z","shell.execute_reply":"2025-06-23T06:48:20.071338Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# no torch\nerrors = []\nfor j in range(5):\n    error = np.mean(np.abs(all_seis[Z][0, j, :, :].cpu().numpy() - seis_notorch[j, :, :])) # MAE\n    errors.append(error)\n    print(f\"Receiver{j+1} Error : {error:.6f}\")\nprint(f\"Mean Error : {np.mean(errors):.6f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:48:20.772842Z","iopub.execute_input":"2025-06-23T06:48:20.773338Z","iopub.status.idle":"2025-06-23T06:48:20.780765Z","shell.execute_reply.started":"2025-06-23T06:48:20.773319Z","shell.execute_reply":"2025-06-23T06:48:20.780053Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Torch","metadata":{}},{"cell_type":"code","source":"import torch\n\ndef ricker_t(f, dt, nt):\n    \"\"\"\n    Generate a differentiable Ricker wavelet using PyTorch.\n    Supports autograd for f and dt.\n\n    Args:\n        f (Tensor): Dominant frequency (can require grad).\n        dt (Tensor): Time step (can require grad).\n        nt (int): Number of time steps.\n\n    Returns:\n        Tuple[Tensor, Tensor]:\n            - w: Wavelet (nt,)\n            - tw: Time axis (nt,)\n    \"\"\"\n    # Approximate the number of wavelet samples based on frequency and time step\n    approx_nw = 2.2 / (f * dt)\n    nw_float = 2.0 * torch.floor(approx_nw / 2.0) + 1  # Ensure odd size\n    nw = nw_float.to(dtype=torch.int32).item()  # Convert to integer for indexing\n\n    # Center index of wavelet\n    nc = nw // 2\n    k = torch.arange(nw, dtype=torch.float32, device=f.device if isinstance(f, torch.Tensor) else None)\n\n    # Compute wavelet samples\n    alpha = (nc - k) * f * dt * torch.pi\n    beta = alpha ** 2\n    w0 = (1.0 - 2.0 * beta) * torch.exp(-beta)\n\n    if nt < nw:\n        raise ValueError(\"nt must be >= nw\")\n\n    # Zero-pad the wavelet to full length\n    w = torch.zeros(nt, dtype=w0.dtype, device=w0.device)\n    w[:nw] = w0\n    w = w.clone()  # Clone to preserve autograd behavior\n\n    # Time axis\n    tw = torch.arange(nt, dtype=torch.float32, device=w.device) * dt\n    return w, tw\n\ndef padvel_t(v0, nbc):\n    \"\"\"\n    Apply edge-replication padding to a 2D tensor.\n\n    Args:\n        v0 (Tensor): Input velocity model (H, W)\n        nbc (int): Number of boundary cells to pad.\n\n    Returns:\n        Tensor: Padded velocity model (H+2*nbc, W+2*nbc)\n    \"\"\"\n    pad = torch.nn.ReplicationPad2d(nbc)\n    v0 = v0.unsqueeze(0).unsqueeze(0)  # Add batch and channel dimensions\n    out = pad(v0)\n    return out.squeeze(0).squeeze(0)  # Remove batch and channel dimensions\n\ndef expand_source_t(s0, nt):\n    \"\"\"\n    Zero-pad a 1D source wavelet to a fixed time length.\n\n    Args:\n        s0 (array-like): Input wavelet.\n        nt (int): Target number of time steps.\n\n    Returns:\n        Tensor: Zero-padded wavelet (nt,)\n    \"\"\"\n    s0 = torch.as_tensor(s0, dtype=torch.float32).flatten()\n    ns = s0.shape[0]\n\n    if ns > nt:\n        raise ValueError(\"s0 is longer than nt\")\n\n    s = torch.zeros(nt, dtype=s0.dtype, device=s0.device)\n    s[:ns] = s0\n    return s.clone()\n\ndef AbcCoef2D_t(vel, nbc, dx):\n    \"\"\"\n    Generate absorbing boundary damping coefficients for a 2D velocity model.\n\n    Args:\n        vel (Tensor): Velocity model with padding (nzbc, nxbc)\n        nbc (int): Number of boundary cells\n        dx (float): Spatial resolution\n\n    Returns:\n        Tensor: Damping coefficients (nzbc, nxbc)\n    \"\"\"\n    nzbc, nxbc = vel.shape\n    nz = nzbc - 2 * nbc\n    nx = nxbc - 2 * nbc\n\n    # Minimum velocity used to scale damping\n    velmin = torch.min(vel)\n    a = (nbc - 1) * dx  # Effective thickness of absorbing layer\n    kappa = 3.0 * velmin * torch.log(torch.tensor(1e7, dtype=vel.dtype, device=vel.device)) / (2.0 * a)\n\n    # Damping profile from edge to interior\n    idx = torch.arange(nbc, dtype=vel.dtype, device=vel.device)\n    damp1d = kappa * ((idx * dx / a) ** 2)\n\n    # Initialize damping matrix\n    damp = torch.zeros((nzbc, nxbc), dtype=vel.dtype, device=vel.device)\n\n    # Left and right edges\n    damp[:, :nbc] = damp1d.flip(0).unsqueeze(0)\n    damp[:, nx + nbc:] = damp1d.unsqueeze(0)\n\n    # Top and bottom edges\n    damp[:nbc, nbc:nx + nbc] = damp1d.flip(0).unsqueeze(1)\n    damp[nz + nbc:, nbc:nx + nbc] = damp1d.unsqueeze(1)\n\n    return damp\n\ndef corrected_reception(p, igx, igz):\n    \"\"\"\n    Extract receiver data from pressure field using integer grid indices.\n\n    Args:\n        p (Tensor): Pressure field (B, H, W)\n        igx, igz (Tensor): Receiver grid indices (B, N)\n\n    Returns:\n        Tensor: Receiver data (B, N)\n    \"\"\"\n    B, H, W = p.shape\n    nx = igx.shape[1]\n    batch_idx = torch.arange(B, device=p.device).view(B, 1).expand(B, nx)\n    return p[batch_idx, igz.long(), igx.long()]\n\ndef prepare_geom_multi_source(source_positions, dx, nbc, nx, device):\n    \"\"\"\n    Prepare grid indices for multiple sources and receivers.\n\n    Args:\n        source_positions (list): Source x positions in index units\n        dx (float): Grid spacing\n        nbc (int): Number of absorbing cells\n        nx (int): Number of receiver x positions\n        device (str): Device to use\n\n    Returns:\n        Dict[str, Tensor]: Contains source and receiver positions and indices.\n    \"\"\"\n    dtype = torch.float32\n    ns = len(source_positions)\n\n    sx = torch.tensor(source_positions, dtype=dtype, device=device) * dx\n    sz = torch.full((ns,), dx, dtype=dtype, device=device)\n    isx = sx / dx + nbc\n    isz = sz / dx + nbc\n    isx_i = torch.round(isx).long()\n    isz_i = torch.round(isz).long()\n\n    gx = torch.arange(nx, dtype=dtype, device=device) * dx\n    gz = torch.full((nx,), dx, dtype=dtype, device=device)\n    igx = torch.round(gx / dx + nbc).long()\n    igz = torch.round(gz / dx + nbc).long()\n\n    igx_all = igx.unsqueeze(0).expand(ns, -1)\n    igz_all = igz.unsqueeze(0).expand(ns, -1)\n\n    return {\n        \"isx\": isx, \"isz\": isz,\n        \"isx_i\": isx_i, \"isz_i\": isz_i,\n        \"igx\": igx_all, \"igz\": igz_all\n    }\n\ndef a2d_mod_abc24_t_batched(vel, nbc, dx, nt, dt, s, isx_i, isz_i, igx, igz, isFS=False):\n    \"\"\"\n    Batched version of 2D wave propagation using finite difference and ABC.\n\n    Args:\n        vel (Tensor): Velocity model (nz, nx)\n        nbc (int): Padding size\n        dx, dt (float): Space and time steps\n        nt (int): Number of time steps\n        s (Tensor): Source wavelet (nt,)\n        isx_i, isz_i (Tensor): Source indices (B,)\n        igx, igz (Tensor): Receiver indices (B, nx)\n\n    Returns:\n        Tensor: Seismograms (B, nt, nx)\n    \"\"\"\n    B = isx_i.shape[0]\n    nx = igx.shape[1]\n    device = vel.device\n    dtype = vel.dtype\n\n    # Padding and damping\n    v_pad = padvel_t(vel, nbc)\n    abc = AbcCoef2D_t(v_pad, nbc, dx)\n\n    # Finite difference coefficients\n    alpha = (v_pad * dt / dx) ** 2\n    kappa = abc * dt\n    temp1 = 2 + 2 * (-2.5) * alpha - kappa\n    temp2 = 1 - kappa\n    beta_dt = (v_pad * dt) ** 2\n\n    # Expand for batch\n    v_pad = v_pad.expand(B, -1, -1)\n    alpha = alpha.expand(B, -1, -1)\n    temp1 = temp1.expand(B, -1, -1)\n    temp2 = temp2.expand(B, -1, -1)\n    beta_dt = beta_dt.expand(B, -1, -1)\n\n    p0 = torch.zeros_like(v_pad)\n    p1 = torch.zeros_like(v_pad)\n    seis = torch.zeros((B, nt, nx), dtype=dtype, device=device)\n    s = expand_source_t(s, nt)\n\n    batch_indices = torch.arange(B, device=device)\n\n    # Main time loop\n    for it in range(nt):\n        c2, c3 = 4.0 / 3.0, -1.0 / 12.0\n        p = (temp1 * p1 - temp2 * p0 +\n             alpha * (\n                 c2 * (torch.roll(p1, 1, dims=2) + torch.roll(p1, -1, dims=2) +\n                        torch.roll(p1, 1, dims=1) + torch.roll(p1, -1, dims=1)) +\n                 c3 * (torch.roll(p1, 2, dims=2) + torch.roll(p1, -2, dims=2) +\n                        torch.roll(p1, 2, dims=1) + torch.roll(p1, -2, dims=1))\n             ))\n\n        # Inject source\n        p[batch_indices, isz_i, isx_i] += beta_dt[batch_indices, isz_i, isx_i] * s[it]\n        # Record receivers\n        seis[:, it] = p[batch_indices.view(-1, 1), igz, igx]\n        p0, p1 = p1, p\n\n    return seis\n\nclass SeismicGeometry:\n    \"\"\"\n    Object-oriented wrapper for 2D seismic simulation with multiple sources.\n\n    Improvements over the basic version:\n    - Differentiable PyTorch implementation (autograd-compatible)\n    - Batched wave propagation for multiple sources\n    - Object encapsulation for reuse and clarity\n    \"\"\"\n    def __init__(\n        self,\n        source_positions=[0, 17, 34, 52, 69],\n        nbc=120,\n        nx=70,\n        dx=10.0,\n        freq=15.0,\n        dt=1e-3,\n        nt=1001,\n        isFS=False,\n        device=\"cuda\",\n        dtype=torch.float32\n    ):\n        self.device = device\n        self.dtype = dtype\n        self.source_positions = source_positions\n        self.geom = prepare_geom_multi_source(\n            source_positions, dx, nbc, nx, device\n        )\n        self.nbc = nbc\n        self.dx = torch.tensor(dx, dtype=dtype, device=device)\n        self.dt = torch.tensor(dt, dtype=dtype, device=device)\n        self.freq = torch.tensor(freq, dtype=dtype, device=device)\n        self.nt = int(nt)\n        self.s, _ = ricker_t(self.freq, self.dt, self.nt)\n        self.isFS = False\n\n    def simulate(self, vel):\n        \"\"\"\n        Run the seismic simulation.\n\n        Args:\n            vel (Tensor): Velocity model (70, 70)\n\n        Returns:\n            Tensor: Seismograms (5, 1001, 70)\n        \"\"\"\n        return a2d_mod_abc24_t_batched(\n            vel, self.nbc, self.dx, self.nt, self.dt, self.s,\n            self.geom[\"isx_i\"], self.geom[\"isz_i\"],\n            self.geom[\"igx\"], self.geom[\"igz\"],\n            self.isFS\n        )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:48:39.687525Z","iopub.execute_input":"2025-06-23T06:48:39.687791Z","iopub.status.idle":"2025-06-23T06:48:39.709846Z","shell.execute_reply.started":"2025-06-23T06:48:39.687771Z","shell.execute_reply":"2025-06-23T06:48:39.709126Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"geom = SeismicGeometry()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:48:43.294398Z","iopub.execute_input":"2025-06-23T06:48:43.294671Z","iopub.status.idle":"2025-06-23T06:48:43.520261Z","shell.execute_reply.started":"2025-06-23T06:48:43.294650Z","shell.execute_reply":"2025-06-23T06:48:43.519720Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"Z = 0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T07:13:25.575894Z","iopub.execute_input":"2025-06-23T07:13:25.576604Z","iopub.status.idle":"2025-06-23T07:13:25.580052Z","shell.execute_reply.started":"2025-06-23T07:13:25.576581Z","shell.execute_reply":"2025-06-23T07:13:25.579321Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\nseis_torch = geom.simulate(all_vels[Z][0,0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:49:44.669522Z","iopub.execute_input":"2025-06-23T06:49:44.670072Z","iopub.status.idle":"2025-06-23T06:49:45.115196Z","shell.execute_reply.started":"2025-06-23T06:49:44.670056Z","shell.execute_reply":"2025-06-23T06:49:45.114623Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"errors = []\nfor j in range(5):\n    error = np.mean(np.abs(all_seis[Z][0, j, :, :].cpu().numpy() - seis_torch[j, :-1, :].cpu().numpy())) # MAE\n    errors.append(error)\n    print(f\"Receiver{j+1} Error : {error:.6f}\")\nprint(f\"Mean Error : {np.mean(errors):.6f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T06:49:45.116248Z","iopub.execute_input":"2025-06-23T06:49:45.116895Z","iopub.status.idle":"2025-06-23T06:49:45.130287Z","shell.execute_reply.started":"2025-06-23T06:49:45.116876Z","shell.execute_reply":"2025-06-23T06:49:45.129639Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Check MAE on 10 % of all files","metadata":{}},{"cell_type":"code","source":"errors = []\nfor Z in tqdm(range(len(all_vels))):\n    for II in range(0, all_vels[Z].shape[0], 10):  # check only 10 %\n        seis_torch = geom.simulate(all_vels[Z][II,0])\n        errors.append(np.mean(np.abs(all_seis[Z][II, j, :, :].cpu().numpy() - seis_torch[j, :-1, :].cpu().numpy())))\nprint(f\"Mean Error : {np.mean(errors):.6f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T07:17:58.463027Z","iopub.execute_input":"2025-06-23T07:17:58.463557Z","execution_failed":"2025-06-23T07:25:14.698Z"}},"outputs":[],"execution_count":null}]}