{"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":"none","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# DeepWave Forward Propagation Loss Function\n\nAfter I published [this notebook](https://www.kaggle.com/code/fpeccia/pytorch-forward-propagation-loss-function) trying to implement a forward propagation loss function in PyTorch, [this comment](https://www.kaggle.com/code/fpeccia/pytorch-forward-propagation-loss-function/comments#3198310) took me to the [DeepWave](https://github.com/ar4/deepwave) library (which was also mentioned [here](https://www.kaggle.com/code/tpmeli/geo-lit-review-winning-strategies-starter)). This notebook is my attempt to implement a loss function using that library. \n\nThe results are not numerically exact when compared with the Python forward propagation. Any help will be appreciated!","metadata":{}},{"cell_type":"code","source":"!pip install deepwave","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-09T13:24:43.884590Z","iopub.execute_input":"2025-05-09T13:24:43.884947Z","iopub.status.idle":"2025-05-09T13:25:01.539039Z","shell.execute_reply.started":"2025-05-09T13:24:43.884918Z","shell.execute_reply":"2025-05-09T13:25:01.537806Z"}},"outputs":[],"execution_count":null},{"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\nimport deepwave","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-09T13:25:01.540592Z","iopub.execute_input":"2025-05-09T13:25:01.540884Z","iopub.status.idle":"2025-05-09T13:25:06.108660Z","shell.execute_reply.started":"2025-05-09T13:25:01.540856Z","shell.execute_reply":"2025-05-09T13:25:06.107682Z"}},"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 padvel(v0, nbc):\n    v_padded = np.pad(v0, ((nbc, nbc), (nbc, nbc)), mode='edge')\n    nz, nx = v_padded.shape\n    v = np.zeros((nz + 1, nx + 1))\n    v[1:, 1:] = v_padded\n    return v\n\t\ndef expand_source(s0, nt):\n    s0 = np.asarray(s0).flatten()\n    s = np.zeros(nt + 1)\n    s[1:len(s0) + 1] = s0\n    return s\n\t\ndef adjust_sr(coord, dx, nbc):\n    isx = int(round(coord['sx'] / dx)) + 1 + nbc\n    isz = int(round(coord['sz'] / dx)) + 1 + nbc\n    igx = (np.round(np.array(coord['gx']) / dx) + 1 + nbc).astype(int)\n    igz = (np.round(np.array(coord['gz']) / dx) + 1 + nbc).astype(int)\n\n    if abs(coord['sz']) < 0.5:\n        isz += 1\n    igz = igz + (np.abs(np.array(coord['gz'])) < 0.5).astype(int)\n    return isx, isz, igx, igz\n\t\ndef AbcCoef2D(vel, nbc, dx):\n    nzbc, nxbc = vel.shape[1] - 1, vel.shape[0] - 1  # 실제 사이즈\n    velmin = np.min(vel[1:, 1:])\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(1, nbc + 1) - 1) * dx / a) ** 2)\n    damp = np.zeros((nzbc + 1, nxbc + 1))\n\n    for iz in range(1, nzbc + 1):\n        damp[iz, 1:nbc + 1] = damp1d[::-1]\n        damp[iz, nx + nbc + 1 : nx + 2 * nbc + 1] = damp1d\n\n    for ix in range(nbc + 1, nbc + nx + 1):\n        damp[1:nbc + 1, ix] = damp1d[::-1]\n        damp[nz + nbc + 1 : nz + 2 * nbc + 1, ix] = damp1d\n\n    return damp\n\t\ndef a2d_mod_abc24(v, nbc, dx, nt, dt, s, coord, isFS):\n    ng = len(coord['gx'])\n    seis = np.zeros((nt + 1, ng))  # 1-based time axis\n\n    c1 = -2.5\n    c2 = 4.0 / 3.0\n    c3 = -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    p0 = np.zeros_like(v)\n    p1 = np.zeros_like(v)\n\n    # Time Loop (1-based)\n    for it in range(1, nt + 1):\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        # Source\n        p[isz, isx] += beta_dt[isz, isx] * s[it]\n\n        # Free Surface\n        if isFS:\n            p[nbc, :] = 0.0\n            p[nbc - 1 : nbc + 1, :] = -p[nbc + 1 : nbc + 3, :]\n\n        for ig in range(ng):\n            seis[it, ig] = p[igz[ig], igx[ig]]\n\n        p0, p1 = p1, p\n\n    return seis\n    \ndef vel_to_seis(vel):\n    \"\"\"\n    vel : (70, 70)\n    output : (5, 1001, 70)\n    \"\"\"\n    # 1. 모델 및 파라미터 설정\n    nz = 70\n    nx = 70\n    dx = 10\n    nbc = 120\n    nt = 1001\n    dt = 1e-3\n    freq = 15\n    isFS = False  # 자유표면 사용 여부\n    \n    # 2. Ricker 파형 생성\n    s, _ = ricker(freq, dt)\n    \n    # 3. 소스 및 수신기 설정\n    coord = {}\n    coord['sz'] = 1 * dx\n    coord['gx'] = np.arange(1, nx + 1) * dx\n    coord['gz'] = np.ones_like(coord['gx']) * dx\n    \n    # 4. 파동장 시뮬레이션 수행 : 소스 위치만 바꿔서\n    seis_data = []\n    for source_x in [1, 18, 35, 53, 70]:\n        coord['sx'] = source_x * dx        \n        \n        # 시뮬레이션\n        seis = a2d_mod_abc24(vel, nbc, dx, nt, dt, s, coord, isFS)\n\n        seis_data += [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-05-09T13:25:06.110535Z","iopub.execute_input":"2025-05-09T13:25:06.111022Z","iopub.status.idle":"2025-05-09T13:25:06.141103Z","shell.execute_reply.started":"2025-05-09T13:25:06.110963Z","shell.execute_reply":"2025-05-09T13:25:06.139597Z"}},"outputs":[],"execution_count":null},{"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, dx=10, dt=1e-3, nt=1001, nx=70, debug=False):\n        super(WavePropagationLoss, self).__init__()\n        self.dx = dx\n        self.dt = dt\n        self.nt = nt\n        self.nx = nx\n        self.debug = debug\n\n        rickers = np.zeros((5, 1, nt + 1), dtype=np.float32)\n        for i in range(5):\n            rickers[i,0,:148] = deepwave.wavelets.ricker(freq, 148, dt, 74*dt)\n\n        self.rickers = torch.from_numpy(rickers)\n        \n        receiver_locations = np.zeros((5,nx,2))\n        for source in range(5):\n            for i in range(nx):\n                receiver_locations[source,i,:] = [0,i]\n\n        self.receiver_locations = torch.from_numpy(receiver_locations)\n        \n        source_locations = np.zeros((5,1,2))\n        source_locations[0,0,:] = [0,0]\n        source_locations[1,0,:] = [0,17]\n        source_locations[2,0,:] = [0,34]\n        source_locations[3,0,:] = [0,52]\n        source_locations[4,0,:] = [0,69]\n\n        self.source_locations = torch.from_numpy(source_locations)\n\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\n        if self.debug:\n            torch_seis_data = []\n            \n        losses = []\n        \n        for batch in range(batch_size):\n            target_seis = target_seis_batch[batch]\n\n            pred_seis = -deepwave.scalar(\n                v_batch[batch,0], grid_spacing=self.dx, dt=1e-3,\n                source_amplitudes=self.rickers,\n                source_locations=self.source_locations,\n                receiver_locations=self.receiver_locations\n            )[-1].transpose(2,1)\n\n            if self.debug:\n                torch_seis_data.append(pred_seis[:, 2:,:])\n            losses.append(F.l1_loss(pred_seis[:, 2:,:], target_seis))\n        loss = torch.stack(losses).mean()\n        if self.debug:\n            return loss, torch.cat(torch_seis_data)\n        else:\n            return loss\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-09T13:25:06.142523Z","iopub.execute_input":"2025-05-09T13:25:06.143103Z","iopub.status.idle":"2025-05-09T13:25:06.168209Z","shell.execute_reply.started":"2025-05-09T13:25:06.143068Z","shell.execute_reply":"2025-05-09T13:25:06.167030Z"}},"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-05-09T13:25:06.169335Z","iopub.execute_input":"2025-05-09T13:25:06.169719Z","iopub.status.idle":"2025-05-09T13:25:06.190360Z","shell.execute_reply.started":"2025-05-09T13:25:06.169689Z","shell.execute_reply":"2025-05-09T13:25:06.189286Z"},"_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[:, 2:, :])\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, 2:, :])) # MAE\n        errors.append(error)\n        print(f\"Receiver{j+1} Error : {error:.6f}\")\n\n    loss_fn = WavePropagationLoss(debug=True)\n    vel_data_torch = torch.from_numpy(vel_data[:1,:,:,:]).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< DEEPWAVE 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 DeepWave 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-05-09T13:25:06.191509Z","iopub.execute_input":"2025-05-09T13:25:06.191815Z","iopub.status.idle":"2025-05-09T13:27:47.387652Z","shell.execute_reply.started":"2025-05-09T13:25:06.191790Z","shell.execute_reply":"2025-05-09T13:27:47.386471Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null}]}