{"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"},{"sourceId":11672208,"sourceType":"datasetVersion","datasetId":4402985},{"sourceId":243855450,"sourceType":"kernelVersion"}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport scipy.io\nimport matplotlib.pyplot as plt\n\nimport numpy as np\nfrom scipy.ndimage import convolve1d\n\nimport cupy as cp\nimport cupyx.scipy.ndimage as ndimage\nfrom multiprocessing import Process","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:24.137469Z","iopub.execute_input":"2025-06-08T16:58:24.137781Z","iopub.status.idle":"2025-06-08T16:58:30.064627Z","shell.execute_reply.started":"2025-06-08T16:58:24.137752Z","shell.execute_reply":"2025-06-08T16:58:30.06376Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"rng = np.random.default_rng(seed=4000+1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:30.065379Z","iopub.execute_input":"2025-06-08T16:58:30.065895Z","iopub.status.idle":"2025-06-08T16:58:30.072913Z","shell.execute_reply.started":"2025-06-08T16:58:30.065861Z","shell.execute_reply":"2025-06-08T16:58:30.07169Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_name = 1\nenv = 'Kaggle'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:30.075024Z","iopub.execute_input":"2025-06-08T16:58:30.075423Z","iopub.status.idle":"2025-06-08T16:58:30.112301Z","shell.execute_reply.started":"2025-06-08T16:58:30.075394Z","shell.execute_reply":"2025-06-08T16:58:30.110655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import multiprocessing as mp\nmp.set_start_method('fork', force=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:30.11343Z","iopub.execute_input":"2025-06-08T16:58:30.113776Z","iopub.status.idle":"2025-06-08T16:58:30.13702Z","shell.execute_reply.started":"2025-06-08T16:58:30.113733Z","shell.execute_reply":"2025-06-08T16:58:30.135813Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, gc\nimport tensorflow as tf\n\nimport pandas as pd\nimport numpy as np\nfrom sklearn.model_selection import KFold\nimport sklearn\nimport matplotlib.pyplot as plt\nimport pickle\nimport shutil\n\nimport time\n\nimport scipy.stats as stats\nimport math\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:30.138072Z","iopub.execute_input":"2025-06-08T16:58:30.138484Z","iopub.status.idle":"2025-06-08T16:58:47.63416Z","shell.execute_reply.started":"2025-06-08T16:58:30.138435Z","shell.execute_reply":"2025-06-08T16:58:47.633432Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# generate vel","metadata":{}},{"cell_type":"code","source":"# Initialize 70x70 velocity map c0 with horizontal strips\ndef initialize_c0(rng):\n    # Random number of strips (3 to 7)\n    num_strips = rng.integers(low=8, high=9)\n\n    # Random widths summing to 70\n    widths = rng.multinomial(70, np.ones(num_strips)/num_strips)\n    while any(widths < 1):\n        widths = rng.multinomial(70, np.ones(num_strips)/num_strips)\n    widths = widths.astype(int)\n\n    # Random values between 1500 and 4500, sorted ascending\n    strip_values = rng.integers(low=1500, high=4501, size=num_strips)\n    strip_values = np.sort(strip_values)  # Higher strips have higher values\n\n    # Initialize 70x70 array\n    c0 = np.zeros((70, 70))\n    current_row = 0\n    for i in range(num_strips):\n        c0[current_row:current_row + widths[i], :] = strip_values[i]\n        current_row += widths[i]\n\n    return c0, num_strips, widths, strip_values\n\n# Vel Family: ci(x, y) = ci−1(x, y + ai * sin(2π * ki * x))\ndef apply_fault_family(c_prev, rng, num_iterations):\n    c0 = c_prev.copy()\n    c_current = c_prev.copy()\n    height, width = c_current.shape\n    # Create x coordinates normalized to [0, 1]\n    x = np.linspace(0, 1, width)\n    X = np.tile(x, (height, 1))  # Shape (70, 70)\n    Y_pixels = np.tile(np.arange(height)[:, None], (1, width))  # Shape (70, 70)\n\n    masks_list = []\n    for iteration in range(num_iterations):\n        theta_deg = rng.uniform(0, 90)\n        theta_rad = np.deg2rad(theta_deg)\n        m = np.tan(theta_rad)\n        if m <= 0:\n            b = rng.uniform(0.0, 70.0-70.0*m)   # Intercept of fi(x)\n        else:\n            b = rng.uniform(-70.0*m, 70)\n        # Compute fault boundary fi(x) = m * x + b\n        fi_x = m * X*70 + b  # Shape (70, 70), in pixel space\n        mask = Y_pixels-70 >= fi_x  # Boolean mask for y >= fi(x)\n\n        masks_list.append(mask)\n\n    sorted_masks = []\n\n    for k in range(len(masks_list)):\n        mask_curr = masks_list[k].copy()\n        for j in range(len(sorted_masks)+1):\n            sorted_masks_temp = sorted_masks[:j].copy()+[mask_curr.copy()]+sorted_masks[j:].copy()\n            mask_counter = -1+np.zeros((height, width))\n            for i in range(len(sorted_masks_temp)):\n                mask_counter = np.where(sorted_masks_temp[i], mask_counter, mask_counter*0+i)\n            mask_counter = mask_counter[70:140]\n            #print(len(np.unique(mask_counter)))\n            if len(np.unique(mask_counter))>k+1:\n                break\n        sorted_masks = sorted_masks[:j].copy()+[mask_curr.copy()]+sorted_masks[j:].copy()\n    #plt.imshow(mask_counter)\n    #plt.show()\n    sorted_masks =sorted_masks[::-1]\n\n    for i in range(len(sorted_masks)):\n        # Random parameters\n        ai = rng.uniform(2.0, 10.0)  # Amplitude for row shifts\n        ki = rng.uniform(1.5, 3.5)   # Frequency for oscillations\n\n        si = rng.uniform(-30.0, 30.0)  # Fault shift\n        s_prime_i = int(rng.uniform(-20.0, 20.0))  # Additional fault shift\n\n        mask = sorted_masks[i]\n\n        # Compute source y coordinates\n        shift = ai * np.sin(2 * np.pi * ki * X) + si  # Shape (70, 70)\n        y_source = Y_pixels + shift  # Shape (70, 70)\n        y_source = np.clip(y_source, 0, height - 1)  # Handle boundaries\n        y_source = np.round(y_source).astype(int)  # Nearest integer\n\n        # Map values using vectorized indexing\n        c_new = c_current.copy()  # Start with ci−1\n        # Apply c0 where mask is True\n        shifted = np.take_along_axis(c0, y_source, axis=0)\n        shifted = np.concatenate([shifted*0+shifted[:,0:1], shifted, shifted*0+shifted[:,-1:]], axis = 1)\n        shifted = shifted[:, 70+s_prime_i:140+s_prime_i]\n        c_new[mask] = shifted[mask]\n        c_current = c_new.copy()\n\n    return c_current","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T18:10:50.072157Z","iopub.execute_input":"2025-06-08T18:10:50.072519Z","iopub.status.idle":"2025-06-08T18:10:50.084226Z","shell.execute_reply.started":"2025-06-08T18:10:50.072496Z","shell.execute_reply":"2025-06-08T18:10:50.083305Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_fault():\n    c1, num_strips, widths, strip_values = initialize_c0(rng)\n    c0 = np.min(c1)+c1.copy()*0\n    c2 = np.max(c1)+c1.copy()*0\n    c0 = np.concatenate([c0,c1,c2], axis = 0)\n    fault_family = apply_fault_family(c0, rng, num_iterations=7)\n    fault_family = fault_family[70:-70]\n    return fault_family\n\ndef get_fault_final():\n    fault_1 = get_fault()\n    fault_2 = get_fault()\n\n    height, width = fault_1.shape\n    Y_pixels = np.tile(np.arange(height)[:, None], (1, width))  # Shape (70, 70)\n\n    # Create x coordinates normalized to [0, 1]\n    x = np.linspace(0, 1, width)\n    X = np.tile(x, (height, 1))  # Shape (70, 70)\n\n    theta_deg = rng.uniform(45, 135)\n    theta_rad = np.deg2rad(theta_deg)\n    m = np.tan(theta_rad)\n\n    if m <= 0:\n        b = rng.uniform(70.0, -70.0*m)   # Intercept of fi(x)\n    else:\n        b = rng.uniform(70-70.0*m, 0)\n\n    # Compute fault boundary fi(x) = m * x + b\n    fi_x = m * X*70 + b  # Shape (70, 70), in pixel space\n    mask = Y_pixels >= fi_x  # Boolean mask for y >= fi(x)\n\n    fault = np.where(mask, fault_1, fault_2)\n    return fault\n\nfault = get_fault_final().astype(np.float32)\nplt.imshow(fault)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T18:12:14.608608Z","iopub.execute_input":"2025-06-08T18:12:14.608903Z","iopub.status.idle":"2025-06-08T18:12:14.80458Z","shell.execute_reply.started":"2025-06-08T18:12:14.608878Z","shell.execute_reply":"2025-06-08T18:12:14.803632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(np.unique(fault))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:48.026937Z","iopub.status.idle":"2025-06-08T16:58:48.027265Z","shell.execute_reply.started":"2025-06-08T16:58:48.027102Z","shell.execute_reply":"2025-06-08T16:58:48.027121Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\nc0_list = []\nfor i in range(5000):\n    if i%10000==0:\n        print(i)\n    fault_family = get_fault_final()\n    c0_list.append(fault_family)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:48.029968Z","iopub.status.idle":"2025-06-08T16:58:48.030359Z","shell.execute_reply.started":"2025-06-08T16:58:48.030185Z","shell.execute_reply":"2025-06-08T16:58:48.030201Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"vels = c0_list\nvels = np.asarray(vels).astype(np.float32)\nprint(len(vels))\nprint(np.min(vels))\nprint(np.max(vels))\nplt.imshow(vels[0])\nplt.show()\nplt.imshow(vels[1])\nplt.show()\nplt.imshow(vels[2])\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:48.033911Z","iopub.status.idle":"2025-06-08T16:58:48.034296Z","shell.execute_reply.started":"2025-06-08T16:58:48.034112Z","shell.execute_reply":"2025-06-08T16:58:48.034129Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport json\n\nsave_folder_name = f'models/model {model_name}'\n\nif env == 'Kaggle':\n    base_folder = '/kaggle/working'\n    save_folder = '/kaggle/working/save_folder'\n    try:\n        os.mkdir(save_folder)\n    except Exception as e:\n        print('exception error:')\n        print(e)\n    f = open('/kaggle/input/kaggle-json/kaggle.json')\n    kaggle_json = json.load(f)\n    KAGGLE_USERNAME = kaggle_json['username']\n    KAGGLE_KEY = kaggle_json['key']\n    os.environ[\"KAGGLE_USERNAME\"] = KAGGLE_USERNAME\n    os.environ[\"KAGGLE_KEY\"] = KAGGLE_KEY\nelif env == 'Colab':\n    from google.colab import drive\n    drive.mount('/content/drive')\n    save_folder = '/content/save_folder'\n    try:\n        os.mkdir(save_folder)\n    except Exception as e:\n        print('exception error:')\n        print(e)\n\n    f = open('/content/drive/MyDrive/kaggle/kaggle_auth/kaggle.json')\n    kaggle_json = json.load(f)\n\n    KAGGLE_USERNAME = kaggle_json['username']\n    KAGGLE_KEY = kaggle_json['key']\n\n    os.environ[\"KAGGLE_USERNAME\"] = KAGGLE_USERNAME\n    os.environ[\"KAGGLE_KEY\"] = KAGGLE_KEY","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:48.035694Z","iopub.status.idle":"2025-06-08T16:58:48.035972Z","shell.execute_reply.started":"2025-06-08T16:58:48.03585Z","shell.execute_reply":"2025-06-08T16:58:48.035863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if env == 'Kaggle':\n    !pip install Kaggle","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:48.037477Z","iopub.status.idle":"2025-06-08T16:58:48.038048Z","shell.execute_reply.started":"2025-06-08T16:58:48.037842Z","shell.execute_reply":"2025-06-08T16:58:48.037865Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# From https://www.kaggle.com/code/jaewook704/waveform-inversion-vel-to-seis ver 2\ndef ricker(f, dt, nt=None):\n    nw = int(2.2 / f / dt)\n    nw = 2 * (nw // 2) + 1\n    nc = nw // 2          \n\n    k = np.arange(1, nw + 1)\n    alpha = (nc - k + 1) * f * dt * np.pi\n    beta = alpha ** 2\n    w0 = (1.0 - 2.0 * beta) * np.exp(-beta)\n\n    if nt is not None:\n        if nt < len(w0):\n            raise ValueError(\"nt is smaller than condition!\")\n        w = np.zeros(nt)\n        w[:len(w0)] = w0\n    else:\n        w = w0\n\n    if nt is not None:\n        tw = np.arange(len(w)) * dt\n    else:\n        tw = np.arange(len(w0)) * dt\n\n    return w, tw\n\n\n# padvel function\ndef padvel(v0, nbc):\n    # First step: Pad horizontally (left and right)\n    v = np.concatenate([\n        np.tile(v0[:, :, 0:1], (1, 1, nbc)),  # Replicate first column nbc times\n        v0,                              # Original array\n        np.tile(v0[:, :, -1:], (1, 1, nbc))   # Replicate last column nbc times\n    ], axis = 2)  # Shape: (70, 310)\n    \n    # Second step: Pad vertically (top and bottom)\n    v = np.concatenate([\n        np.tile(v[:, 0:1, :], (1, nbc, 1)),  # Replicate first row nbc times\n        v,                              # Current array\n        np.tile(v[:, -1:, :], (1, nbc, 1))   # Replicate last row nbc times\n    ], axis = 1)  # Shape: (310, 310)\n    \n    return v\n\n# AbcCoef2D function\ndef AbcCoef2D(vel, nbc, dx):\n    # Get dimensions of vel\n    bs, nzbc, nxbc = vel.shape  # Shape: (310, 310)\n    \n    # Minimum velocity\n    velmin = np.min(vel, axis = (1,2))  # Scalar\n    \n    # Original grid sizes\n    nz = nzbc - 2 * nbc  # Scalar: 70\n    nx = nxbc - 2 * nbc  # Scalar: 70\n    \n    # Boundary layer width\n    a = (nbc - 1) * dx  # Scalar: 1190.0\n    \n    # Damping scaling factor\n    kappa = 3.0 * velmin * np.log(10000000.0) / (2.0 * a)  # Scalar\n    \n    # 1D damping array\n    damp1d = np.tile(kappa[:, None], (1, nbc)) * np.tile(((np.arange(0, nbc) * dx / a) ** 2)[None, ], (len(vel), 1)) # Shape: (120,)\n    \n    # Initialize 2D damping array\n    damp = np.zeros((bs, nzbc, nxbc))  # Shape: (310, 310)\n    \n    # Fill left and right boundaries (zones 1, 4, 7 and 3, 6, 9)\n    for iz in range(nzbc):\n        damp[:, iz, 0:nbc] = damp1d[:, ::-1]  # Reverse damp1d\n        damp[:, iz, nx + nbc:nx + 2 * nbc] = damp1d[:, :]\n    \n    \n    # Fill top and bottom boundaries (zones 2 and 8)\n    for ix in range(nbc, nbc + nx):\n        damp[:, 0:nbc, ix] = damp1d[:, ::-1]  # Reverse damp1d\n        damp[:, nz + nbc:nz + 2 * nbc, ix] = damp1d\n    \n    return damp\n\n# expand_source\ndef expand_source(s0, nt):\n    nt0 = len(s0)  # nt0 = 147\n    if nt0 < nt:\n        s = np.zeros((nt, 1))  # Shape: (1001, 1)\n        s[:nt0, 0] = s0  # Copy first 147 elements\n    else:\n        s = s0[:, None]  # Ensure column vector\n    return s\n\n# adjust_sr\ndef adjust_sr(coord, dx, nbc):\n    isx = int(np.round(coord['sx'] / dx)) + 1 + nbc  # 156\n    isz = int(np.round(coord['sz'] / dx)) + 1 + nbc  # 121\n    igx = (np.round(coord['gx'] / dx) + 1 + nbc).astype(int)  # [122, ..., 191]\n    igz = (np.round(coord['gz'] / dx) + 1 + nbc).astype(int)  # [122, ...]\n    if abs(coord['sz']) < 0.5:\n        isz += 1  # 122\n    igz = igz + (abs(coord['gz']) < 0.5).astype(int)  # No change\n    return isx, isz, igx, igz\n\n\n\ndef get_seis(vel, idx):\n    seis_list = []\n    sources_x = [0,17,34,52,69]\n    for source_X in sources_x:\n        # Grid parameters\n        nz = 70\n        nx = 70\n        dx = 10.0\n        nbc = 120\n        nt = 1000\n        dt = 0.001\n        freq = 15\n        isFS = False\n        \n        # Generate Ricker wavelet\n        s, _ = ricker(freq, dt)  # Ignore tw for now, as main.m doesn't use it\n        \n        # Source and receiver coordinates\n        coord = {\n            'sx': (source_X) * dx,           # Source x-coordinate\n            'sz': 1.0 * dx,            # Source z-coordinate\n            'gx': np.arange(0, nx + 0) * dx,  # Receiver x-coordinates (1:nx)*dx\n            'gz': np.ones(nx) * dx   # Receiver z-coordinates (ones(size(gx))*dx\n        }\n        \n        v = vel.copy()\n        # Translated code\n        seis = np.zeros((len(v), nt, len(coord['gx'])))  # Shape: (1001, 70)\n        ng = len(coord['gx'])                    # Scalar: 70\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  # Shape: (310, 310)\n        kappa = abc * dt           # Shape: (310, 310)\n        temp1 = 2 + 2 * c1 * alpha - kappa  # Shape: (310, 310)\n        temp2 = 1 - kappa           # Shape: (310, 310)\n        beta_dt = (v * dt) ** 2     # Shape: (310, 310)\n        s = expand_source(s, nt)    # Shape: (1001, 1)\n        isx, isz, igx, igz = adjust_sr(coord, dx, nbc)  # Scalars, (70,)\n        p1 = np.zeros_like(v)       # Shape: (310, 310)\n        p0 = np.zeros_like(v)       # Shape: (310, 310)\n        bs, nzbc, nxbc = v.shape        # Scalars: (310, 310)\n        nzp = nzbc - nbc            # Scalar: 190\n        nxp = nxbc - nbc            # Scalar: 190\n        kernel = np.array([c3, c2, 0, c2, c3])\n\n\n        \n        p1 = cp.asarray(p1, dtype=cp.float32)          # Shape: (13, 310, 310)\n        p0 = cp.asarray(p0, dtype=cp.float32)          # Shape: (13, 310, 310)\n        kernel = cp.asarray(kernel, dtype=cp.float32)  # Shape: (5,)\n        temp1 = cp.asarray(temp1, dtype=cp.float32)    # Shape: (13, 310, 310)\n        temp2 = cp.asarray(temp2, dtype=cp.float32)    # Shape: (13, 310, 310)\n        alpha = cp.asarray(alpha, dtype=cp.float32)    # Shape: (13, 310, 310)\n        beta_dt = cp.asarray(beta_dt, dtype=cp.float32) # Shape: (13, 310, 310)\n        s = cp.asarray(s, dtype=cp.float32)            # Shape: (1000, 1)\n        seis = cp.asarray(seis, dtype=cp.float32)      # Shape: (13, 1000, 70)\n    \n        # igz and igx remain NumPy arrays for indexing\n        # Precompute adjusted indices (igz - 1, igx - 1)\n        igz_m1 = igz - 1  # Shape: (70,)\n        igx_m1 = igx - 1  # Shape: (70,)\n        # Time looping\n        \n\n        '''\n        p1 = p1[0]\n        p0 = p0[0]\n        temp1 = temp1[0]\n        temp2 = temp2[0]\n        alpha = alpha[0]\n        beta_dt = beta_dt[0]\n        seis = seis[0]\n        '''\n\n        '''\n        # Time looping\n        for it in range(1, nt + 1):\n            horizontal_part = convolve1d(p1, kernel, axis=1, mode='wrap')\n            vertical_part = convolve1d(p1, kernel, axis=0, mode='wrap')\n            p = temp1 * p1 - temp2 * p0 + alpha * (horizontal_part + vertical_part)\n            p[isz - 1, isx - 1] += beta_dt[isz - 1, isx - 1] * s[it - 1, 0]\n                \n            # Free surface (skipped since isFS = False)\n            if isFS:\n                p[nbc, :] = 0.0\n                p[nbc - 1:nbc - 3:-1, :] = -p[nbc + 1:nbc + 3, :]\n            \n            # Record seismograms\n            for ig in range(ng):\n                seis[it - 1, ig] = p[igz[ig] - 1, igx[ig] - 1]\n            \n            # Update wavefields\n            p0 = p1.copy()\n            p1 = p.copy()\n    \n        seis_list.append(seis)\n    return seis_list\n        '''\n        # igz and igx remain NumPy arrays for indexing\n        # Precompute adjusted indices (igz - 1, igx - 1)\n        igz_m1 = igz - 1  # Shape: (70,)\n        igx_m1 = igx - 1  # Shape: (70,)\n        # Time looping\n        for it in range(1, nt + 1):\n            # Perform 1D convolutions along horizontal and vertical axes\n            horizontal_part = ndimage.convolve1d(p1, kernel, axis=2, mode='wrap')\n            vertical_part = ndimage.convolve1d(p1, kernel, axis=1, mode='wrap')\n            \n            # Update wavefield\n            p = temp1 * p1 - temp2 * p0 + alpha * (horizontal_part + vertical_part)\n            p[:, isz - 1, isx - 1] += beta_dt[:, isz - 1, isx - 1] * s[it - 1, 0]\n            \n            # Free surface condition (executes only if isFS is True)\n            if isFS:\n                p[nbc, :] = 0.0\n                p[nbc - 1:nbc - 3:-1, :] = -p[nbc + 1:nbc + 3, :]\n            \n            # Record seismograms using advanced indexing (replacing the inner loop)\n            seis[:, it - 1, :] = p[:, igz_m1, igx_m1]\n            \n            # Update wavefields for next iteration\n            p0 = p1.copy()\n            p1 = p.copy()\n    \n        seis_list.append(seis.get())\n    seis_list = np.asarray(seis_list)\n\n    \n    seis = tf.transpose(seis_list, [1,2,3,0])\n    seis = tf.transpose(tf.image.resize(seis[:,:999,:,:],(288,70), method='area'), [0,3,1,2])\n    seis = seis*10\n    seis = tf.cast(seis, tf.bfloat16)\n    \n    pickle.dump(vel, open(f'vel_{idx}.p', 'bw'))\n    pickle.dump(seis, open(f'seis_{idx}.p', 'bw'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:48.039425Z","iopub.status.idle":"2025-06-08T16:58:48.039725Z","shell.execute_reply.started":"2025-06-08T16:58:48.039602Z","shell.execute_reply":"2025-06-08T16:58:48.039614Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to run on each GPU\ndef run_on_gpu(vel, idx, gpu_id):\n    # Ensure each process uses the correct GPU\n    with cp.cuda.Device(gpu_id):\n        get_seis(vel, idx)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:48.040913Z","iopub.status.idle":"2025-06-08T16:58:48.041226Z","shell.execute_reply.started":"2025-06-08T16:58:48.041064Z","shell.execute_reply":"2025-06-08T16:58:48.041075Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch_size = 1000\nvels_all = vels\nprint(len(vels_all))\nprint(len(vels_all)//batch_size)\nprint(vels_all.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:48.043573Z","iopub.status.idle":"2025-06-08T16:58:48.044015Z","shell.execute_reply.started":"2025-06-08T16:58:48.043785Z","shell.execute_reply":"2025-06-08T16:58:48.043804Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\nfor idx in range(0, len(vels_all)//batch_size, 2):\n    print(idx)\n    processes = []\n    for gpu_id in range(2):  # Assuming 2 GPUs\n        vel = vels_all[(idx+gpu_id)*batch_size:(idx+gpu_id+1)*batch_size]\n        if len(vel)>0:\n            p = Process(target=run_on_gpu, args=(vel, idx+gpu_id, gpu_id))\n            processes.append(p)\n            p.start()\n\n    # Wait for all processes to complete\n    for p in processes:\n        p.join()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-08T16:58:48.045506Z","iopub.status.idle":"2025-06-08T16:58:48.045829Z","shell.execute_reply.started":"2025-06-08T16:58:48.045689Z","shell.execute_reply":"2025-06-08T16:58:48.045702Z"}},"outputs":[],"execution_count":null}]}