{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# PyTorch Forward Propagation Loss Function\n\nThis notebook implements a PyTorch version of the forward propagation presented [here](https://www.kaggle.com/code/jaewook704/waveform-inversion-vel-to-seis). My goal was to use it as a loss function for a PINN, but it is too slow to be actually usable. Any contribution is welcome!","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn.functional as F\nfrom torch.autograd import Variable\nfrom math import exp","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-12T13:56:23.508426Z","iopub.execute_input":"2025-06-12T13:56:23.508932Z","iopub.status.idle":"2025-06-12T13:56:26.905516Z","shell.execute_reply.started":"2025-06-12T13:56:23.508898Z","shell.execute_reply":"2025-06-12T13:56:26.904644Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## MATLAB code converted to Python","metadata":{"execution":{"iopub.status.busy":"2025-04-16T13:47:55.989917Z","iopub.execute_input":"2025-04-16T13:47:55.990275Z","iopub.status.idle":"2025-04-16T13:47:55.996695Z","shell.execute_reply.started":"2025-04-16T13:47:55.99025Z","shell.execute_reply":"2025-04-16T13:47:55.995305Z"}}},{"cell_type":"code","source":"# https://arxiv.org/pdf/2111.02926\n# https://csim.kaust.edu.sa/files/SeismicInversion/Chapter.FD/lab.FD2.8/lab.html\n\ndef ricker(f, dt, nt=None):\n    nw = int(2.2 / f / dt)\n    nw = 2 * (nw // 2) + 1\n    nc = nw // 2 + 1  # 중심 인덱스를 1-based 기준으로 설정\n\n    k = np.arange(1, nw + 1)  # 1-based index\n    alpha = (nc - k) * f * dt * np.pi\n    beta = alpha ** 2\n    w0 = (1.0 - 2.0 * beta) * np.exp(-beta)\n\n    # 1-based wavelet 생성\n    if nt is not None:\n        if nt < len(w0):\n            raise ValueError(\"nt is smaller than condition!\")\n        w = np.zeros(nt + 1)  # dummy 포함\n        w[1:len(w0) + 1] = w0\n    else:\n        w = np.zeros(len(w0) + 1)\n        w[1:] = w0\n\n    # 1-based time axis 생성\n    if nt is not None:\n        tw = np.arange(1, len(w)) * dt\n    else:\n        tw = np.arange(1, len(w)) * dt\n\n    return w, tw\n\ndef AbcCoef2D(vel, nbc, dx):\n    \"\"\"\n    Calculates coefficients for a 2D Absorbing Boundary Condition (ABC).\n    This is a Python/NumPy translation of the provided MATLAB function.\n\n    Args:\n        vel (np.ndarray): The padded 2D velocity model.\n        nbc (int): The number of padding cells (boundary width).\n        dx (float): The spatial grid interval.\n\n    Returns:\n        np.ndarray: A 2D array of damping coefficients.\n    \"\"\"\n    nzbc, nxbc = vel.shape\n    velmin = np.min(vel)\n    nz = nzbc - 2 * nbc\n    nx = nxbc - 2 * nbc\n\n    if nbc <= 1:\n        return np.zeros_like(vel)\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), dtype=np.float64)\n\n    # Fill left and right damping zones\n    for iz in range(nzbc):\n        damp[iz, 0:nbc] = damp1d[::-1]\n        damp[iz, nx + nbc : nx + 2 * nbc] = damp1d\n\n    # Fill top and bottom damping zones\n    for ix in range(nbc, nbc + nx):\n        damp[0:nbc, ix] = damp1d[::-1]\n        damp[nbc + nz : nz + 2 * nbc, ix] = damp1d\n        \n    return damp\n\ndef padvel(v0, nbc):\n    \"\"\"\n    Pads the velocity model by extending the edge values outward.\n    \"\"\"\n    return np.pad(v0, pad_width=nbc, mode='edge')\n\ndef expand_source(s0, nt):\n    \"\"\"\n    Ensures the source time function has length 'nt'.\n    \"\"\"\n    nt0 = s0.size\n    if nt0 < nt:\n        s = np.zeros(nt, dtype=np.float64)\n        s[:nt0] = s0\n        return s\n    else:\n        return s0[:nt].astype(np.float64)\n\ndef adjust_sr(coord, dx, nbc):\n    \"\"\"\n    Converts physical source/receiver coordinates to grid indices.\n    \"\"\"\n    # MATLAB's round(x.5) rounds away from zero. NumPy's np.round(x.5)\n    # rounds to the nearest even integer. Using np.floor(x + 0.5) for\n    # positive numbers emulates MATLAB's behavior.\n    round_to_int = lambda x: np.floor(x + 0.5).astype(int)\n\n    isx = round_to_int(coord['sx'] / dx) + nbc\n    isz = round_to_int(coord['sz'] / dx) + nbc\n    igx = round_to_int(coord['gx'] / dx) + nbc\n    igz = round_to_int(coord['gz'] / dx) + nbc\n\n    if np.abs(coord['sz']) < 0.5:\n        isz += 1\n        \n    igz += (np.abs(coord['gz']) < 0.5).astype(int)\n    \n    return isx, isz, igx, igz\n\ndef a2d_mod_abc24(v, nbc, dx, nt, dt, s, coord, isFS):\n    \"\"\"\n    Performs a 2D acoustic wave finite-difference simulation (4th order).\n    \"\"\"\n    n_receivers = coord['gx'].size\n    seis = np.zeros((nt, n_receivers), dtype=np.float64)\n    \n    c1 = -2.5\n    c2 = 4.0 / 3.0\n    c3 = -1.0 / 12.0\n    \n    v_padded = padvel(v, nbc)\n    abc = AbcCoef2D(v_padded, nbc, dx)\n    \n    alpha = (v_padded * dt / dx)**2\n    kappa = abc * dt\n    temp1 = 2 + 2 * c1 * alpha - kappa\n    temp2 = 1 - kappa\n    beta_dt = (v_padded * dt)**2\n    \n    s = expand_source(s, nt)\n    isx, isz, igx, igz = adjust_sr(coord, dx, nbc)\n\n    p0 = np.zeros_like(v_padded, dtype=np.float64)\n    p1 = np.zeros_like(v_padded, dtype=np.float64)\n    \n    for it in range(nt):\n        laplacian = (\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        p = temp1 * p1 - temp2 * p0 + alpha * laplacian\n        p[isz, isx] += beta_dt[isz, isx] * s[it]\n        \n        if isFS:\n            p[nbc, :] = 0.0\n            p[nbc-1, :] = -p[nbc+1, :]\n            p[nbc-2, :] = -p[nbc+2, :]\n\n        seis[it, :] = p[igz, igx]\n        \n        p0 = p1.copy()\n        p1 = p.copy()\n        \n    return seis\n    \ndef vel_to_seis(vel, method='abc24'):\n    \"\"\"\n    Runs the simulation for multiple sources and collects the seismograms.\n    \n    Args:\n        vel (np.ndarray): The (70, 70) velocity model.\n        method (str): The simulation function to use ('abc24').\n        \n    Returns:\n        np.ndarray: Stacked seismogram data of shape (5, 1001, 70).\n    \"\"\"\n    # 1. Model and Simulation Parameters\n    nz, nx = vel.shape\n    dx = 10.0\n    nbc = 120\n    nt = 1001\n    dt = 1e-3\n    freq = 15.0\n    isFS = False  # Use free surface condition or not\n\n    # 2. Generate Ricker wavelet source\n    s, _ = ricker(freq, dt)\n    \n    # 3. Setup Receiver Coordinates\n    # Receivers are placed at every grid point horizontally at a fixed depth.\n    coord = {}\n    coord['sz'] = 1 * dx\n    coord['gx'] = np.arange(nx) * dx\n    coord['gz'] = np.ones(nx) * dx\n    \n    # 4. Loop over source positions and run simulation\n    seis_data = []\n    source_x_locations = [0, 17, 34, 52, 69] # Using 0-based indices now\n    \n    for sx_idx in source_x_locations:\n        coord['sx'] = sx_idx * dx\n        \n        if method == 'abc24':\n            seis = a2d_mod_abc24(vel, nbc, dx, nt, dt, s, coord, isFS)\n        else:\n            raise ValueError(f\"Invalid method: {method}\")\n\n        seis_data.append(seis)\n        \n    return np.stack(seis_data, axis=0)\n    \n    \ndef plot_seis(seis):\n    \"\"\"\n    seis : (5, 1000, 70)\n    \"\"\"\n    fig,ax=plt.subplots(1,5,figsize=(20,5))\n    timesteps = seis.shape[1]\n    ax[0].imshow(seis[0, :, :],extent=[0,70,timesteps,0],aspect='auto',cmap='gray',vmin=-0.5,vmax=0.5)\n    ax[1].imshow(seis[1, :, :],extent=[0,70,timesteps,0],aspect='auto',cmap='gray',vmin=-0.5,vmax=0.5)\n    ax[2].imshow(seis[2, :, :],extent=[0,70,timesteps,0],aspect='auto',cmap='gray',vmin=-0.5,vmax=0.5)\n    ax[3].imshow(seis[3, :, :],extent=[0,70,timesteps,0],aspect='auto',cmap='gray',vmin=-0.5,vmax=0.5)\n    ax[4].imshow(seis[4, :, :],extent=[0,70,timesteps,0],aspect='auto',cmap='gray',vmin=-0.5,vmax=0.5)\n    for axis in ax:\n        axis.set_xticks(range(0, 70, 10))\n        axis.set_xticklabels(range(0, 700, 100))\n        axis.set_yticks(range(0, timesteps*2, timesteps))\n        axis.set_yticklabels(range(0, timesteps*2, timesteps))\n        axis.set_ylabel('Time (s)' if timesteps > 500 else 'Time (ms)', fontsize=12)\n        axis.set_xlabel('Offset (m)', fontsize=12)\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-12T13:56:26.906754Z","iopub.execute_input":"2025-06-12T13:56:26.907151Z","iopub.status.idle":"2025-06-12T13:56:26.928021Z","shell.execute_reply.started":"2025-06-12T13:56:26.907127Z","shell.execute_reply":"2025-06-12T13:56:26.927144Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## PyTorch forward loss","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass WavePropagationLoss(nn.Module):\n    def __init__(self, freq=15, nbc=120, dx=10, dt=1e-3, nt=1001, nx=70, isFS=False):\n        super(WavePropagationLoss, self).__init__()\n        self.nbc = nbc\n        self.dx = dx\n        self.dt = dt\n        self.nt = nt\n        self.nx = nx\n        self.isFS = isFS\n        self.s0 = torch.from_numpy(self.ricker(freq, dt))\n\n    def ricker(self, f, dt):\n        nw = int(2.2 / f / dt)\n        nw = 2 * (nw // 2) + 1\n        nc = nw // 2 + 1  # 중심 인덱스를 1-based 기준으로 설정\n    \n        k = np.arange(1, nw + 1)  # 1-based index\n        alpha = (nc - k) * f * dt * np.pi\n        beta = alpha ** 2\n        w0 = (1.0 - 2.0 * beta) * np.exp(-beta)\n    \n        # 1-based wavelet 생성\n        w = np.zeros(len(w0) + 1)\n        w[1:] = w0\n    \n        return w\n    \n    def forward(self, v_batch, target_seis_batch):\n        \"\"\"\n        v: torch.Tensor [batch, nz, nx]  (velocity model)\n        target_seis: torch.Tensor [batch, 5, nt+1, n_receivers] (observed seismograms)\n\n        Returns:\n            loss: scalar tensor\n        \"\"\"\n        batch_size = v_batch.shape[0]\n        source_x_idxs = [0, 17, 34, 52, 69]\n        losses = []\n\n        s = self.s0.to(v_batch.device).unsqueeze(0).repeat(batch_size,1)\n\n        torch_seis_data = []\n        for j in range(5):\n            target_seis = target_seis_batch[:,j]\n\n            coord = {}\n            coord['sz'] = 1 * self.dx\n            coord['gx'] = np.arange(0, self.nx) * self.dx\n            coord['gz'] = np.ones_like(coord['gx']) * self.dx\n            coord['sx'] = source_x_idxs[j] * self.dx \n                \n            pred_seis = self.simulate_batch(v_batch, s, coord)\n            torch_seis_data.append(pred_seis[:, 1:,:])\n            losses.append(F.l1_loss(pred_seis[:, 1:,:], target_seis[:,:,:]))\n        loss = torch.stack(losses).mean()\n        return loss, torch.cat(torch_seis_data)\n    \n    def simulate_batch(self, v_batch, s_batch, coord):\n        batch_size, nz, nx = v_batch.shape\n        ng = len(coord['gx'])\n        device = v_batch.device\n\n        # Prepare constants\n        c1 = -2.5\n        c2 = 4.0 / 3.0\n        c3 = -1.0 / 12.0\n        #c1 = -205.0/72.0 \n        #c2 = 8.0/5.0\n        #c3 = -1.0/5.0\n        #c4 = 8.0/315.0\n        #c5 = -1.0/560.0;\n\n        # Pad velocity\n        v_batch = self.padvel(v_batch)  # [batch_size, nz_p, nx_p]\n        abc_batch = self.AbcCoef2D(v_batch)\n\n        alpha = (v_batch * self.dt / self.dx) ** 2\n        kappa = abc_batch * self.dt\n        temp1 = 2 + 2 * c1 * alpha - kappa\n        temp2 = 1 - kappa\n        beta_dt = (v_batch * self.dt) ** 2\n\n        s_batch = self.expand_source_batch(s_batch)  # [batch_size, nt+1]\n        isx, isz, igx, igz = self.adjust_sr(device, coord)\n\n        nz_p, nx_p = v_batch.shape[-2], v_batch.shape[-1]\n\n        p0 = torch.zeros_like(v_batch)\n        p1 = torch.zeros_like(v_batch)\n        seis_batch = torch.zeros((batch_size, self.nt, ng), device=device)\n\n        # Build Laplacian kernel for conv2d: shape (1, 1, 3, 3)\n        #laplace_kernel = torch.tensor([[0, c2, 0],\n        #                               [c2, c1 * 4, c2],\n        #                               [0, c2, 0]], device=device).view(1, 1, 3, 3)\n    \n        # Expand to batch: group conv\n        #laplace_kernel = laplace_kernel.repeat(batch_size, 1, 1, 1)  # (B, 1, 3, 3)\n    \n        # Reshape for conv2d: (B, 1, Z, X)\n        #def laplacian(u):\n        #    return F.conv2d(u.unsqueeze(1), laplace_kernel, padding=1, groups=batch_size).squeeze(1)\n            \n        for it in range(self.nt):\n            # This for loop is the bottleneck.\n            # The laplacian calculated using convolutions is actually quite slow, so this idea was discarded. \n            #lap_u = self.laplacian_9pt(p1)\n            #p = temp1 * p1 - temp2 * p0 + alpha * lap_u\n\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                     #c4 * (torch.roll(p1, 3, dims=2) + torch.roll(p1, -3, dims=2) +\n                     #      torch.roll(p1, 3, dims=1) + torch.roll(p1, -3, dims=1)) +\n                     #c5 * (torch.roll(p1, 4, dims=2) + torch.roll(p1, -4, dims=2) +\n                     #      torch.roll(p1, 4, dims=1) + torch.roll(p1, -4, dims=1))\n                 ))\n\n            # Source injection (vectorized)\n            p[torch.arange(batch_size), isz, isx] += beta_dt[torch.arange(batch_size), isz, isx] * s_batch[:, it]\n\n            if self.isFS:\n                p[:, self.nbc, :] = 0.0\n                p[:, self.nbc-1:self.nbc+1, :] = -p[:, self.nbc+1:self.nbc+3, :]\n\n            # Receiver sampling (vectorized)\n            #print(torch.max(p[torch.arange(batch_size).unsqueeze(1), igz.unsqueeze(0), igx.unsqueeze(0)].view(-1)))\n            #seis_batch[:, it, :] = p[torch.arange(batch_size).unsqueeze(1), igz.unsqueeze(0), igx.unsqueeze(0)]\n            for ig in range(ng):\n                seis_batch[:, it, ig] = p[torch.arange(batch_size).unsqueeze(1), igz[ig], igx[ig]]\n            # Record receivers: vectorized gather\n            #batch_idx = torch.arange(batch_size, device=device).unsqueeze(1).expand(-1, G)  # (B, G)\n            #seis[:, it, :] = p[batch_idx, igz, igx]\n\n            p0, p1 = p1, p\n\n        return seis_batch\n\n    #def padvel(self, v0):\n    #    v_padded = torch.squeeze(F.pad(torch.unsqueeze(v0,0), (self.nbc, self.nbc, self.nbc, self.nbc), mode='replicate'))\n    #    nz, nx = v_padded.shape\n    #    v = torch.zeros((nz + 1, nx + 1), device=v0.device, dtype=v0.dtype)\n    #    v[1:, 1:] = v_padded\n    #    return v\n\n    def padvel(self, v_batch):\n        # Pad with replicate mode\n        v_padded = F.pad(v_batch, (self.nbc, self.nbc, self.nbc, self.nbc), mode='replicate')\n        #batch_size, nz_p, nx_p = v_padded.shape\n        #v = torch.zeros((batch_size, nz_p + 1, nx_p + 1), device=v_batch.device, dtype=v_batch.dtype)\n        #v[:, 1:, 1:] = v_padded\n        return v_padded\n\n    #def expand_source(self, s0):\n    #    s0 = torch.as_tensor(s0, dtype=torch.float32, device=s0.device if isinstance(s0, torch.Tensor) else 'cpu').flatten()\n    #    s = torch.zeros(self.nt + 1, device=s0.device, dtype=s0.dtype)\n    #    s[1:len(s0) + 1] = s0\n    #    return s\n\n    def expand_source_batch(self, s_batch):\n        batch_size = s_batch.shape[0]\n        device = s_batch.device\n        s = torch.zeros((batch_size, self.nt), device=device)\n        for b in range(batch_size):\n            ns = s_batch[b].numel()\n            s[b, :ns] = s_batch[b].flatten()\n        return s\n\n    def adjust_sr(self, device, coord):\n        def round_away_from_zero(x):\n            return torch.sign(x) * torch.floor(torch.abs(x) + 0.5)\n    \n        #device = self.dx.device if isinstance(self.dx, torch.Tensor) else 'cpu'\n        sx = torch.tensor(coord['sx'], dtype=torch.float32, device=device)\n        sz = torch.tensor(coord['sz'], dtype=torch.float32, device=device)\n        gx = torch.tensor(coord['gx'], dtype=torch.float32, device=device)\n        gz = torch.tensor(coord['gz'], dtype=torch.float32, device=device)\n\n        isx = round_away_from_zero(sx / self.dx).int() + self.nbc\n        isz = round_away_from_zero(sz / self.dx).int() + self.nbc\n        igx = (round_away_from_zero(gx / self.dx) + self.nbc).int()\n        igz = (round_away_from_zero(gz / self.dx) + self.nbc).int()\n\n        if torch.abs(sz) < 0.5:\n            isz += 1\n        igz += (torch.abs(gz) < 0.5).int()\n\n        return isx, isz, igx, igz\n\n    def AbcCoef2D(self, vel_batch):\n        batch_size = vel_batch.shape[0]\n        nzbc, nxbc = vel_batch.shape[2], vel_batch.shape[1]\n        #velmin = torch.min(vel_batch[:, 1:, 1:], dim=(1,2)).values\n        #velmin = vel_batch[:, 1:, 1:].view(-1,nzbc*nxbc).min(dim=1).values\n        velmin = vel_batch.reshape(batch_size, -1).min(dim=1).values\n        nz = nzbc - 2 * self.nbc\n        nx = nxbc - 2 * self.nbc\n\n        a = (self.nbc - 1) * self.dx\n        kappa = 3.0 * velmin * torch.log(torch.tensor(1e7, device=vel_batch.device)) / (2.0 * a)\n\n        damp1d = []\n        for k in kappa:\n            d = k * ((torch.arange(0, self.nbc, device=vel_batch.device) * self.dx / a) ** 2)\n            damp1d.append(d)\n        damp1d = torch.stack(damp1d)  # [batch_size, nbc]\n        \n        #damp1d = kappa * (((torch.arange(1, self.nbc + 1, device=vel.device, dtype=vel.dtype) - 1) * self.dx / a) ** 2)\n        damp = torch.zeros((batch_size, nzbc, nxbc), device=vel_batch.device, dtype=vel_batch.dtype)\n\n        for iz in range(nzbc):\n            damp[:, iz, :self.nbc] = torch.flip(damp1d, dims=[1])\n            damp[:, iz, nx + self.nbc : nx + 2 * self.nbc] = damp1d\n\n        for ix in range(self.nbc, self.nbc + nx):\n            damp[:, :self.nbc, ix] = torch.flip(damp1d, dims=[1])\n            damp[:, nz + self.nbc : nz + 2 * self.nbc, ix] = damp1d\n\n        return damp","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-12T13:56:26.929507Z","iopub.execute_input":"2025-06-12T13:56:26.929766Z","iopub.status.idle":"2025-06-12T13:56:26.952799Z","shell.execute_reply.started":"2025-06-12T13:56:26.929745Z","shell.execute_reply":"2025-06-12T13:56:26.952028Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Comparison","metadata":{}},{"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-12T13:56:26.953477Z","iopub.execute_input":"2025-06-12T13:56:26.953692Z","iopub.status.idle":"2025-06-12T13:56:26.971695Z","shell.execute_reply.started":"2025-06-12T13:56:26.953674Z","shell.execute_reply":"2025-06-12T13:56:26.970911Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\nfor i, tp in enumerate(types):\n    seis_data = np.load(seis_files[i])\n    vel_data = np.load(vel_files[i])\n    seis_data_sim = vel_to_seis(vel_data[0][0])\n    print(f\"TYPE : {tp}\")\n    print(\"\\n< VELOCITY MAP >\")\n    fig, ax = plt.subplots(1, 1, figsize=(5, 2.5))\n    ax.imshow(vel_data[0, 0])\n    plt.show()\n\n    print(f\"\\n< FORWARD SIMULATION (from velocity map) : shape={seis_data_sim.shape} >\")\n    plot_seis(seis_data_sim[:, 1:, :])\n    \n    print(\"\\n< INPUT (origin) >\")\n    plot_seis(seis_data[0])\n    \n    print(\"\\n< ERROR >\")\n    errors = []\n    for j in range(5):\n        error = np.mean(np.abs(seis_data[0, j, :, :] - seis_data_sim[j, 1:, :])) # MAE\n        errors.append(error)\n        print(f\"Receiver{j+1} Error : {error:.6f}\")\n\n    loss_fn = WavePropagationLoss()#, timesteps_to_compare=100)\n    vel_data_torch = torch.from_numpy(vel_data[0]).to('cpu')\n    seis_data_torch = torch.from_numpy(seis_data[:1,:,:,:]).to('cpu')\n    torch_error, torch_seis_data = loss_fn(vel_data_torch, seis_data_torch)\n\n    print(f\"\\n< TORCH FORWARD SIMULATION (from velocity map) : shape={torch_seis_data.shape} >\")\n    plot_seis(torch_seis_data)\n\n    print(f\"Mean Error : {np.mean(errors):.6f} - Mean Torch Error : {torch_error:.6f}\")\n    print(\"#########################################\")\n    print()\n\n    # if i==2:\n    #     break\n    # break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-12T13:56:26.972528Z","iopub.execute_input":"2025-06-12T13:56:26.972824Z","execution_failed":"2025-06-12T13:57:07.599Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null}]}