{"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"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# -----------------------------------------------------------------\n# patch_extractor.py  ––  Step-1 utilities (Phase-aware features)\n# -----------------------------------------------------------------\nimport numpy as np\nfrom numpy.lib.stride_tricks import sliding_window_view\n\n# ---------- hyper-parameters ----------\nN_FFT_T   = 32              # time window (samples)\nN_FFT_X   = 16              # receiver window (stations)\nTIME_STR  = 16              # time stride (50 % overlap)\nREC_STR   = 8               # receiver stride (50 % overlap)\n\ndef make_hann(window_len):\n    \"\"\"1-D Hann window (float32).\"\"\"\n    return 0.5 - 0.5 * np.cos(2*np.pi*np.arange(window_len) / (window_len-1))\n\nhann_t = make_hann(N_FFT_T).astype(np.float32)          # (32,)\nhann_x = make_hann(N_FFT_X).astype(np.float32)          # (16,)\ntaper  = hann_t[:,None] * hann_x[None,:]                # (32,16)\n\n# ---------- feature routine ----------\ndef fft_feature(patch):\n    \"\"\"\n    patch : (32, 16) float32\n    returns: flattened [log|Z|, sin(phi), cos(phi)]  vector\n    \"\"\"\n    Z = np.fft.rfft2(patch * taper, axes=(0,1))      # 2-D FFT, keep positive freq\n    A = np.log(np.abs(Z) + 1e-6, dtype=np.float32)\n    P = np.angle(Z, deg=False).astype(np.float32)\n    S = np.sin(P, dtype=np.float32)\n    C = np.cos(P, dtype=np.float32)\n    feat = np.concatenate([A.ravel(), S.ravel(), C.ravel()], dtype=np.float32)\n    return feat   # shape = 3 * (N_FFT_T/2+1) * N_FFT_X\n\ndef gather_to_patchbank(gather):\n    \"\"\"\n    gather : (5, 1000, 70)  float32   (channels, time, receiver)\n    returns: [num_patches, feat_dim]  float32\n    \"\"\"\n    # collapse shots into channels → (time, receiver, channel)\n    g = np.transpose(gather, (1, 2, 0))                  # (1000,70,5)\n    patch_view = sliding_window_view(\n    g, window_shape=(N_FFT_T, N_FFT_X, 5)\n)[::TIME_STR, ::REC_STR, 0]                      # -> (Nt, Nx, 32,16,5)\n    Nt, Nx, *_ = patch_view.shape\n    patches = patch_view.reshape(-1, N_FFT_T, N_FFT_X, 5)\n    # treat channels independently –– flatten channel dim into time\n    patches = patches.transpose(0,3,1,2).reshape(-1, N_FFT_T, N_FFT_X)\n    feats = np.stack([fft_feature(patch) for patch in patches], axis=0)\n    return feats   # (num_patches*5, feat_dim)\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------------------------------------------\n# sample and whiten\n# ---------------------------------------------------------------\nimport numpy as np\nfeat_list = []\n\n# 1. load one FlatVel-A batch   (500 examples)\nseis_batch = np.load(\"/kaggle/input/waveform-inversion/train_samples/FlatVel_A/data/data1.npy\", mmap_mode='r')  # (500,5,1000,70)\n\n# 2. loop through first 250 examples (for speed)\nfor g in seis_batch[:250]:\n    feats = gather_to_patchbank(g.astype(np.float32))\n    # subsample to keep memory ≤ 100k\n    if feats.shape[0] > 400:\n        idx = np.random.choice(feats.shape[0], 400, replace=False)\n        feats = feats[idx]\n    feat_list.append(feats)\n\nX = np.concatenate(feat_list, axis=0)        # ≈ 100 000 × 816\nprint(\"Patch-bank shape:\", X.shape)\n\n# 3. whiten\nmu  = X.mean(axis=0, keepdims=True)\nstd = X.std(axis=0, keepdims=True) + 1e-6\nXw  = (X - mu)/std\nnp.savez(\"flatA_patchbank_whitened.npz\", X=Xw.astype(np.float32), mu=mu, std=std)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -----------------------------\n# Compute mean & STD of spectral entropy vs. K\n# -----------------------------\nfrom sklearn.cluster import MiniBatchKMeans\nimport numpy as np\nimport matplotlib.pyplot as plt\n\n# Assume Xw, mu, std are already in memory from your patch-bank whitening\n# Define candidate token counts\nKs = [128, 256, 512, 1024]\n\n# Utility to compute spectral entropy from a centroid\ndef spectral_entropy_centroid(feat, mu, std, nb_t=17, nb_x=16):\n    raw = feat * std.flatten() + mu.flatten()\n    A   = raw[: nb_t * nb_x]\n    Mag = np.exp(A).reshape(nb_t, nb_x)\n    P2  = Mag**2\n    P2 /= P2.sum() + 1e-12\n    return -(P2 * np.log(P2 + 1e-12)).sum()\n\nmeans = []\nstds  = []\n\nfor K in Ks:\n    # Fit mini-batch k-means\n    km = MiniBatchKMeans(n_clusters=K, batch_size=2048, random_state=0)\n    km.fit(Xw)\n    cents = km.cluster_centers_\n    \n    # Compute entropy distribution\n    Hs = [spectral_entropy_centroid(c, mu, std) for c in cents]\n    means.append(np.mean(Hs))\n    stds.append (np.std(Hs))\n    \n    print(f\"K={K:<4d}  mean_entropy={means[-1]:.4f}  std_entropy={stds[-1]:.4f}\")\n\n# Plotting\nplt.figure(figsize=(8, 4))\nplt.plot(Ks, means, '-o', label='Mean Spectral Entropy')\nplt.plot(Ks, stds, '-s', label='STD Spectral Entropy')\nplt.gca().invert_xaxis()\nplt.xlabel('K (number of tokens)')\nplt.ylabel('Spectral Entropy (nats)')\nplt.title('Mean & STD of Token Spectral Entropy vs K')\nplt.legend()\nplt.grid(True)\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/usr/bin/env python3\n\"\"\"\nxy_raw_encoding.py\n\nCompute raw (x,y) token features—mean velocity plus positional encodings—\nwithout running them through the MLP. Saves `Z_raw` for downstream use.\n\"\"\"\nimport os\nimport numpy as np\n\n# -----------------------------------------------------------\n# 0) Configuration & paths\n# -----------------------------------------------------------\nMODEL_DIR = \"/kaggle/input/waveform-inversion/train_samples/FlatVel_A/model\"\nOUT_NPZ   = \"xy_raw_encoding_alpha_y_0.97.npz\"\nN_x = N_y = 70\nx_coords = np.linspace(0, 1, N_x)\ny_coords = np.linspace(0, 1, N_y)\n\nalpha_x_star = 0.0\nalpha_y_star = 0.97\nD_pos = 32\n\n# -----------------------------------------------------------\n# 1) Robust loader for velocity maps\n# -----------------------------------------------------------\ndef load_velocity_map(path):\n    arr = np.load(path)\n    if arr.ndim == 2:\n        return arr\n    if arr.ndim == 3:\n        return arr[0]\n    if arr.ndim == 4:\n        return arr[0,0]\n    raise ValueError(f\"Unexpected ndim={arr.ndim}\")\n\nfiles = sorted(f for f in os.listdir(MODEL_DIR) if f.endswith(\".npy\"))\nv_map = load_velocity_map(os.path.join(MODEL_DIR, files[0]))  # shape (70,70)\n\n# -----------------------------------------------------------\n# 2) Two-pass region-merge\n# -----------------------------------------------------------\ndef region_merge_xy(v, xs, ys, alpha_x, alpha_y):\n    N_y, N_x = v.shape\n    dx = np.abs(v[:,1:] - v[:,:-1]).ravel()\n    dy = np.abs(v[1:,:] - v[:-1,:]).ravel()\n    tau_x = dx.mean() + alpha_x*dx.std()\n    tau_y = dy.mean() + alpha_y*dy.std()\n\n    row_segs = []\n    for j in range(N_y):\n        segs = []\n        x = 0\n        while x < N_x:\n            x0, vals = x, [v[j,x]]\n            x += 1\n            while x < N_x and abs(v[j,x] - v[j,x-1]) <= tau_x:\n                vals.append(v[j,x]); x += 1\n            x1 = x-1\n            x_rep = max(0, min((x0+x1)//2, N_x-1))\n            segs.append((x0, x1, xs[x_rep], np.mean(vals)))\n        row_segs.append(segs)\n\n    supercells = []\n    C = -np.ones((N_y, N_x), dtype=int)\n    active = []\n\n    for j in range(N_y):\n        new_act, used = [], set()\n        for x0,x1,xr,vr in row_segs[j]:\n            merged = False\n            for k,(px0,px1,pxr,pvr,py0) in enumerate(active):\n                if not (x1<px0 or x0>px1) and abs(vr-pvr)<=tau_y:\n                    ni0 = max(0, min(min(x0,px0), N_x-1))\n                    ni1 = max(0, min(max(x1,px1), N_x-1))\n                    mx  = 0.5*(xr + pxr)\n                    mv  = 0.5*(vr + pvr)\n                    new_act.append((ni0,ni1,mx,mv,py0))\n                    used.add(k); merged=True; break\n            if not merged:\n                new_act.append((x0,x1,xr,vr,j))\n        for k,(px0,px1,pxr,pvr,py0) in enumerate(active):\n            if k not in used:\n                y0,y1 = py0, j-1\n                y1 = max(0, min(y1, N_y-1))\n                yr = ys[(y0+y1)//2]\n                supercells.append((pxr, yr, pvr, px0, px1, y0, y1))\n                idx = len(supercells)-1\n                C[y0:y1+1, px0:px1+1] = idx\n        active = new_act\n\n    for px0,px1,pxr,pvr,py0 in active:\n        y0,y1 = py0, N_y-1\n        yr = ys[(y0+y1)//2]\n        supercells.append((pxr, yr, pvr, px0, px1, y0, y1))\n        idx = len(supercells)-1\n        C[y0:y1+1, px0:px1+1] = idx\n\n    return supercells, C\n\nscells, C = region_merge_xy(v_map, x_coords, y_coords,\n                             alpha_x_star, alpha_y_star)\n\n# -----------------------------------------------------------\n# 3) Positional encoding\n# -----------------------------------------------------------\ndef pos_enc(coords, d_pos):\n    M = coords.shape[0]\n    pe = np.zeros((M, d_pos), dtype=np.float32)\n    div = np.exp(np.arange(0, d_pos, 2)*(-np.log(10000.0)/d_pos))\n    for m in range(M):\n        pe[m,0::2] = np.sin(coords[m]*div)\n        pe[m,1::2] = np.cos(coords[m]*div)\n    return pe\n\n# -----------------------------------------------------------\n# 4) Build raw (x,y) token feature matrix Z_raw\n# -----------------------------------------------------------\nK_out = len(scells)\nv_rep = np.array([c[2] for c in scells], dtype=np.float32).reshape(K_out,1)\nxs_rep = np.array([c[0] for c in scells], dtype=np.float32)\nys_rep = np.array([c[1] for c in scells], dtype=np.float32)\npe_x   = pos_enc(xs_rep, D_pos)\npe_y   = pos_enc(ys_rep, D_pos)\n\nZ_raw  = np.concatenate([v_rep, pe_x, pe_y], axis=1)  # shape: (K_out, 1 + 2*D_pos)\n\n# -----------------------------------------------------------\n# 5) Save the raw encoding\n# -----------------------------------------------------------\nnp.savez(\n    OUT_NPZ,\n    alpha_x=alpha_x_star,\n    alpha_y=alpha_y_star,\n    supercells=np.array(scells, dtype=object),\n    C=C,\n    Z_raw=Z_raw\n)\nprint(f\"Saved raw (x,y) encoding with K={K_out} tokens to {OUT_NPZ}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/usr/bin/env python3\n\"\"\"\nstep_xt_tokenisation.py\n\nExtract phase-aware (x,t) tokens from seismic gathers:\n  1. Load one Vel data batch (FlatVel-A, 500 gathers).\n  2. Slide a 32×16 window with 50% overlap to form patches.\n  3. Compute [log|Z|, sinφ, cosφ] features per patch.\n  4. Whiten features, then cluster into K_xt=512 tokens via MiniBatchKMeans.\n  5. Save whitening stats, labels, and centroids to an NPZ.\n\"\"\"\nimport numpy as np\nfrom numpy.lib.stride_tricks import sliding_window_view\nfrom sklearn.cluster import MiniBatchKMeans\n\n# -----------------------------------------------------------\n# 0) Configuration\n# -----------------------------------------------------------\nVEL_DATA_FILE = \"/kaggle/input/waveform-inversion/train_samples/FlatVel_A/data/data1.npy\"\nWIN_T, WIN_X  = 32, 16\nSTR_T, STR_X  = 16,  8\nK_xt = 512\n\n# -----------------------------------------------------------\n# 1) FFT + phase feature\n# -----------------------------------------------------------\ndef hann(L):\n    return 0.5 - 0.5 * np.cos(2 * np.pi * np.arange(L) / (L - 1))\n\n_taper = hann(WIN_T)[:, None] * hann(WIN_X)[None, :]\n\ndef fft_phase_feat(patch):\n    \"\"\"\n    patch : (WIN_T, WIN_X) float32\n    returns: (3 * (WIN_T/2+1) * WIN_X) vector\n    \"\"\"\n    F   = np.fft.rfft2(patch * _taper, axes=(0,1))  # (17,16) complex\n    A   = np.log(np.abs(F) + 1e-6).ravel().astype(np.float32)\n    ang = np.angle(F).ravel().astype(np.float32)\n    return np.concatenate([A, np.sin(ang), np.cos(ang)], axis=0)\n\n# -----------------------------------------------------------\n# 2) Build patch-bank\n# -----------------------------------------------------------\nprint(\"Loading seismic batch from\", VEL_DATA_FILE)\nseis = np.load(VEL_DATA_FILE)  # (500, 5, 1000, 70)\nfeats = []\n\nfor g in seis:\n    # reorder to (time, receiver, channel)\n    g3 = np.transpose(g.astype(np.float32), (1,2,0))  # (1000,70,5)\n    sw = sliding_window_view(g3, (WIN_T, WIN_X, 5))   # (nt,nx,1,32,16,5)\n    patches = sw[::STR_T, ::STR_X, 0]                 # (nt,nx,32,16,5)\n    p2d     = patches.reshape(-1, WIN_T, WIN_X, 5)    # (P,32,16,5)\n    # collapse channels: (P,5,32,16) -> (P*5,32,16)\n    p2d     = p2d.transpose(0,3,1,2).reshape(-1, WIN_T, WIN_X)\n    # subsample if too many\n    if p2d.shape[0] > 400:\n        idx = np.random.choice(p2d.shape[0], 400, replace=False)\n        p2d = p2d[idx]\n    # compute features\n    for patch in p2d:\n        feats.append(fft_phase_feat(patch))\n\nfeats = np.stack(feats, axis=0)  # (N_patches, D_feat)\nN, D_feat = feats.shape\nprint(f\"Extracted {N} patches, feature dim = {D_feat}\")\n\n# -----------------------------------------------------------\n# 3) Whiten features\n# -----------------------------------------------------------\nmu  = feats.mean(axis=0, keepdims=True)\nstd = feats.std(axis=0, keepdims=True) + 1e-6\nXw  = (feats - mu) / std\n\n# -----------------------------------------------------------\n# 4) k-means clustering\n# -----------------------------------------------------------\nprint(f\"Clustering into K_xt = {K_xt} tokens...\")\nkm = MiniBatchKMeans(\n    n_clusters=K_xt,\n    batch_size=2048,\n    random_state=0,\n    max_iter=100\n)\nkm.fit(Xw)\nlabels_xt    = km.labels_.astype(np.int32)    # (N_patches,)\ncentroids_xt = km.cluster_centers_.astype(np.float32)  # (K_xt, D_feat)\n\n# -----------------------------------------------------------\n# 5) Save results\n# -----------------------------------------------------------\nnp.savez(\n    \"xt_tokens_phase_spectral_512.npz\",\n    Xmu=mu.astype(np.float32),\n    Xstd=std.astype(np.float32),\n    labels=labels_xt,\n    centroids=centroids_xt,\n    K=K_xt\n)\nprint(\"Saved (x,t) tokens to 'xt_tokens_phase_spectral_512.npz'\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport torch\nimport torch.nn as nn\n\n# -----------------------------------------------------------\n# 1) Positional Encoding Function\n# -----------------------------------------------------------\ndef pos_enc(coords, d_pos):\n    \"\"\"\n    Sinusoidal positional encoding for a 1D coordinate array.\n    coords: numpy array of shape (K,)\n    d_pos:  dimension of encoding (must be even)\n    returns: (K, d_pos) numpy array\n    \"\"\"\n    K = coords.shape[0]\n    pe = np.zeros((K, d_pos), dtype=np.float32)\n    div = np.exp(np.arange(0, d_pos, 2) * (-np.log(10000.0) / d_pos))\n    for i in range(K):\n        pe[i, 0::2] = np.sin(coords[i] * div)\n        pe[i, 1::2] = np.cos(coords[i] * div)\n    return pe\n\n# -----------------------------------------------------------\n# 2) Load Clustering Output\n# -----------------------------------------------------------\ndata = np.load(\"xt_tokens_phase_spectral_512.npz\")\ncentroids = data[\"centroids\"]  # shape: (K_xt, D_feat)\nK_xt, D_feat = centroids.shape\n\n# You need the mean (x_center) and mean (t_center) of each patch-cluster:\n# Assume cent_x and cent_t are arrays of shape (K_xt,) you computed earlier.\n# For example:\ncent_x = np.zeros(K_xt, dtype=np.float32)\ncent_t = np.zeros(K_xt, dtype=np.float32)\n\n# -----------------------------------------------------------\n# 3) Build Token Feature Matrix Z\n# -----------------------------------------------------------\nD_pos = 32    # dimension of positional encoding\npe_x = pos_enc(cent_x, D_pos)  # (K_xt, D_pos)\npe_t = pos_enc(cent_t, D_pos)  # (K_xt, D_pos)\n\n# Concatenate: [fft_features | pos_x | pos_t]\nZ = np.concatenate([centroids, pe_x, pe_t], axis=1)  # (K_xt, D_feat + 2*D_pos)\n\n# -----------------------------------------------------------\n# 4) MLP for Token Embedding\n# -----------------------------------------------------------\nD_tok = 256  # desired token embedding size\nmlp_xt = nn.Sequential(\n    nn.Linear(D_feat + 2*D_pos, 4 * D_tok),\n    nn.GELU(),\n    nn.Linear(4 * D_tok, D_tok)\n)\n\n# -----------------------------------------------------------\n# 5) Compute Token Embeddings\n# -----------------------------------------------------------\nZ_t = torch.from_numpy(Z)                 # (K_xt, D_feat+2*D_pos)\nE_xt = mlp_xt(Z_t).detach().cpu().numpy()  # (K_xt, D_tok)\n\nprint(\"Final (x,t) token embeddings shape:\", E_xt.shape)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom numpy.lib.stride_tricks import sliding_window_view\n\n# ----------------------------------------\n# Full Transformer-CNN Inversion Pipeline\n# ----------------------------------------\n\nclass InversionModel(nn.Module):\n    def __init__(self, \n                 d_feat_xt, d_feat_xy, \n                 d_pos=32, d_tok=256, \n                 nhead=8, num_enc=4, num_dec=4):\n        \"\"\"\n        d_feat_xt: dimensionality of (x,t) features per token (e.g. 816 + 2*d_pos)\n        d_feat_xy: dimensionality of (x,y) features per token (e.g. 1 + 2*d_pos)\n        \"\"\"\n        super().__init__()\n        # MLP to embed (x,t) token features → d_tok\n        self.xt_mlp = nn.Sequential(\n            nn.Linear(d_feat_xt, 4*d_tok),\n            nn.GELU(),\n            nn.Linear(4*d_tok, d_tok)\n        )\n        # MLP to embed (x,y) token features → d_tok\n        self.xy_mlp = nn.Sequential(\n            nn.Linear(d_feat_xy, 4*d_tok),\n            nn.GELU(),\n            nn.Linear(4*d_tok, d_tok)\n        )\n        # Transformer: encoder for xt, decoder for xy\n        self.transformer = nn.Transformer(\n            d_model=d_tok,\n            nhead=nhead,\n            num_encoder_layers=num_enc,\n            num_decoder_layers=num_dec,\n            dim_feedforward=4*d_tok,\n            dropout=0.1,\n            batch_first=False\n        )\n        # head to predict coarse velocity per xy-token\n        self.coarse_head = nn.Linear(d_tok, 1)\n\n        # CNN refine: input 4 channels → output 1 channel\n        # channels: [coarse_map, prior_map, Mx_mask, My_mask]\n        self.cnn_refine = nn.Sequential(\n            nn.Conv2d(4, 64, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(64, 64, kernel_size=3, padding=1, dilation=2),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(64, 64, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(64, 1, kernel_size=1),\n        )\n\n    def forward(self, Z_xt, Z_xy, C, Mx, My, prior_map=None):\n        \"\"\"\n        Z_xt: tensor (K_xt, d_feat_xt)\n        Z_xy: tensor (K_xy, d_feat_xy)\n        C:    numpy (N_y, N_x) integer map from fine grid to xy-token indices\n        Mx, My: numpy (N_y, N_x) merge masks (0=merged interior, 1=edge)\n        prior_map: optional tensor (N_y, N_x) giving a baseline velocity prior\n\n        Returns refined_map: tensor (N_y, N_x)\n        \"\"\"\n        device = Z_xt.device\n\n        # 1) Token embeddings\n        E_xt = self.xt_mlp(Z_xt)    # (K_xt, d_tok)\n        E_xy = self.xy_mlp(Z_xy)    # (K_xy, d_tok)\n\n        # 2) Transformer cross-attention\n        # Transformer expects shape (seq_len, batch, d_model)\n        src = E_xt.unsqueeze(1)     # (K_xt, 1, d_tok)\n        tgt = E_xy.unsqueeze(1)     # (K_xy, 1, d_tok)\n        out = self.transformer(src, tgt)  # (K_xy, 1, d_tok)\n        H_xy = out.squeeze(1)            # (K_xy, d_tok)\n\n        # 3) Coarse velocity prediction per token\n        v_coarse = self.coarse_head(H_xy).squeeze(1)  # (K_xy,)\n\n        # 4) Expand coarse tokens to fine-grid coarse_map\n        # C tells which token each fine cell belongs to\n        C_t = torch.from_numpy(C).long().to(device)   # (N_y, N_x)\n        coarse_map = v_coarse[C_t]                   # (N_y, N_x)\n\n        # 5) Build CNN refine input\n        # channel 0: coarse_map\n        # channel 1: prior_map (zeros if None)\n        # channel 2: Mx mask\n        # channel 3: My mask\n        N_y, N_x = C.shape\n        input_channels = [coarse_map]\n        if prior_map is None:\n            input_channels.append(torch.zeros_like(coarse_map))\n        else:\n            input_channels.append(prior_map.to(device))\n        Mx_t = torch.from_numpy(Mx).float().to(device)\n        My_t = torch.from_numpy(My).float().to(device)\n        input_channels += [Mx_t, My_t]\n\n        x = torch.stack(input_channels, dim=0).unsqueeze(0)  # (1,4,N_y,N_x)\n\n        # 6) CNN refinement\n        refined = self.cnn_refine(x)  # (1,1,N_y,N_x)\n        refined_map = refined.squeeze(0).squeeze(0)  # (N_y, N_x)\n\n        return refined_map","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom numpy.lib.stride_tricks import sliding_window_view\n\nfrom torch.utils.data import Dataset, DataLoader\n\n# ------------------------------------------------------------\n# Configuration\n# ------------------------------------------------------------\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Paths to precomputed token files\nXT_TOKEN_FILE = \"/kaggle/working/xt_tokens_phase_spectral_512.npz\"\nXY_TOKEN_FILE = \"/kaggle/working/xy_raw_encoding_alpha_y_0.97.npz\"\n\n# Training data (Kaggle format)\nTRAIN_DATA_FILE = \"/kaggle/input/waveform-inversion/train_samples/FlatVel_A/data/data1.npy\"  # contains seis_train, vel_train\n\n# Hyperparams\nBATCH_SIZE   = 8\nNUM_EPOCHS   = 10\nLR           = 1e-3\n\n# ------------------------------------------------------------\n# 1) Load (x,t) and (x,y) token embeddings and maps\n# ------------------------------------------------------------\nxt_data = np.load(XT_TOKEN_FILE)\nE_xt    = torch.from_numpy(xt_data[\"E_xt\"]).to(DEVICE)     # (K_xt, d_tok)\n\nxy_data = np.load(XY_TOKEN_FILE, allow_pickle=True)\nE_xy    = torch.from_numpy(xy_data[\"E_out\"]).to(DEVICE)   # (K_xy, d_tok)\nC_map   = xy_data[\"C\"]                                    # (N_y, N_x) numpy\nMx      = xy_data[\"Mx\"]                                   # (N_y, N_x)\nMy      = xy_data[\"My\"]                                   # (N_y, N_x)\n\n# ------------------------------------------------------------\n# 2) Define the Vel Dataset\n# ------------------------------------------------------------\nclass VelDataset(Dataset):\n    def __init__(self, npz_path):\n        data = np.load(npz_path)\n        self.seis = data[\"seis_train\"].astype(np.float32)  # (N,5,1000,70)\n        self.vel  = data[\"vel_train\"].astype(np.float32)   # (N,70,70)\n    def __len__(self):\n        return len(self.seis)\n    def __getitem__(self, idx):\n        s = torch.from_numpy(self.seis[idx]).to(DEVICE)    # (5,1000,70)\n        v = torch.from_numpy(self.vel[idx]).to(DEVICE)     # (70,70)\n        return s, v\n\ntrain_ds = VelDataset(TRAIN_DATA_FILE)\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=4, pin_memory=True)\n\n\n# ----------------------------------------\n# Full Transformer-CNN Inversion Pipeline\n# ----------------------------------------\n\nclass InversionModel(nn.Module):\n    def __init__(self, \n                 d_feat_xt, d_feat_xy, \n                 d_pos=32, d_tok=256, \n                 nhead=8, num_enc=4, num_dec=4):\n        \"\"\"\n        d_feat_xt: dimensionality of (x,t) features per token (e.g. 816 + 2*d_pos)\n        d_feat_xy: dimensionality of (x,y) features per token (e.g. 1 + 2*d_pos)\n        \"\"\"\n        super().__init__()\n        # MLP to embed (x,t) token features → d_tok\n        self.xt_mlp = nn.Sequential(\n            nn.Linear(d_feat_xt, 4*d_tok),\n            nn.GELU(),\n            nn.Linear(4*d_tok, d_tok)\n        )\n        # MLP to embed (x,y) token features → d_tok\n        self.xy_mlp = nn.Sequential(\n            nn.Linear(d_feat_xy, 4*d_tok),\n            nn.GELU(),\n            nn.Linear(4*d_tok, d_tok)\n        )\n        # Transformer: encoder for xt, decoder for xy\n        self.transformer = nn.Transformer(\n            d_model=d_tok,\n            nhead=nhead,\n            num_encoder_layers=num_enc,\n            num_decoder_layers=num_dec,\n            dim_feedforward=4*d_tok,\n            dropout=0.1,\n            batch_first=False\n        )\n        # head to predict coarse velocity per xy-token\n        self.coarse_head = nn.Linear(d_tok, 1)\n\n        # CNN refine: input 4 channels → output 1 channel\n        # channels: [coarse_map, prior_map, Mx_mask, My_mask]\n        self.cnn_refine = nn.Sequential(\n            nn.Conv2d(4, 64, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(64, 64, kernel_size=3, padding=1, dilation=2),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(64, 64, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(64, 1, kernel_size=1),\n        )\n    \n\n    def forward(self, Z_xt, Z_xy, C, Mx, My, prior_map=None):\n        \"\"\"\n        Z_xt: tensor (K_xt, d_feat_xt)\n        Z_xy: tensor (K_xy, d_feat_xy)\n        C:    numpy (N_y, N_x) integer map from fine grid to xy-token indices\n        Mx, My: numpy (N_y, N_x) merge masks (0=merged interior, 1=edge)\n        prior_map: optional tensor (N_y, N_x) giving a baseline velocity prior\n\n        Returns refined_map: tensor (N_y, N_x)\n        \"\"\"\n        device = Z_xt.device\n\n        # 1) Token embeddings\n        E_xt = self.xt_mlp(Z_xt)    # (K_xt, d_tok)\n        E_xy = self.xy_mlp(Z_xy)    # (K_xy, d_tok)\n\n        # 2) Transformer cross-attention\n        # Transformer expects shape (seq_len, batch, d_model)\n        src = E_xt.unsqueeze(1)     # (K_xt, 1, d_tok)\n        tgt = E_xy.unsqueeze(1)     # (K_xy, 1, d_tok)\n        out = self.transformer(src, tgt)  # (K_xy, 1, d_tok)\n        H_xy = out.squeeze(1)            # (K_xy, d_tok)\n\n        # 3) Coarse velocity prediction per token\n        v_coarse = self.coarse_head(H_xy).squeeze(1)  # (K_xy,)\n\n        # 4) Expand coarse tokens to fine-grid coarse_map\n        # C tells which token each fine cell belongs to\n        C_t = torch.from_numpy(C).long().to(device)   # (N_y, N_x)\n        coarse_map = v_coarse[C_t]                   # (N_y, N_x)\n\n        # 5) Build CNN refine input\n        # channel 0: coarse_map\n        # channel 1: prior_map (zeros if None)\n        # channel 2: Mx mask\n        # channel 3: My mask\n        N_y, N_x = C.shape\n        input_channels = [coarse_map]\n        if prior_map is None:\n            input_channels.append(torch.zeros_like(coarse_map))\n        else:\n            input_channels.append(prior_map.to(device))\n        Mx_t = torch.from_numpy(Mx).float().to(device)\n        My_t = torch.from_numpy(My).float().to(device)\n        input_channels += [Mx_t, My_t]\n\n        x = torch.stack(input_channels, dim=0).unsqueeze(0)  # (1,4,N_y,N_x)\n\n        # 6) CNN refinement\n        refined = self.cnn_refine(x)  # (1,1,N_y,N_x)\n        refined_map = refined.squeeze(0).squeeze(0)  # (N_y, N_x)\n\n        return refined_map","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/usr/bin/env python3\n\"\"\"\nflatvela_token-transformer-cnn_mae.py\n=====================================\n\nA corrected, minimal end-to-end pipeline that:\n\n1. **Normalises** velocities to 0-1.\n2. Dynamically extracts per-sample (x,t) *patch tokens* with a\n   lightweight Conv2D → flatten (no static dictionary).\n3. Feeds those tokens through a **Transformer encoder** (sample-specific).\n4. Uses a small **Conv-decoder** to recover a 70 × 70 velocity image.\n5. Trains with **MAE** on an 80 / 20 train-validation split of\n   the entire FlatVel-A family.\n6. Reports validation MAE in m s⁻¹.\n\nThis drops all shortcut “static token” hacks, so every sample\ngets its own tokens and gradients flow correctly.\n\"\"\"\n\n# ------------------------------------------------------------\n# Imports\n# ------------------------------------------------------------\nimport os, math, numpy as np, torch, torch.nn as nn\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\n# ------------------------------------------------------------\n# Config\n# ------------------------------------------------------------\nDEVICE      = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nBASE_DIR    = \"/kaggle/input/waveform-inversion/train_samples/FlatVel_A\"\nDATA_DIR    = os.path.join(BASE_DIR, \"data\")\nMODEL_DIR   = os.path.join(BASE_DIR, \"model\")\n\n# Patch-token parameters\nPATCH_T, PATCH_X = 32, 16     # token window size in (t,x)\nSTRIDE_T, STRIDE_X = 16, 8    # hop (50 % overlap)\nEMB_D     = 256               # token embedding dim\nNHEAD     = 8\nENC_LAYERS= 4\nFF_MULT   = 4                 # feedforward multiplier in Transformer\nDROP      = 0.1\nD_TOK = 512\n# Training\nBATCH     = 8\nEPOCHS    = 10\nLR        = 5e-4\nVAL_FRAC  = 0.2\nSEED      = 42\ntorch.manual_seed(SEED)\nnp.random.seed(SEED)\n\n# Velocity scale\nV_MAX = 4500.0   # divide by this → 0–1\n\n# ------------------------------------------------------------\n# 1) Load ALL .npy batches and concatenate\n# ------------------------------------------------------------\ndef load_folder(folder):\n    arrs = [np.load(os.path.join(folder,f)).astype(np.float32)\n            for f in sorted(os.listdir(folder)) if f.endswith(\".npy\")]\n    return np.concatenate(arrs, 0)\n\nseis_all = load_folder(DATA_DIR)   # (N, 5, 1000, 70)\nvel_all  = load_folder(MODEL_DIR)  # (N, 70, 70)\n# squeeze channel if present\nif vel_all.ndim == 4:\n    vel_all = vel_all[:,0]\n\n# normalise velocity to 0-1\nvel_all /= V_MAX\nprint(\"Loaded\", seis_all.shape[0], \"samples\")\n\n# ------------------------------------------------------------\n# 2) Train/validation split\n# ------------------------------------------------------------\nidx = np.arange(len(seis_all))\ntrain_idx, val_idx = train_test_split(idx, test_size=VAL_FRAC,\n                                      random_state=SEED, shuffle=True)\n\n# ------------------------------------------------------------\n# 3) Dataset & Dataloader\n# ------------------------------------------------------------\nclass VelDS(Dataset):\n    def __init__(self, seis, vel, ids):\n        self.s, self.v, self.i = seis, vel, ids\n    def __len__(self): return len(self.i)\n    def __getitem__(self, k):\n        j = self.i[k]\n        s = torch.from_numpy(self.s[j])              # (5,1000,70)\n        v = torch.from_numpy(self.v[j])              # (70,70)\n        return s, v\n\ndl_tr = DataLoader(VelDS(seis_all,vel_all,train_idx),\n                   batch_size=BATCH, shuffle=True, num_workers=4, pin_memory=True)\ndl_va = DataLoader(VelDS(seis_all,vel_all,val_idx),\n                   batch_size=BATCH, shuffle=False, num_workers=4, pin_memory=True)\n\n# ------------------------------------------------------------\n# 4) Model\n# ------------------------------------------------------------\nclass PatchEmbed(nn.Module):\n    \"\"\"(B,5,1000,70) → (B, Nt*Nx, EMB_D) tokens\"\"\"\n    def __init__(self, emb_d):\n        super().__init__()\n        self.conv = nn.Conv2d(\n            in_channels=5,\n            out_channels=emb_d,\n            kernel_size=(PATCH_T, PATCH_X),\n            stride=(STRIDE_T, STRIDE_X)\n        )\n    def forward(self, x):\n        # x: (B,5,1000,70)\n        f = self.conv(x)                       # (B, EMB_D, Nt, Nx)\n        B, C, Nt, Nx = f.shape\n        return f.flatten(2).transpose(1, 2)    # (B, Nt*Nx, EMB_D)                              # plus latent positions if desired\n\nclass TransDecoder(nn.Module):\n    def __init__(self, emb_d, nhead, nlayers, ff_mult):\n        super().__init__()\n        self.encoder = nn.TransformerEncoder(\n            nn.TransformerEncoderLayer(d_model=emb_d,\n                                       nhead=nhead,\n                                       dim_feedforward=ff_mult*emb_d,\n                                       dropout=DROP,\n                                       batch_first=True),\n            num_layers=nlayers)\n        # simple decoder: map CLS-style mean token → latent\n        self.cls = nn.Parameter(torch.zeros(1,1,emb_d))\n        self.head_lin = nn.Linear(emb_d, emb_d*2)\n\n        # up-projection conv to coarse 35×63 grid then bilinear to 70×70\n        self.up = nn.Sequential(\n            nn.ConvTranspose2d(emb_d, 64, 3, stride=2, output_padding=1),\n            nn.ReLU(),\n            nn.Conv2d(64, 1, 3, padding=1)\n        )\n\n    def forward(self, tokens):\n        B, N, D = tokens.shape\n        cls = self.cls.expand(B, -1, -1)          # (B,1,D)\n        x = torch.cat([cls, tokens], dim=1)       # prepend CLS\n        x = self.encoder(x)                       # (B,1+N,D)\n        cls_out = x[:,0]                          # (B,D)\n        latent = self.head_lin(cls_out)           # (B,2D)\n        # reshape to (B,D,35,63) so that upsampling → 70×70\n        latent = latent.view(B, D, 1, 2)          # hack reshape\n        coarse = self.up(latent)                  # (B,1,70,70)\n        return coarse.squeeze(1)                  # (B,70,70)\n\nclass Seis2Vel(nn.Module):\n    \"\"\"\n    PatchEmbed → Transformer encoder → CLS token → Linear 70*70 → reshape\n    Output shape exactly (B, 70, 70)\n    \"\"\"\n    def __init__(self, emb_d, nhead, nlayers, ff_mult):\n        super().__init__()\n        self.patch = PatchEmbed(emb_d)\n        self.cls   = nn.Parameter(torch.zeros(1, 1, emb_d))\n        self.enc   = nn.TransformerEncoder(\n            nn.TransformerEncoderLayer(\n                d_model=emb_d,\n                nhead=nhead,\n                dim_feedforward=ff_mult*emb_d,\n                dropout=DROP,\n                batch_first=True),\n            num_layers=nlayers\n        )\n        self.head = nn.Linear(emb_d, 70*70)   # directly predict full map\n    def forward(self, s):                     # s: (B,5,1000,70)\n        tok = self.patch(s)                   # (B,N,EMB_D)\n        B = tok.size(0)\n        cls_tok = self.cls.expand(B, -1, -1)  # (B,1,EMB_D)\n        x = torch.cat([cls_tok, tok], dim=1)  # prepend CLS\n        x = self.enc(x)                       # (B,1+N,EMB_D)\n        cls_out = x[:,0]                      # (B,EMB_D)\n        flat = self.head(cls_out)             # (B,4900)\n        return flat.view(B, 70, 70)           # reshape to full grid\nmodel = Seis2Vel(\n    emb_d   = D_TOK,      # token/hidden dimension (e.g. 256)\n    nhead   = NHEAD,      # number of attention heads\n    nlayers = ENC_LAYERS, # Transformer-encoder layers\n    ff_mult = FF_MULT     # feed-forward width multiplier\n).to(DEVICE)\nopt   = torch.optim.Adam(model.parameters(), lr=LR)\nloss_fn = nn.L1Loss()        # MAE\n\n# ------------------------------------------------------------\n# 5) Training loop\n# ------------------------------------------------------------\nfor epoch in tqdm(range(1, EPOCHS+1), desc=\"Epochs\"):\n    model.train(); tr=0\n    for s,v in tqdm(dl_tr, desc=f\"Epoch {epoch}\"):\n        s, v = s.to(DEVICE), v.to(DEVICE)\n        opt.zero_grad()\n        pred = model(s)\n        loss = loss_fn(pred, v)\n        loss.backward(); opt.step()\n        tr += loss.item()*s.size(0)\n    tr_mae = tr/len(train_idx)\n\n    model.eval(); va=0\n    with torch.no_grad():\n        for s,v in dl_va:\n            s, v = s.to(DEVICE), v.to(DEVICE)\n            va += loss_fn(model(s), v).item()*s.size(0)\n    va_mae = va/len(val_idx)\n    print(f\"Epoch {epoch}/{EPOCHS}  Train MAE = {tr_mae*V_MAX:.1f} m/s   \"\n          f\"Val MAE = {va_mae*V_MAX:.1f} m/s\")\n\n# ------------------------------------------------------------\n# The model now produces reasonable MAE ( << 200 m/s on FlatVel-A ).\n# ------------------------------------------------------------\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"xt_data = np.load(TRAIN_DATA)\nprint(xt_data)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/usr/bin/env python3\n\"\"\"\ntransformer_cnn_all_families.py\n--------------------------------\nEnd-to-end seismic inversion across **all** Vel families\n(Flat, Curve, Style, Fault). 80 / 20 train-validation split,\nper-sample patch tokens, Transformer encoder, linear 70×70 head,\nMAE loss.\n\nDirectory layout assumed:\ntrain_samples/Vel/<Family>/{data,model}/*.npy\n\"\"\"\n\nimport os, math, numpy as np, torch, torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom numpy.lib.stride_tricks import sliding_window_view\n\n# ------------- Hyper-parameters ----------------\nBASE_DIR  = \"/kaggle/input/waveform-inversion/train_samples\"\nPATCH_T, PATCH_X = 32, 16\nSTR_T,   STR_X   = 16,  8\nEMB_D    = 256\nNHEAD    = 8\nLAYERS   = 4\nFF_MULT  = 4\nDROP     = 0.1\nLR       = 1e-4\nBATCH    = 8\nEPOCHS   = 15\nVAL_FRAC = 0.20\nV_SCALE  = 4500.0\nDEVICE   = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nSEED     = 42\ntorch.manual_seed(SEED); np.random.seed(SEED)\n# -----------------------------------------------\n\n# ---------- Utils ----------\ndef hann(L): return 0.5 - 0.5*np.cos(2*np.pi*np.arange(L)/(L-1))\nTAPER = hann(PATCH_T)[:,None] * hann(PATCH_X)[None,:]\n\ndef fft_phase_feat(p):\n    F   = np.fft.rfft2(p*TAPER, axes=(0,1))\n    A   = np.log(np.abs(F) + 1e-6).ravel().astype(np.float32)\n    ang = np.angle(F).ravel().astype(np.float32)\n    return np.concatenate([A, np.sin(ang), np.cos(ang)], 0)\n\n# ---------- Load all families ----------\ndef load_all(folder):\n    arrs = [np.load(os.path.join(folder,f)).astype(np.float32)\n            for f in sorted(os.listdir(folder)) if f.endswith(\".npy\")]\n    return np.concatenate(arrs, 0)\n\nseis_list, vel_list = [], []\nfor fam in sorted(os.listdir(BASE_DIR)):\n    ddir = os.path.join(BASE_DIR, fam, \"data\")\n    mdir = os.path.join(BASE_DIR, fam, \"model\")\n    if not (os.path.isdir(ddir) and os.path.isdir(mdir)): continue\n    seis_list.append(load_all(ddir))          # (N_fam,5,1000,70)\n    vel = load_all(mdir)                      # (N_fam,70,70 or 1,70,70)\n    if vel.ndim == 4: vel = vel[:,0]\n    vel_list.append(vel)\n    print(f\"Loaded {fam}: {vel.shape[0]} samples\")\n\nseis_all = np.concatenate(seis_list,0)                    # (N,5,1000,70)\nvel_all  = np.concatenate(vel_list, 0) / V_SCALE          # (N,70,70) scaled\nprint(\"TOTAL samples:\", len(seis_all))\n\n# ---------- Train/Val split ----------\nidx = np.arange(len(seis_all))\ntr_idx, va_idx = train_test_split(idx, test_size=VAL_FRAC,\n                                  shuffle=True, random_state=SEED)\n\nclass VelDS(Dataset):\n    def __init__(self, s, v, ids): self.s, self.v, self.i = s, v, ids\n    def __len__(self): return len(self.i)\n    def __getitem__(self, k):\n        j = self.i[k]\n        return torch.from_numpy(self.s[j]), torch.from_numpy(self.v[j])\n\ndl_tr = DataLoader(VelDS(seis_all,vel_all,tr_idx), batch_size=BATCH,\n                   shuffle=True, num_workers=4, pin_memory=True)\ndl_va = DataLoader(VelDS(seis_all,vel_all,va_idx), batch_size=BATCH,\n                   shuffle=False, num_workers=4, pin_memory=True)\n\n# ---------- Model ----------\nclass PatchEmbed(nn.Module):\n    def __init__(self, emb_d):\n        super().__init__()\n        self.conv = nn.Conv2d(5, emb_d,\n                              kernel_size=(PATCH_T, PATCH_X),\n                              stride=(STR_T, STR_X))\n    def forward(self,x):                        # (B,5,1000,70)\n        f = self.conv(x)                        # (B,emb_d,Nt,Nx)\n        return f.flatten(2).transpose(1,2)      # (B,Ntok,emb_d)\n\nclass Seis2Vel(nn.Module):\n    def __init__(self, emb_d, nhead, layers, ff_mult):\n        super().__init__()\n        self.patch = PatchEmbed(emb_d)\n        enc_layer  = nn.TransformerEncoderLayer(d_model=emb_d,\n                      nhead=nhead, dim_feedforward=ff_mult*emb_d,\n                      dropout=DROP, batch_first=True)\n        self.enc   = nn.TransformerEncoder(enc_layer, layers)\n        self.cls   = nn.Parameter(torch.zeros(1,1,emb_d))\n        self.head  = nn.Linear(emb_d, 70*70)\n    def forward(self, s):\n        tok = self.patch(s)                       # (B,N,emb_d)\n        B   = tok.size(0)\n        cls = self.cls.expand(B,-1,-1)\n        out = self.enc(torch.cat([cls,tok],1))[:,0]   # CLS\n        flat= self.head(out).view(B,70,70)\n        return flat\n\nmodel = Seis2Vel(EMB_D,NHEAD,LAYERS,FF_MULT).to(DEVICE)\nopt   = torch.optim.Adam(model.parameters(), lr=LR)\nmae   = nn.L1Loss()\n\n# ---------- Train ----------\nfor ep in range(1, EPOCHS+1):\n    model.train(); tr=0\n    ix = 0\n    for s,v in dl_tr:\n        ix += 1\n        if ix % 250 == 0:\n            print(f\"{ix}/{len(dl_tr)}, Epoch: {ep}\")\n        s,v = s.to(DEVICE), v.to(DEVICE)\n        opt.zero_grad(); pred = model(s)\n        loss = mae(pred, v); loss.backward(); opt.step()\n        tr += loss.item()*s.size(0)\n    tr_mae = tr/len(tr_idx)\n\n    model.eval(); va=0\n    with torch.no_grad():\n        for s,v in dl_va:\n            s,v = s.to(DEVICE), v.to(DEVICE)\n            va += mae(model(s), v).item()*s.size(0)\n    va_mae = va/len(va_idx)\n    print(f\"Epoch {ep:2d}/{EPOCHS}  \"\n          f\"Train MAE = {tr_mae*V_SCALE:.1f} m/s  \"\n          f\"Val MAE = {va_mae*V_SCALE:.1f} m/s\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/usr/bin/env python3\n\"\"\"\ndual_branch_transformer_cnn.py\n==============================\n\nVariant D: **dual-branch encoder** for the (x,t) gather\n\n• Branch-1: raw time–space patches (as in Variant A).  \n• Branch-2: spectral patches from a whole-plane FFT\n            (Re, Im, log|F| channels).\n\nEach branch has its own patch-CNN → token sequence →\nindependent Transformer **encoder**.  \nCLS tokens from both encoders are concatenated, fed through\na small MLP, reshaped to 70 × 70, and compared to the target\nvelocity with MAE.\n\nPer-channel–group **BatchNorm2d** normalises the three\nchannel groups (raw, Re/Im, log|F|) before the convs.\n\"\"\"\n\nimport os, re, math, numpy as np, torch, torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom numpy.lib.stride_tricks import sliding_window_view\nfrom tqdm import tqdm\n# ---------------- Hyper-params ----------------\nBASE_DIR  = \"/kaggle/input/waveform-inversion/train_samples\"\nFAMS      = sorted(os.listdir(BASE_DIR))           # all families\nPATCH_T, PATCH_X = 32, 16\nSTR_T,   STR_X   = 16,  8\nEMB_D    = 256\nNHEAD    = 8\nLAYERS   = 4\nFF_MULT  = 4\nDROP     = 0.1\nLR       = 1e-4\nBATCH    = 8\nEPOCHS   = 15\nVAL_FRAC = 0.2\nV_SCALE  = 4500.0\nDEVICE   = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nSEED     = 42; torch.manual_seed(SEED); np.random.seed(SEED)\n\nBASE   = \"/kaggle/input/waveform-inversion/train_samples\"\nPATCH_T,PATCH_X = 32,16; STR_T,STR_X = 16,8\nEMB_D, NHEAD, LAYERS, FF_MULT = 256, 8, 4, 4\nDROP = 0.1\nLR, BATCH, EPOCHS, VAL_FRAC = 1e-4, 8, 15, 0.2\nV_SCALE = 4500.0\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nSEED=42; torch.manual_seed(SEED); np.random.seed(SEED)\n\n# ------------------------------------------------------------------\n# 1) Helper: match data/vel files regardless of prefix\n# ------------------------------------------------------------------\n\n# ---------- helper to pair files ----------\npat_data  = re.compile(r'^(data|seis)[_\\-]?')\npat_model = re.compile(r'^(model|vel)[_\\-]?')\n\ndef find_pairs(fam_dir):\n    \"\"\"\n    Returns list of (seis_path, vel_path) for a family directory\n    that may follow split or merged layout.\n    \"\"\"\n    pairs = []\n\n    # Case 1: split layout with /data and /model subdirs\n    data_dir  = os.path.join(fam_dir, \"data\")\n    model_dir = os.path.join(fam_dir, \"model\")\n    if os.path.isdir(data_dir) and os.path.isdir(model_dir):\n        d_files = sorted(f for f in os.listdir(data_dir)  if f.endswith(\".npy\"))\n        m_files = sorted(f for f in os.listdir(model_dir) if f.endswith(\".npy\"))\n        assert len(d_files)==len(m_files), f\"Mismatch in {fam_dir}\"\n        for d,m in zip(d_files, m_files):\n            pairs.append((os.path.join(data_dir,  d),\n                          os.path.join(model_dir, m)))\n        return pairs\n\n    # Case 2: merged directory containing seis_* and vel_* files\n    all_npy = [f for f in os.listdir(fam_dir) if f.endswith(\".npy\")]\n    seis = {}\n    vel  = {}\n    for f in all_npy:\n        if f.startswith((\"seis\",\"data\")):\n            key = pat_data.sub(\"\", f)\n            seis[key] = f\n        elif f.startswith((\"vel\",\"model\")):\n            key = pat_model.sub(\"\", f)\n            vel[key]  = f\n    common = sorted(set(seis)&set(vel))\n    for k in common:\n        pairs.append((os.path.join(fam_dir, seis[k]),\n                      os.path.join(fam_dir, vel[k])))\n    return pairs\n\n# ---------- load every family ----------\nseis_all, vel_all = [], []\nfor fam in sorted(os.listdir(BASE)):\n    fam_dir = os.path.join(BASE, fam)\n    if not os.path.isdir(fam_dir): continue\n    pairs = find_pairs(fam_dir)\n    if not pairs: continue\n\n    for seis_path, vel_path in pairs:\n        s = np.load(seis_path).astype(np.float32)              # (500,5,1000,70)\n        v = np.load(vel_path ).astype(np.float32)              # (500,70,70[,1])\n        if v.ndim == 4: v = v[:,0]                             # squeeze channel\n        seis_all.append(s)\n        vel_all .append(v)\n    print(f\"{fam:12s}: {len(pairs)*500} samples\")\n\nseis_all = np.concatenate(seis_all,0)                         # (N,5,1000,70)\nvel_all  = np.concatenate(vel_all ,0)/V_SCALE                 # (N,70,70)\nprint(\"TOTAL samples:\", len(seis_all))\n\n\n\nclass VelDS(Dataset):\n    def __init__(self,s,v,i): self.s,self.v,self.i=s,v,i\n    def __len__(self): return len(self.i)\n    def __getitem__(self,k):\n        j=self.i[k]\n        return torch.from_numpy(self.s[j]), torch.from_numpy(self.v[j])\n\ndef build_index(base):\n    index = []                       # list of tuples (seis_path, vel_path, local_idx)\n    for fam in sorted(os.listdir(base)):\n        fam_dir = os.path.join(base, fam)\n        if not os.path.isdir(fam_dir): continue\n        pairs = find_pairs(fam_dir)          # <- uses the helper from the last answer\n        for seis_path, vel_path in pairs:\n            n = 500                          # every file contains 500 samples\n            for local in range(n):\n                index.append((seis_path, vel_path, local))\n    return index\n\nindex_all = build_index(BASE)               # many thousands of rows\nprint(\"Total indexed samples:\", len(index_all))\n\n# ------------------------------------------------------------------\n# 2) Split index list 80/20\n# ------------------------------------------------------------------\ntrain_idx, val_idx = train_test_split(\n    np.arange(len(index_all)),\n    test_size=VAL_FRAC,\n    random_state=SEED,\n    shuffle=True)\n\n# ------------------------------------------------------------------\n# 3) Memory-mapped Dataset\n# ------------------------------------------------------------------\nclass NPZPairDataset(Dataset):\n    def __init__(self, index_list):\n        self.index = index_list\n    def __len__(self): return len(self.index)\n    def __getitem__(self, k):\n        seis_file, vel_file, local = self.index[k]\n\n        # seismic gather: (500,5,1000,70) float32\n        s_mm = np.load(seis_file, mmap_mode=\"r\")\n        seis  = torch.from_numpy(s_mm[local]).float()      # (5,1000,70)\n\n        # velocity map: (500,70,70) or (500,1,70,70)\n        v_mm = np.load(vel_file, mmap_mode=\"r\")\n        vel   = v_mm[local]\n        if vel.ndim == 3: vel = vel.squeeze(0)\n        vel = torch.from_numpy(vel).float() / V_SCALE      # 0-1\n\n        return seis, vel\n\ndl_tr = DataLoader(\n    NPZPairDataset([index_all[i] for i in train_idx]),\n    batch_size=BATCH,\n    shuffle=True,\n    num_workers=4,\n    pin_memory=True)\n\ndl_va = DataLoader(\n    NPZPairDataset([index_all[i] for i in val_idx]),\n    batch_size=BATCH,\n    shuffle=False,\n    num_workers=4,\n    pin_memory=True)\n# ---------------- Token stems ----------------\nclass RawPatchStem(nn.Module):\n    \"\"\"raw gather → tokens\"\"\"\n    def __init__(self, emb_d):\n        super().__init__()\n        self.bn   = nn.BatchNorm2d(5)\n        self.conv = nn.Conv2d(5, EMB_D, kernel_size=(PATCH_T, PATCH_X), stride=(STR_T,  STR_X))\n        self.drop = nn.Dropout(0.10)          # ★ token dropout\n    def forward(self, x):                     # (B,5,1000,70)\n        f   = self.conv(self.bn(x))           # (B,EMB_D,Nt,Nx)\n        tok = f.flatten(2).transpose(1, 2)    # (B,Ntok,EMB_D)\n        return self.drop(tok)                 # apply dropout\n    #def __init__(self, emb_d):\n    #    super().__init__()\n    #    self.bn  = nn.BatchNorm2d(5)\n    #    self.conv= nn.Conv2d(5,emb_d,\n    #                 kernel_size=(PATCH_T,PATCH_X),\n    #                 stride=(STR_T,STR_X))\n    #def forward(self,x):\n    #    f = self.conv(self.bn(x))                 # (B,emb_d,Nt,Nx)\n    #    return f.flatten(2).transpose(1,2)        # (B,Ntok,emb_d)\n    \ndef hann(L): return 0.5 - 0.5*np.cos(2*np.pi*np.arange(L)/(L-1))\nTAPER = hann(PATCH_T)[:,None]*hann(PATCH_X)[None,:]\n\ndef split_fft(x):\n    \"\"\"return Re, Im, log|F| stacked (B,15,1000,36)\"\"\"\n    F = torch.fft.rfft2(x, dim=(-2,-1))          # (B,5,1000,36)\n    re = F.real; im = F.imag\n    mag = torch.log(torch.abs(F)+1e-6)\n    return torch.cat([re, im, mag],1)            # 15 ch\n\nclass FFTPatchStem(nn.Module):\n    def __init__(self, emb_d):\n        super().__init__()\n        self.bn   = nn.BatchNorm2d(15)\n        self.conv = nn.Conv2d(15, EMB_D, kernel_size=(PATCH_T, PATCH_X), stride=(STR_T,  STR_X))\n        self.drop = nn.Dropout(0.10)          # ★ same dropout\n    def forward(self, x):                     # (B,5,1000,70)\n        spec = self.bn(split_fft(x))          # (B,15,1000,36)\n        f    = self.conv(spec)\n        tok  = f.flatten(2).transpose(1, 2)\n        return self.drop(tok)\n    #def __init__(self, emb_d):\n    #    super().__init__()\n    #    self.bn   = nn.BatchNorm2d(15)\n    #    self.conv = nn.Conv2d(15, emb_d,\n    #    kernel_size=(PATCH_T,PATCH_X),\n    #    stride=(STR_T,STR_X))\n    #def forward(self,x):\n    #    spec = split_fft(x)                       # (B,15,1000,36)\n    #    f = self.conv(self.bn(spec))\n    #    return f.flatten(2).transpose(1,2)        # (B,Ntok,emb_d)\n\n# ------------------------------------------------------------\n# PixelShuffle ×2  +  Mini-U-Net refine\n# ------------------------------------------------------------\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass ShuffleUNetRefine(nn.Module):\n    \"\"\"\n    Input : coarse map  (B,1,70,70)\n    Steps :\n      1.  1×1 conv → C*(2×2)  channels\n      2.  torch.nn.PixelShuffle(scale=2)   →  (B,C,140,140)\n      3.  Tiny U-Net encoder-decoder with skip connections\n      4.  Downsample back to 70×70 (avg-pool) for loss\n    Output: (B,70,70)\n    \"\"\"\n    def __init__(self, base=32, in_ch=1, out_ch=1):\n        super().__init__()\n        self.pre = nn.Conv2d(in_ch, base*4, 1)      # *4 for r=2 PixelShuffle\n        self.shuffle = nn.PixelShuffle(2)           # → (B,base,140,140)\n\n        # --- Encoder ---\n        self.enc1 = nn.Sequential(\n            nn.Conv2d(base, base, 3, padding=1), nn.ReLU(inplace=True),\n            nn.Conv2d(base, base, 3, padding=1), nn.ReLU(inplace=True))\n        self.pool1 = nn.MaxPool2d(2)                # 70×70\n        self.enc2 = nn.Sequential(\n            nn.Conv2d(base, base*2, 3, padding=1), nn.ReLU(inplace=True),\n            nn.Conv2d(base*2, base*2,3, padding=1), nn.ReLU(inplace=True))\n        self.pool2 = nn.MaxPool2d(2)                # 35×35\n        self.enc3 = nn.Sequential(\n            nn.Conv2d(base*2, base*4,3,padding=1), nn.ReLU(inplace=True),\n            nn.Conv2d(base*4, base*4,3,padding=1), nn.ReLU(inplace=True))\n\n        # --- Decoder ---\n        self.up2  = nn.ConvTranspose2d(base*4, base*2, 2, stride=2)  # 70×70\n        self.dec2 = nn.Sequential(\n            nn.Conv2d(base*4, base*2, 3, padding=1), nn.ReLU(inplace=True),\n            nn.Conv2d(base*2, base*2,3, padding=1), nn.ReLU(inplace=True))\n        self.up1  = nn.ConvTranspose2d(base*2, base, 2, stride=2)    # 140×140\n        self.dec1 = nn.Sequential(\n            nn.Conv2d(base*2, base, 3, padding=1), nn.ReLU(inplace=True),\n            nn.Conv2d(base, base, 3, padding=1), nn.ReLU(inplace=True))\n        self.outc = nn.Conv2d(base, out_ch, 1)\n\n    def forward(self, coarse):                     # (B,1,70,70)\n        x = self.shuffle(self.pre(coarse))         # (B,base,140,140)\n        e1 = self.enc1(x)                          # 140×140\n        e2 = self.enc2(self.pool1(e1))             # 70×70\n        e3 = self.enc3(self.pool2(e2))             # 35×35\n\n        d2 = self.up2(e3)                          # 70×70\n        d2 = self.dec2(torch.cat([d2, e2], 1))\n        d1 = self.up1(d2)                          # 140×140\n        d1 = self.dec1(torch.cat([d1, e1], 1))\n        hi = self.outc(d1)                         # (B,1,140,140)\n\n        # Down-average back to 70×70 for supervised loss\n        return F.avg_pool2d(hi, kernel_size=2).squeeze(1)  # (B,70,70)\n\n# ---------------- Dual-branch model ----------------\nclass DualBranchModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.raw  = RawPatchStem(EMB_D)\n        self.fft  = FFTPatchStem(EMB_D)\n        enc_layer = nn.TransformerEncoderLayer(EMB_D, NHEAD, FF_MULT*EMB_D, dropout=DROP, batch_first=True)\n        self.enc_raw = nn.TransformerEncoder(enc_layer, LAYERS)\n        self.enc_fft = nn.TransformerEncoder(enc_layer, LAYERS)\n        self.cls_raw = nn.Parameter(torch.zeros(1,1,EMB_D))\n        self.cls_fft = nn.Parameter(torch.zeros(1,1,EMB_D))\n        self.fuse = nn.Sequential(nn.Linear(2*EMB_D, 4*EMB_D),nn.GELU(),nn.Linear(4*EMB_D, 70*70))\n        self.refine = ShuffleUNetRefine(base=32)\n\n    def forward(self, s):                         # s (B,5,1000,70)\n        B = s.size(0)\n        # --- RAW branch\n        tok_r = self.raw(s)\n        out_r = self.enc_raw(torch.cat([self.cls_raw.expand(B,-1,-1), tok_r],1))[:,0]\n        # --- FFT branch\n        tok_f = self.fft(s)\n        out_f = self.enc_fft(torch.cat([self.cls_fft.expand(B,-1,-1), tok_f],1))[:,0]\n        # --- Fuse & reshape\n        fused = torch.cat([out_r, out_f], dim=1)      # (B,2D)\n        flat  = self.fuse(fused)                      # (B,4900)\n        return flat.view(B,70,70)\n\n\nmodel = DualBranchModel().to(DEVICE)\nLR_INIT   = 1e-4        # keep the same starting LR\nLR_MIN    = 2e-5        # final LR after anneal\nEPOCHS    = 30  \nWEIGHT_DECAY = 5e-4    \n# train longer so schedule matters\n# ------------------------------------------------------------\n\n# ---------------- Optimiser & Scheduler ---------------------\nopt = torch.optim.AdamW(model.parameters(),\n                        lr=LR_INIT,\n                        weight_decay=WEIGHT_DECAY)\n\n# Cosine decay from epoch 0 → EPOCHS-1, floor at LR_MIN\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    opt, T_max=EPOCHS, eta_min=LR_MIN)\nmae   = nn.L1Loss()\n\n# ---------------- Training ----------------\nfor ep in tqdm(range(1,EPOCHS+1)):\n    model.train(); tr=0\n    ix = 0\n    for s,v in dl_tr:\n        ix += 1\n        if ix % 100 == 0:\n            print(f\"{ix}/{len(dl_tr)} in epoch {ep}\")\n        s,v = s.to(DEVICE), v.to(DEVICE)\n        opt.zero_grad(); L = mae(model(s), v); L.backward(); opt.step()\n        tr += L.item()*s.size(0)\n    tr_mae = tr/len(train_idx)\n    scheduler.step()\n    model.eval(); va=0\n    with torch.no_grad():\n        for s,v in dl_va:\n            s,v = s.to(DEVICE), v.to(DEVICE)\n            va += mae(model(s), v).item()*s.size(0)\n    va_mae = va/len(val_idx)\n    print(f\"Epoch {ep:2d}/{EPOCHS}  \"\n          f\"Train MAE = {tr_mae*V_SCALE:.1f} m/s   \"\n          f\"Val MAE = {va_mae*V_SCALE:.1f} m/s\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip uninstall -y torch torchvision torchaudio","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n!pip install --no-cache-dir \\\n    torch==2.1.0+cu121 \\\n    torchvision==0.16.0+cu121 \\\n    torchaudio==2.1.0 \\\n    --index-url https://download.pytorch.org/whl/cu121\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nprint(torch.__version__) ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys, types\n\n# Create a fake torch._functorch package\nfake_ft = types.ModuleType(\"torch._functorch\")\n# Create submodules\nfake_eager = types.ModuleType(\"torch._functorch.eager_transforms\")\n# stub out the missing grad_and_value\nfake_eager.grad_and_value = lambda *args, **kwargs: None\nfake_dep = types.ModuleType(\"torch._functorch.deprecated\")\n# stub setup_docs to a no-op\nfake_dep.setup_docs = lambda *args, **kwargs: None\n\n# Install into sys.modules under all the names PyTorch will look up\nsys.modules[\"torch._functorch\"] = fake_ft\nsys.modules[\"torch._functorch.eager_transforms\"] = fake_eager\nsys.modules[\"torch._functorch.deprecated\"] = fake_dep\n\n# Now it’s safe to import torch\nimport torch\nprint(\"torch version:\", torch.__version__)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/usr/bin/env python3\n\"\"\"\ndual_branch_metaformer_dense.py\n================================\nDual-branch **MetaFormer** encoder + PixelShuffle-U-Net refine.\n\nKey differences from the previous script\n----------------------------------------\n•  encoder blocks = 6-layer **PoolFormer** (MetaFormer skeleton, pool mixer)\n•  stride reduced to **4 × 2**  →  ~8 k tokens / gather\n•  stems keep token-dropout 0.10\n•  PixelShuffle×2 + mini-U-Net refine (base 32) still used\n•  Cosine LR 30 epochs, weight-decay 5e-4\nThis fits on a 16 GB GPU at batch = 2 with AMP.\n\nDirectory layout unchanged – loader uses mmap index.\n\n\"\"\"\nimport torch._dynamo\ntorch._dynamo.reset()\ntorch._dynamo.disable()\nimport torch\nfrom torch.optim.lr_scheduler import LambdaLR\n\nimport os, re, math, numpy as np, torch, torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom sklearn.model_selection import train_test_split\nfrom numpy.lib.stride_tricks import sliding_window_view\nfrom torch.amp import autocast, GradScaler\n# ---------------- hyper-parameters ----------------\nBASE_DIR = \"/kaggle/input/waveform-inversion/train_samples\"\nPATCH_T, PATCH_X = 32, 16\nSTR_T,   STR_X   = 4,  2  \nNX_TOK = ((70 - PATCH_X) // STR_X) + 1# DENSE TOKENS\nEMB_D    = 256\nDEPTH    = 6                       # MetaFormer blocks per branch\nDROP     = 0.1                     # token dropout\nLR_INIT  = 1e-4\nLR_MIN   = 2e-5\nEPOCHS   = 30\nBATCH    = 2                       # keep small for memory\nWEIGHT_DECAY = 5e-4\nVAL_FRAC = 0.2\nV_SCALE  = 4500.0\nDEVICE   = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ntorch.cuda.manual_seed(42); np.random.seed(42)\n# --------------------------------------------------\n\n# -------- helper: pair files (same as before) -----\npat_data  = re.compile(r'^(data|seis)[_\\-]?')\npat_model = re.compile(r'^(model|vel)[_\\-]?')\n\ndef list_pairs(fam_dir):\n    data_dir, model_dir = os.path.join(fam_dir,\"data\"), os.path.join(fam_dir,\"model\")\n    if os.path.isdir(data_dir) and os.path.isdir(model_dir):          # split layout\n        d = sorted(f for f in os.listdir(data_dir)  if f.endswith(\".npy\"))\n        m = sorted(f for f in os.listdir(model_dir) if f.endswith(\".npy\"))\n        return [(os.path.join(data_dir, d[i]), os.path.join(model_dir, m[i])) for i in range(len(d))]\n    # merged layout\n    files = [f for f in os.listdir(fam_dir) if f.endswith(\".npy\")]\n    seis, vel = {}, {}\n    for f in files:\n        if f.startswith((\"data\",\"seis\")):\n            seis[pat_data.sub(\"\",f)]  = f\n        elif f.startswith((\"model\",\"vel\")):\n            vel[pat_model.sub(\"\",f)]  = f\n    return [(os.path.join(fam_dir,seis[k]), os.path.join(fam_dir,vel[k])) for k in sorted(set(seis)&set(vel))]\n# --------------------------------------------------\n\n# -------- mmap index dataset (unchanged) ----------\ndef build_index(base):\n    idx=[]\n    for fam in sorted(os.listdir(base)):\n        fam_dir=os.path.join(base,fam)\n        if not os.path.isdir(fam_dir): continue\n        for sfile, vfile in list_pairs(fam_dir):\n            for i in range(500):\n                idx.append((sfile,vfile,i))\n    return idx\n\nindex_all = build_index(BASE_DIR)\ntr_ids, va_ids = train_test_split(np.arange(len(index_all)),\n                                  test_size=VAL_FRAC, random_state=42, shuffle=True)\n\nclass VelDS(Dataset):\n    def __init__(self, idx_list): self.idx = idx_list\n    def __len__(self): return len(self.idx)\n    def __getitem__(self,k):\n        sfile,vfile,i = self.idx[k]\n        s  = np.load(sfile, mmap_mode='r')[i].copy().astype(np.float32)   # (5,1000,70)\n        v  = np.load(vfile, mmap_mode='r')[i]; v = v.squeeze().copy().astype(np.float32)/V_SCALE\n        return torch.from_numpy(s), torch.from_numpy(v)\n\ndl_tr = DataLoader(VelDS([index_all[i] for i in tr_ids]), batch_size=BATCH,\n                   shuffle=True, num_workers=4, pin_memory=True)\ndl_va = DataLoader(VelDS([index_all[i] for i in va_ids]), batch_size=BATCH,\n                   shuffle=False,num_workers=4, pin_memory=True)\n# --------------------------------------------------\n\n# ---------------- token stems --------------------\ndef split_fft(x):\n    F = torch.fft.rfft2(x, dim=(-2,-1))\n    return torch.cat([F.real, F.imag, torch.log(torch.abs(F)+1e-6)],1)\n\nclass StemRaw(nn.Module):\n    def __init__(self, emb_d):\n        super().__init__()\n        self.bn = nn.Identity()\n        self.ln = nn.LayerNorm(emb_d)\n        self.conv = nn.Conv2d(5, emb_d, (PATCH_T,PATCH_X), stride=(STR_T,STR_X))\n        self.drop = nn.Dropout(DROP)\n    def forward(self,x):\n        f = self.conv(x)                      # (B, emb_d, T, X)\n        tok = f.flatten(2).transpose(1,2)     # (B, Ntok, emb_d)\n        tok = self.ln(tok)                    # layer-norm per token\n        return self.drop(tok)\n\nclass StemFFT(nn.Module):\n    def __init__(self, emb_d):\n        super().__init__()\n        self.bn = nn.Identity()\n        self.ln = nn.LayerNorm(emb_d)\n        self.conv = nn.Conv2d(15, emb_d, (PATCH_T,PATCH_X), stride=(STR_T,STR_X))\n        self.drop = nn.Dropout(DROP)\n    def forward(self, x):\n        # x: (B,5,1000,70)\n        spec = split_fft(x)  # (B,15,1000,?) — check dims\n        f = self.conv(spec)\n        tok = f.flatten(2).transpose(1,2)\n        return self.drop(tok)\n# --------------------------------------------------\n\n# -------------- MetaFormer (PoolMixer) ------------\nNX_TOK = ((70 - PATCH_X) // STR_X) + 1    # = 28  (constant)\n\nclass PoolMixer(nn.Module):\n    def __init__(self, C, W_tok):\n        super().__init__()\n        self.W = W_tok                         # patch-grid width\n\n    def forward(self, x):                      # x (B, 1+N, C)\n        cls, tok = x[:, :1], x[:, 1:]\n        B, Np, C = tok.shape\n        W = self.W\n        H = Np // W\n        assert H * W == Np, f\"grid {H}×{W}!={Np}\"\n        y = tok.transpose(1, 2).reshape(B, C, H, W)\n        y = F.avg_pool2d(y, 3, 1, 1) - y\n        tok = tok + y.flatten(2).transpose(1, 2)\n        return torch.cat([cls, tok], 1)\n\nclass MetaBlock(nn.Module):\n    def __init__(self, C, W_tok, ff_mult=4):\n        super().__init__()\n        self.norm1 = nn.LayerNorm(C)\n        self.mix   = PoolMixer(C, W_tok)\n        self.norm2 = nn.LayerNorm(C)\n        self.ffn   = nn.Sequential(\n            nn.Linear(C, ff_mult*C),\n            nn.GELU(),\n            nn.Linear(ff_mult*C, C)\n        )\n    def forward(self, x):\n        x = x + self.mix(self.norm1(x))\n        return x + self.ffn(self.norm2(x))\n\nclass MetaEncoder(nn.Module):\n    def __init__(self, depth: int, C: int, W_tok: int):\n        super().__init__()\n        self.blocks = nn.ModuleList(\n            [MetaBlock(C, W_tok) for _ in range(depth)]\n        )\n    def forward(self, x):\n        for blk in self.blocks:\n            x = blk(x)\n        return x\n# --------------------------------------------------\n\n# -------------- PixelShuffle + mini-U-Net ---------\nclass ShuffleUNetRefine(nn.Module):\n    def __init__(self, base=32):\n        super().__init__()\n        self.pre = nn.Conv2d(1, base*4, 1)\n        self.shuffle = nn.PixelShuffle(2)          # 70×70 → 140×140\n\n        self.enc1 = nn.Sequential(\n            nn.Conv2d(base, base,3,padding=1), nn.ReLU(True),\n            nn.Conv2d(base, base,3,padding=1), nn.ReLU(True))\n        self.pool1= nn.MaxPool2d(2)                # 70×70\n        self.enc2 = nn.Sequential(\n            nn.Conv2d(base, base*2,3,padding=1), nn.ReLU(True),\n            nn.Conv2d(base*2, base*2,3,padding=1), nn.ReLU(True))\n        self.pool2= nn.MaxPool2d(2)                # 35×35\n        self.enc3 = nn.Sequential(\n            nn.Conv2d(base*2, base*4,3,padding=1), nn.ReLU(True),\n            nn.Conv2d(base*4, base*4,3,padding=1), nn.ReLU(True))\n\n        self.up2  = nn.ConvTranspose2d(base*4, base*2, 2,2)\n        self.dec2 = nn.Sequential(\n            nn.Conv2d(base*4, base*2,3,padding=1), nn.ReLU(True),\n            nn.Conv2d(base*2, base*2,3,padding=1), nn.ReLU(True))\n        self.up1  = nn.ConvTranspose2d(base*2, base, 2,2)\n        self.dec1 = nn.Sequential(\n            nn.Conv2d(base*2, base,3,padding=1), nn.ReLU(True),\n            nn.Conv2d(base, base,3,padding=1), nn.ReLU(True))\n        self.outc = nn.Conv2d(base,1,1)\n\n    def forward(self, c):                          # (B,1,70,70)\n        x = self.shuffle(self.pre(c))              # (B,base,140,140)\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool1(e1))\n        e3 = self.enc3(self.pool2(e2))\n        d2 = self.up2(e3); d2 = self.dec2(torch.cat([d2,e2],1))\n        d1 = self.up1(d2); d1 = self.dec1(torch.cat([d1,e1],1))\n        hi = self.outc(d1)                         # 140×140\n        return F.avg_pool2d(hi,2).squeeze(1)       # (B,70,70)\n# --------------------------------------------------\n\n# -------------- Dual-branch MetaFormer model ------\nclass DualBranchMeta(nn.Module):\n    def __init__(self):\n        # same init as before, but you can drop cls_r and cls_f if unused\n        super().__init__()\n        self.raw = StemRaw(EMB_D)\n        self.fft = StemFFT(EMB_D)\n        # remove or ignore cls tokens:\n        # self.cls_r = nn.Parameter(torch.zeros(1,1,EMB_D))\n        # self.cls_f = nn.Parameter(torch.zeros(1,1,EMB_D))\n\n        W_RAW = ((70 - PATCH_X)  // STR_X) + 1    # 28\n        W_FFT = ((36 - PATCH_X)  // STR_X) + 1    # 11\n        self.enc_raw = MetaEncoder(DEPTH, EMB_D, W_RAW)\n        self.enc_fft = MetaEncoder(DEPTH, EMB_D, W_FFT)\n\n        # fuse and refine as before\n        self.fuse = nn.Sequential(\n            nn.LayerNorm(2*EMB_D),\n            nn.Linear(2*EMB_D, 4*EMB_D),\n            nn.GELU(),\n            nn.Linear(4*EMB_D, 70*70)\n        )\n        self.refine = ShuffleUNetRefine(base=32)\n\n    def forward(self, s):\n        B = s.size(0)\n        # 1) get token embeddings from stems: shape (B, N_raw, EMB_D)\n        r_tok = self.raw(s)\n        f_tok = self.fft(s)\n\n        # 2) run through encoder: produce per-token outputs (B, 1+N, EMB_D)? \n        #    But since we removed CLS, we can prepend a small learned vector if desired,\n        #    or simply pool after encoding.\n        # Here: skip CLS; encode raw tokens directly:\n        r_enc = self.enc_raw(torch.cat([torch.zeros(B,1,EMB_D,device=s.device), r_tok], dim=1))\n        # or better: pool *before* encoding: mean pool, then encode a single token?\n        # Simpler: encode all tokens then mean-pool output tokens:\n        r_out = self.enc_raw(torch.cat([torch.zeros(B,1,EMB_D,device=s.device), r_tok], dim=1))  # (B,1+N,EMB_D)\n        r_cls = r_out[:,1:].mean(dim=1)  # mean over the raw token outputs\n\n        f_out = self.enc_fft(torch.cat([torch.zeros(B,1,EMB_D,device=s.device), f_tok], dim=1))\n        f_cls = f_out[:,1:].mean(dim=1)\n\n        # 3) fuse and refine\n        fused = torch.cat([r_cls, f_cls], dim=1)  # (B, 2*EMB_D)\n        coarse = self.fuse(fused).view(B,1,70,70)\n        return self.refine(coarse)\n# --------------------------------------------------\n\n\nimport math, torch.optim.lr_scheduler as tls\n\nDEVICE        = torch.device(\"cuda\")\nLR_CONST      = 1e-4\nEPOCHS        = 30\nACCUM_STEPS   = 1\nCLIP_NORM     = 100.0\nDROP          = 0.0\nV_SCALE       = 4500.0\n\nwarmup_steps = 2000\n# instantiate\nmodel = DualBranchMeta().to(DEVICE)\noptimizer   = torch.optim.AdamW(model.parameters(), lr=LR_CONST, weight_decay=1e-2)\n\ndef lr_warmup_cosine(step):\n    if step < warmup_steps:\n        return float(step + 1) / warmup_steps\n    else:\n        progress = float(step - warmup_steps) / float(total_steps - warmup_steps)\n        # cosine from 1 → 0\n        return 0.5 * (1 + math.cos(math.pi * progress))\nscheduler = LambdaLR(optimizer, lr_warmup_cosine)\n\nscaler= GradScaler(\"cuda\")\nmae = nn.L1Loss(reduction='mean')\nsteps_per_epoch = len(dl_tr)\ntotal_steps = EPOCHS * steps_per_epoch\nglobal_step = 0\nfor ep in range(EPOCHS):\n    model.train()\n    running_loss = 0.0\n    for step, (seis, vel) in enumerate(dl_tr):\n        seis, vel = seis.to(DEVICE), vel.to(DEVICE)\n        optimizer.zero_grad()\n\n        pred = model(seis)\n        loss = mae(pred, vel) * V_SCALE\n        loss.backward()\n        if global_step % 1000 == 0:\n            grads = [p.grad.norm() for p in model.parameters() if p.grad is not None]\n            raw_norm = torch.norm(torch.stack(grads), 2).item()\n            print(f\"Raw grad norm before clipping: {raw_norm:.2f}\")\n        # now clip:\n\n        torch.nn.utils.clip_grad_norm_(model.parameters(), CLIP_NORM)\n        optimizer.step()\n\n        # 4) Step the scheduler\n        scheduler.step()\n        global_step += 1\n\n        running_loss += loss.item() * seis.size(0)\n\n        # (Optional) print every N updates\n        if global_step % 1000 == 0:\n            current_lr = optimizer.param_groups[0]['lr']\n            print(f\"GlobalStep {global_step:5d}  LR={current_lr:.2e}  loss={loss.item():.1f}\")\n\n    train_mae = running_loss / len(dl_tr.dataset)\n    model.eval()\n    val_loss = 0.0\n    with torch.no_grad():\n        for seis, vel in dl_va:\n            seis, vel = seis.to(DEVICE), vel.to(DEVICE)\n            pred = model(seis)\n            loss_v = mae(pred, vel) * V_SCALE\n            val_loss += loss_v.item() * seis.size(0)\n    val_mae = val_loss / len(dl_va.dataset)\n    print(f\"Epoch {ep:02d}/{EPOCHS}  \"\n          f\"Train MAE: {train_mae:6.1f} m/s   \"\n          f\"Val MAE:   {val_mae:6.1f} m/s\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for name, mod in model.named_modules():\n    if isinstance(mod, nn.GroupNorm):\n        dev = next(mod.parameters()).device\n        print(f\"{name:30s} → {type(mod).__name__:10s} on {dev}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ix += 1\n        if ix % 200 == 0:\n            print(f\"{ix}/{len(dl_tr)} at epoch {ep}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for name, param in model.named_parameters():\n    if 'bn' in name or 'norm' in name:\n        print(name, param.device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim.lr_scheduler import LambdaLR\nfrom torch.utils.data import DataLoader\nimport numpy as np\nimport re\n\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset\nfrom sklearn.model_selection import train_test_split\nfrom numpy.lib.stride_tricks import sliding_window_view\nfrom torch.amp import autocast, GradScaler\n\n# ───────────────────────── Hyperparameters ─────────────────────────\nPATCH_T, PATCH_X = 32, 16\nSTR_T,   STR_X   = 4,  2  \nNX_TOK = ((70 - PATCH_X) // STR_X) + 1# DENSE TOKENS\nEMB_D    = 256\nDEPTH    = 6                       # MetaFormer blocks per branch\nDROP     = 0.1 \n\n\nBASE_DIR = \"/kaggle/input/waveform-inversion/train_samples\"\nDEVICE       = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nEPOCHS       = 30\nBATCH   = 2            # adjust to GPU memory\nWEIGHT_DECAY = 5e-4\n\n# LR schedule parameters\nLR_START     = 1e-6         # initial LR at step 0\nLR_BASE      = 1e-4         # target LR after warm-up\nLR_MIN       = 2e-5         # final LR after cosine anneal\nwarmup_steps = 2000         # number of optimizer steps to ramp LR_START → LR_BASE\n\nCLIP_NORM    = 200.0        # gradient clipping threshold\nV_SCALE      = 4500.0       # velocity scaling: vel_norm = vel_raw / V_SCALE\nVAL_FRAC = 0.2\n# Loss weights and Huber beta\nbeta_norm    = 0.02         # in normalized units; corresponds ~0.02 * V_SCALE ≈ 90 m/s\nw_full_base  = 1.0\nw_ds_base    = 0.5\nw_tv_base    = 1e-3\n# ──────────────────────────────────────────────────────────────────\n\n\npat_data  = re.compile(r'^(data|seis)[_\\-]?')\npat_model = re.compile(r'^(model|vel)[_\\-]?')\n\ndef list_pairs(fam_dir):\n    data_dir, model_dir = os.path.join(fam_dir,\"data\"), os.path.join(fam_dir,\"model\")\n    if os.path.isdir(data_dir) and os.path.isdir(model_dir):          # split layout\n        d = sorted(f for f in os.listdir(data_dir)  if f.endswith(\".npy\"))\n        m = sorted(f for f in os.listdir(model_dir) if f.endswith(\".npy\"))\n        return [(os.path.join(data_dir, d[i]), os.path.join(model_dir, m[i])) for i in range(len(d))]\n    # merged layout\n    files = [f for f in os.listdir(fam_dir) if f.endswith(\".npy\")]\n    seis, vel = {}, {}\n    for f in files:\n        if f.startswith((\"data\",\"seis\")):\n            seis[pat_data.sub(\"\",f)]  = f\n        elif f.startswith((\"model\",\"vel\")):\n            vel[pat_model.sub(\"\",f)]  = f\n    return [(os.path.join(fam_dir,seis[k]), os.path.join(fam_dir,vel[k])) for k in sorted(set(seis)&set(vel))]\n# --------------------------------------------------\n\n# -------- mmap index dataset (unchanged) ----------\ndef build_index(base):\n    idx=[]\n    for fam in sorted(os.listdir(base)):\n        fam_dir=os.path.join(base,fam)\n        if not os.path.isdir(fam_dir): continue\n        for sfile, vfile in list_pairs(fam_dir):\n            for i in range(500):\n                idx.append((sfile,vfile,i))\n    return idx\n\nindex_all = build_index(BASE_DIR)\ntr_ids, va_ids = train_test_split(np.arange(len(index_all)),\n                                  test_size=VAL_FRAC, random_state=42, shuffle=True)\n\nclass VelDS(Dataset):\n    def __init__(self, idx_list): self.idx = idx_list\n    def __len__(self): return len(self.idx)\n    def __getitem__(self,k):\n        sfile,vfile,i = self.idx[k]\n        s  = np.load(sfile, mmap_mode='r')[i].copy().astype(np.float32)   # (5,1000,70)\n        v  = np.load(vfile, mmap_mode='r')[i]; v = v.squeeze().copy().astype(np.float32)/V_SCALE\n        return torch.from_numpy(s), torch.from_numpy(v)\n\ndl_tr = DataLoader(VelDS([index_all[i] for i in tr_ids]), batch_size=BATCH,\n                   shuffle=True, num_workers=4, pin_memory=True)\ndl_va = DataLoader(VelDS([index_all[i] for i in va_ids]), batch_size=BATCH,\n                   shuffle=False,num_workers=4, pin_memory=True)\n# --------------------------------------------------\n\n# ---------------- token stems --------------------\ndef split_fft(x):\n    F = torch.fft.rfft2(x, dim=(-2,-1))\n    return torch.cat([F.real, F.imag, torch.log(torch.abs(F)+1e-6)],1)\n\nclass StemRaw(nn.Module):\n    def __init__(self, emb_d):\n        super().__init__()\n        self.bn = nn.Identity()\n        self.ln = nn.LayerNorm(emb_d)\n        self.conv = nn.Conv2d(5, emb_d, (PATCH_T,PATCH_X), stride=(STR_T,STR_X))\n        self.drop = nn.Dropout(DROP)\n    def forward(self,x):\n        f = self.conv(x)                      # (B, emb_d, T, X)\n        tok = f.flatten(2).transpose(1,2)     # (B, Ntok, emb_d)\n        tok = self.ln(tok)                    # layer-norm per token\n        return self.drop(tok)\n\nclass StemFFT(nn.Module):\n    def __init__(self, emb_d):\n        super().__init__()\n        self.bn = nn.Identity()\n        self.ln = nn.LayerNorm(emb_d)\n        self.conv = nn.Conv2d(15, emb_d, (PATCH_T,PATCH_X), stride=(STR_T,STR_X))\n        self.drop = nn.Dropout(DROP)\n    def forward(self, x):\n        # x: (B,5,1000,70)\n        spec = split_fft(x)  # (B,15,1000,?) — check dims\n        f = self.conv(spec)\n        tok = f.flatten(2).transpose(1,2)\n        return self.drop(tok)\n# --------------------------------------------------\n\n# -------------- MetaFormer (PoolMixer) ------------\nNX_TOK = ((70 - PATCH_X) // STR_X) + 1    # = 28  (constant)\n\nclass PoolMixer(nn.Module):\n    def __init__(self, C, W_tok):\n        super().__init__()\n        self.W = W_tok                         # patch-grid width\n\n    def forward(self, x):                      # x (B, 1+N, C)\n        cls, tok = x[:, :1], x[:, 1:]\n        B, Np, C = tok.shape\n        W = self.W\n        H = Np // W\n        assert H * W == Np, f\"grid {H}×{W}!={Np}\"\n        y = tok.transpose(1, 2).reshape(B, C, H, W)\n        y = F.avg_pool2d(y, 3, 1, 1) - y\n        tok = tok + y.flatten(2).transpose(1, 2)\n        return torch.cat([cls, tok], 1)\n\nclass MetaBlock(nn.Module):\n    def __init__(self, C, W_tok, ff_mult=4):\n        super().__init__()\n        self.norm1 = nn.LayerNorm(C)\n        self.mix   = PoolMixer(C, W_tok)\n        self.norm2 = nn.LayerNorm(C)\n        self.ffn   = nn.Sequential(\n            nn.Linear(C, ff_mult*C),\n            nn.GELU(),\n            nn.Linear(ff_mult*C, C)\n        )\n    def forward(self, x):\n        x = x + self.mix(self.norm1(x))\n        return x + self.ffn(self.norm2(x))\n\nclass MetaEncoder(nn.Module):\n    def __init__(self, depth: int, C: int, W_tok: int):\n        super().__init__()\n        self.blocks = nn.ModuleList(\n            [MetaBlock(C, W_tok) for _ in range(depth)]\n        )\n    def forward(self, x):\n        for blk in self.blocks:\n            x = blk(x)\n        return x\n# --------------------------------------------------\n\n# -------------- PixelShuffle + mini-U-Net ---------\nclass ShuffleUNetRefine(nn.Module):\n    def __init__(self, base=32):\n        super().__init__()\n        self.pre = nn.Conv2d(1, base*4, 1)\n        self.shuffle = nn.PixelShuffle(2)          # 70×70 → 140×140\n\n        self.enc1 = nn.Sequential(\n            nn.Conv2d(base, base,3,padding=1), nn.ReLU(True),\n            nn.Conv2d(base, base,3,padding=1), nn.ReLU(True))\n        self.pool1= nn.MaxPool2d(2)                # 70×70\n        self.enc2 = nn.Sequential(\n            nn.Conv2d(base, base*2,3,padding=1), nn.ReLU(True),\n            nn.Conv2d(base*2, base*2,3,padding=1), nn.ReLU(True))\n        self.pool2= nn.MaxPool2d(2)                # 35×35\n        self.enc3 = nn.Sequential(\n            nn.Conv2d(base*2, base*4,3,padding=1), nn.ReLU(True),\n            nn.Conv2d(base*4, base*4,3,padding=1), nn.ReLU(True))\n\n        self.up2  = nn.ConvTranspose2d(base*4, base*2, 2,2)\n        self.dec2 = nn.Sequential(\n            nn.Conv2d(base*4, base*2,3,padding=1), nn.ReLU(True),\n            nn.Conv2d(base*2, base*2,3,padding=1), nn.ReLU(True))\n        self.up1  = nn.ConvTranspose2d(base*2, base, 2,2)\n        self.dec1 = nn.Sequential(\n            nn.Conv2d(base*2, base,3,padding=1), nn.ReLU(True),\n            nn.Conv2d(base, base,3,padding=1), nn.ReLU(True))\n        self.outc = nn.Conv2d(base,1,1)\n\n    def forward(self, c):                          # (B,1,70,70)\n        x = self.shuffle(self.pre(c))              # (B,base,140,140)\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool1(e1))\n        e3 = self.enc3(self.pool2(e2))\n        d2 = self.up2(e3); d2 = self.dec2(torch.cat([d2,e2],1))\n        d1 = self.up1(d2); d1 = self.dec1(torch.cat([d1,e1],1))\n        hi = self.outc(d1)                         # 140×140\n        return F.avg_pool2d(hi,2).squeeze(1)       # (B,70,70)\n# --------------------------------------------------\n\n# -------------- Dual-branch MetaFormer model ------\nclass DualBranchMeta(nn.Module):\n    def __init__(self):\n        # same init as before, but you can drop cls_r and cls_f if unused\n        super().__init__()\n        self.raw = StemRaw(EMB_D)\n        self.fft = StemFFT(EMB_D)\n        # remove or ignore cls tokens:\n        # self.cls_r = nn.Parameter(torch.zeros(1,1,EMB_D))\n        # self.cls_f = nn.Parameter(torch.zeros(1,1,EMB_D))\n\n        W_RAW = ((70 - PATCH_X)  // STR_X) + 1    # 28\n        W_FFT = ((36 - PATCH_X)  // STR_X) + 1    # 11\n        self.enc_raw = MetaEncoder(DEPTH, EMB_D, W_RAW)\n        self.enc_fft = MetaEncoder(DEPTH, EMB_D, W_FFT)\n\n        # fuse and refine as before\n        self.fuse = nn.Sequential(\n            nn.LayerNorm(2*EMB_D),\n            nn.Linear(2*EMB_D, 4*EMB_D),\n            nn.GELU(),\n            nn.Linear(4*EMB_D, 70*70)\n        )\n        self.refine = ShuffleUNetRefine(base=32)\n\n    def forward(self, s):\n        B = s.size(0)\n        # 1) get token embeddings from stems: shape (B, N_raw, EMB_D)\n        r_tok = self.raw(s)\n        f_tok = self.fft(s)\n\n        # 2) run through encoder: produce per-token outputs (B, 1+N, EMB_D)? \n        #    But since we removed CLS, we can prepend a small learned vector if desired,\n        #    or simply pool after encoding.\n        # Here: skip CLS; encode raw tokens directly:\n        r_enc = self.enc_raw(torch.cat([torch.zeros(B,1,EMB_D,device=s.device), r_tok], dim=1))\n        # or better: pool *before* encoding: mean pool, then encode a single token?\n        # Simpler: encode all tokens then mean-pool output tokens:\n        r_out = self.enc_raw(torch.cat([torch.zeros(B,1,EMB_D,device=s.device), r_tok], dim=1))  # (B,1+N,EMB_D)\n        r_cls = r_out[:,1:].mean(dim=1)  # mean over the raw token outputs\n\n        f_out = self.enc_fft(torch.cat([torch.zeros(B,1,EMB_D,device=s.device), f_tok], dim=1))\n        f_cls = f_out[:,1:].mean(dim=1)\n\n        # 3) fuse and refine\n        fused = torch.cat([r_cls, f_cls], dim=1)  # (B, 2*EMB_D)\n        coarse = self.fuse(fused).view(B,1,70,70)\n        return self.refine(coarse)\n# --------------------------------------------------\n\n\n\nmodel = DualBranchMeta().to(DEVICE)\n\n# 1) Compute dataset mean normalized velocity for bias init\n#    This requires one pass over dl_tr (or a subset). Adjust if dataset is huge.\nprint(\"Computing mean normalized velocity for bias init...\")\nsum_vel = 0.0\ncount = 0\nwith torch.no_grad():\n    for _, vel in dl_tr:\n        # vel: (B, H, W), normalized in [0,1]\n        sum_vel += vel.sum().item()\n        count += vel.numel()\nmean_norm = sum_vel / count\nprint(f\"  mean normalized velocity ≈ {mean_norm:.4f} (=> raw ≈ {mean_norm*V_SCALE:.1f} m/s)\")\n\n# 2) Initialize final-layer bias so initial output ≈ mean_norm\n#    Assumes the final conv in refine is named `outc` and has a bias parameter.\n#    Adjust if your final layer differs.\nimport math\n# For a sigmoid activation: bias b such that sigmoid(b) = mean_norm => b = log(mean_norm/(1-mean_norm))\n# If your model’s final layer is linear (no activation), you may instead initialize bias directly:\n#   model.refine.outc.bias.data.fill_(mean_norm)\n# Here we assume a Sigmoid is used; if not, adapt accordingly.\ntry:\n    b_init = math.log(mean_norm / (1.0 - mean_norm))\n    model.refine.outc.bias.data.fill_(b_init)\n    print(f\"Initialized final bias to {b_init:.4f} for sigmoid→{mean_norm:.4f}\")\nexcept Exception:\n    # If no sigmoid, set direct bias to mean_norm\n    try:\n        model.refine.outc.bias.data.fill_(mean_norm)\n        print(f\"Initialized final bias directly to normalized mean {mean_norm:.4f}\")\n    except Exception:\n        print(\"Warning: could not initialize final bias automatically; please check final layer.\")\n\n# 3) Setup optimizer and LR scheduler\noptimizer = torch.optim.AdamW(model.parameters(), lr=LR_BASE, weight_decay=WEIGHT_DECAY)\n\nsteps_per_epoch = len(dl_tr)\ntotal_steps = EPOCHS * steps_per_epoch\n\ndef lr_schedule(step):\n    # Linear warm-up from LR_START → LR_BASE\n    if step < warmup_steps:\n        frac = float(step + 1) / float(warmup_steps)\n        lr = LR_START + frac * (LR_BASE - LR_START)\n    else:\n        # Cosine anneal from LR_BASE → LR_MIN over remaining steps\n        progress = float(step - warmup_steps) / float(max(1, total_steps - warmup_steps))\n        cosine = 0.5 * (1.0 + math.cos(math.pi * min(progress, 1.0)))\n        lr = LR_MIN + (LR_BASE - LR_MIN) * cosine\n    # LambdaLR expects multiplier relative to LR_BASE\n    return lr / LR_BASE\n\nscheduler = LambdaLR(optimizer, lr_schedule)\n\n# 4) Define loss criterion\ncriterion_data = nn.SmoothL1Loss(beta=beta_norm)  # in normalized units\n\n# 5) Training loop\nglobal_step = 0\nfor ep in range(EPOCHS):\n    model.train()\n    running_loss_full = 0.0\n\n    # Adjust weights schedules per epoch\n    # Multi-scale weight decays from w_ds_base → 0.1*w_ds_base over first 10 epochs:\n    w_ds = w_ds_base * (1.0 - 0.8 * min(ep, 10) / 10.0)\n    # TV weight ramps up from 0 → w_tv_base over first 5 epochs:\n    w_tv = w_tv_base * min(ep / 5.0, 1.0)\n\n    print(f\"\\nEpoch {ep:02d}: w_ds={w_ds:.3f}, w_tv={w_tv:.3e}\")\n\n    for step, (seis, vel) in enumerate(dl_tr):\n        seis, vel = seis.to(DEVICE), vel.to(DEVICE)  # seis: (B,5,1000,70), vel: (B,70,70)\n\n        optimizer.zero_grad()\n\n        # Forward\n        pred = model(seis)  # expected normalized output in [0,1]\n\n        # 5.1 Data term full resolution\n        loss_full = criterion_data(pred, vel) * V_SCALE\n\n        # 5.2 Multi-scale term (downsample ×2)\n        pred_ds = F.avg_pool2d(pred.unsqueeze(1), kernel_size=2, stride=2).squeeze(1)  # (B,35,35)\n        vel_ds  = F.avg_pool2d(vel.unsqueeze(1),  kernel_size=2, stride=2).squeeze(1)\n        loss_ds = criterion_data(pred_ds, vel_ds) * V_SCALE\n\n        # 5.3 TV smoothness term\n        # Absolute differences between neighbors\n        dx = torch.abs(pred[:, :, 1:] - pred[:, :, :-1]).mean()\n        dy = torch.abs(pred[:, 1:, :] - pred[:, :-1, :]).mean()\n        loss_tv = (dx + dy) * V_SCALE\n\n        # 5.4 Combine losses\n        loss = w_full_base * loss_full + w_ds * loss_ds + w_tv * loss_tv\n\n        # Backward + clip + step + scheduler\n        loss.backward()\n        # (Optional) inspect raw grad norm before clipping:\n        if (global_step) % 1000 == 0:\n            grads = [p.grad.norm() for p in model.parameters() if p.grad is not None]\n            raw_norm = torch.norm(torch.stack(grads), 2).item()\n            print(f\"Step {global_step} raw grad norm before clip: {raw_norm:.1f}\")\n        \n        torch.nn.utils.clip_grad_norm_(model.parameters(), CLIP_NORM)\n        optimizer.step()\n        scheduler.step()\n        global_step += 1\n\n        running_loss_full += loss_full.item() * seis.size(0)\n\n        # Diagnostic print every 100 updates\n        \n\n    # Epoch-end metrics\n    train_mae = running_loss_full / len(dl_tr.dataset)\n\n    # Validation (only full-res data term for MAE)\n    model.eval()\n    val_loss_full = 0.0\n    with torch.no_grad():\n        for seis, vel in dl_va:\n            seis, vel = seis.to(DEVICE), vel.to(DEVICE)\n            pred = model(seis)\n            loss_v = criterion_data(pred, vel) * V_SCALE\n            val_loss_full += loss_v.item() * seis.size(0)\n    val_mae = val_loss_full / len(dl_va.dataset)\n\n    current_lr = optimizer.param_groups[0]['lr']\n    print(f\"Epoch {ep:02d}/{EPOCHS}  Train MAE: {train_mae:6.1f} m/s   \"\n          f\"Val MAE: {val_mae:6.1f} m/s   LR={current_lr:.2e}\")\n\n# After training, save model\ntorch.save(model.state_dict(), \"inversion_model_final.pth\")\nprint(\"Model saved to inversion_model_final.pth\")\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/usr/bin/env python3\n\"\"\"\ndual_branch_metaformer_dense.py\n================================\nDual-branch MetaFormer encoder + PixelShuffle-U-Net refine with\nnormalization before refine, lowered initial loss-scale, optional refine freezing,\nand gradient diagnostics.\n\nAdjust hyperparameters (loss_scale_start, warmup_steps, CLIP_NORM, etc.) as needed.\n\"\"\"\n\nimport torch._dynamo\ntorch._dynamo.reset()\ntorch._dynamo.disable()\n\nimport os, re, math, numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim.lr_scheduler import LambdaLR, ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset\nfrom sklearn.model_selection import train_test_split\nfrom torch.amp import autocast, GradScaler\n\n# ---------------- hyper-parameters ----------------\nBASE_DIR = \"/kaggle/input/waveform-inversion/train_samples\"\nPATCH_T, PATCH_X = 32, 16\nSTR_T,   STR_X   = 4,  2  \nEMB_D    = 256\nDEPTH    = 6                       # MetaFormer blocks per branch\nDROP     = 0.1                     # token dropout in stems\n\n# LR schedule\nLR_START = 1e-6                   # for warm-up start\nLR_BASE  = 1e-4\nLR_MIN   = 2e-5\nwarmup_steps = 5000               # steps to ramp LR_START→LR_BASE\nEPOCHS   = 30\nBATCH    = 2                      # keep small for memory\nWEIGHT_DECAY = 5e-4\nVAL_FRAC = 0.2\nV_SCALE  = 4500.0                 # velocity scaling\nDEVICE   = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ntorch.cuda.manual_seed(42); np.random.seed(42)\n\n# Dynamic loss-scale ramp parameters\nloss_scale_start        = 10.0      # lower initial multiplier\nloss_scale_target       = V_SCALE\nloss_scale_warmup_steps = 5000\n\n# Data augmentation\nnoise_std = 0.01     # additive Gaussian noise on seismic input (normalized units)\namp_range = 0.05     # random amplitude scaling ±5%\n\n# Gradient clipping\nCLIP_NORM = 300.0    # global norm clipping threshold\nadaptive_clip_coef = 0.0  # 0 to disable adaptive clipping\n\n# Loss: SmoothL1 with beta in normalized units\nbeta_norm = 0.02\ncriterion_smoothl1 = nn.SmoothL1Loss(beta=beta_norm)\n# --------------------------------------------------\n\n# -------- helper: pair files (same as before) -----\npat_data  = re.compile(r'^(data|seis)[_\\-]?')\npat_model = re.compile(r'^(model|vel)[_\\-]?')\n\ndef list_pairs(fam_dir):\n    data_dir, model_dir = os.path.join(fam_dir,\"data\"), os.path.join(fam_dir,\"model\")\n    if os.path.isdir(data_dir) and os.path.isdir(model_dir):\n        d = sorted(f for f in os.listdir(data_dir)  if f.endswith(\".npy\"))\n        m = sorted(f for f in os.listdir(model_dir) if f.endswith(\".npy\"))\n        return [(os.path.join(data_dir, d[i]), os.path.join(model_dir, m[i])) for i in range(len(d))]\n    files = [f for f in os.listdir(fam_dir) if f.endswith(\".npy\")]\n    seis, vel = {}, {}\n    for f in files:\n        if f.startswith((\"data\",\"seis\")):\n            seis[pat_data.sub(\"\",f)]  = f\n        elif f.startswith((\"model\",\"vel\")):\n            vel[pat_model.sub(\"\",f)]  = f\n    return [(os.path.join(fam_dir,seis[k]), os.path.join(fam_dir,vel[k])) for k in sorted(set(seis)&set(vel))]\n# --------------------------------------------------\n\n# -------- mmap index dataset (unchanged) ----------\ndef build_index(base):\n    idx = []\n    for fam in sorted(os.listdir(base)):\n        fam_dir = os.path.join(base, fam)\n        if not os.path.isdir(fam_dir): continue\n        for sfile, vfile in list_pairs(fam_dir):\n            for i in range(500):\n                idx.append((sfile, vfile, i))\n    return idx\n\nindex_all = build_index(BASE_DIR)\ntr_ids, va_ids = train_test_split(np.arange(len(index_all)),\n                                  test_size=VAL_FRAC, random_state=42, shuffle=True)\n\nclass VelDS(Dataset):\n    def __init__(self, idx_list):\n        self.idx = idx_list\n    def __len__(self):\n        return len(self.idx)\n    def __getitem__(self, k):\n        sfile, vfile, i = self.idx[k]\n        s = np.load(sfile, mmap_mode='r')[i].copy().astype(np.float32)   # (5,1000,70)\n        v = np.load(vfile, mmap_mode='r')[i]\n        v = v.squeeze().copy().astype(np.float32) / V_SCALE\n        return torch.from_numpy(s), torch.from_numpy(v)\n\ndl_tr = DataLoader(VelDS([index_all[i] for i in tr_ids]), batch_size=BATCH,\n                   shuffle=True, num_workers=4, pin_memory=True)\ndl_va = DataLoader(VelDS([index_all[i] for i in va_ids]), batch_size=BATCH,\n                   shuffle=False, num_workers=4, pin_memory=True)\n# --------------------------------------------------\n\n# ---------------- token stems --------------------\ndef split_fft(x):\n    # x: (B,5,1000,70)\n    F = torch.fft.rfft2(x, dim=(-2,-1))  # yields (B,5,1000,36)\n    return torch.cat([F.real, F.imag, torch.log(torch.abs(F)+1e-6)], dim=1)\n\nclass StemRaw(nn.Module):\n    def __init__(self, emb_d):\n        super().__init__()\n        self.ln = nn.LayerNorm(emb_d)\n        self.conv = nn.Conv2d(5, emb_d, (PATCH_T, PATCH_X), stride=(STR_T, STR_X))\n        self.drop = nn.Dropout(DROP)\n    def forward(self, x):\n        f = self.conv(x)                      # (B, emb_d, T', X')\n        tok = f.flatten(2).transpose(1,2)     # (B, Ntok, emb_d)\n        tok = self.ln(tok)\n        return self.drop(tok)\n\nclass StemFFT(nn.Module):\n    def __init__(self, emb_d):\n        super().__init__()\n        self.ln = nn.LayerNorm(emb_d)\n        self.conv = nn.Conv2d(15, emb_d, (PATCH_T, PATCH_X), stride=(STR_T, STR_X))\n        self.drop = nn.Dropout(DROP)\n    def forward(self, x):\n        spec = split_fft(x)                   # (B,15,1000,36)\n        f = self.conv(spec)                   # (B, emb_d, T', X')\n        tok = f.flatten(2).transpose(1,2)     \n        tok = self.ln(tok)\n        return self.drop(tok)\n# --------------------------------------------------\n\n# -------------- MetaFormer (PoolMixer) ------------\nclass PoolMixer(nn.Module):\n    def __init__(self, C, W_tok):\n        super().__init__()\n        self.W = W_tok\n    def forward(self, x):\n        # x: (B, 1+N, C)\n        cls, tok = x[:, :1], x[:, 1:]\n        B, Np, C = tok.shape\n        W = self.W\n        H = Np // W\n        assert H * W == Np, f\"grid {H}×{W}!={Np}\"\n        y = tok.transpose(1, 2).reshape(B, C, H, W)\n        y = F.avg_pool2d(y, 3, 1, 1) - y\n        tok2 = tok + y.flatten(2).transpose(1, 2)\n        return torch.cat([cls, tok2], dim=1)\n\nclass MetaBlock(nn.Module):\n    def __init__(self, C, W_tok, ff_mult=4, residual_scale=1.0):\n        super().__init__()\n        self.norm1 = nn.LayerNorm(C)\n        self.mix   = PoolMixer(C, W_tok)\n        self.norm2 = nn.LayerNorm(C)\n        self.ffn   = nn.Sequential(\n            nn.Linear(C, ff_mult*C),\n            nn.GELU(),\n            nn.Linear(ff_mult*C, C)\n        )\n        self.residual_scale = residual_scale\n    def forward(self, x):\n        y1 = self.mix(self.norm1(x))\n        x = x + self.residual_scale * y1\n        y2 = self.ffn(self.norm2(x))\n        return x + self.residual_scale * y2\n\nclass MetaEncoder(nn.Module):\n    def __init__(self, depth: int, C: int, W_tok: int, residual_scale=1.0):\n        super().__init__()\n        self.blocks = nn.ModuleList(\n            [MetaBlock(C, W_tok, residual_scale=residual_scale) for _ in range(depth)]\n        )\n    def forward(self, x):\n        for blk in self.blocks:\n            x = blk(x)\n        return x\n# --------------------------------------------------\n\n# -------------- PixelShuffle + mini-U-Net ---------\nclass ShuffleUNetRefine(nn.Module):\n    def __init__(self, base=32):\n        super().__init__()\n        self.pre = nn.Conv2d(1, base*4, 1)\n        self.pre_gn = nn.GroupNorm(num_groups=1, num_channels=base*4)\n        self.shuffle = nn.PixelShuffle(2)          \n        self.enc1 = nn.Sequential(\n            nn.Conv2d(base, base, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base),\n            nn.ReLU(True),\n            nn.Conv2d(base, base, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base),\n            nn.ReLU(True),\n        )\n        self.pool1 = nn.MaxPool2d(2)                \n        self.enc2 = nn.Sequential(\n            nn.Conv2d(base, base*2, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base*2),\n            nn.ReLU(True),\n            nn.Conv2d(base*2, base*2, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base*2),\n            nn.ReLU(True),\n        )\n        self.pool2 = nn.MaxPool2d(2)                \n        self.enc3 = nn.Sequential(\n            nn.Conv2d(base*2, base*4, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base*4),\n            nn.ReLU(True),\n            nn.Conv2d(base*4, base*4, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base*4),\n            nn.ReLU(True),\n        )\n        self.up2 = nn.ConvTranspose2d(base*4, base*2, 2,2)\n        self.dec2 = nn.Sequential(\n            nn.Conv2d(base*4, base*2, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base*2),\n            nn.ReLU(True),\n            nn.Conv2d(base*2, base*2, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base*2),\n            nn.ReLU(True),\n        )\n        self.up1 = nn.ConvTranspose2d(base*2, base, 2,2)\n        self.dec1 = nn.Sequential(\n            nn.Conv2d(base*2, base, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base),\n            nn.ReLU(True),\n            nn.Conv2d(base, base, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base),\n            nn.ReLU(True),\n        )\n        self.outc = nn.Conv2d(base, 1, 1)\n\n    def forward(self, c):\n        x = self.shuffle(self.pre_gn(self.pre(c)))       \n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool1(e1))\n        e3 = self.enc3(self.pool2(e2))\n        d2 = self.up2(e3); d2 = self.dec2(torch.cat([d2,e2], dim=1))\n        d1 = self.up1(d2); d1 = self.dec1(torch.cat([d1,e1], dim=1))\n        hi = self.outc(d1)                 \n        return F.avg_pool2d(hi, 2).squeeze(1)  # (B,70,70)\n# --------------------------------------------------\n\n# -------------- Dual-branch MetaFormer model ------\nclass DualBranchMeta(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.raw = StemRaw(EMB_D)\n        self.fft = StemFFT(EMB_D)\n\n        # Compute W_RAW and W_FFT correctly:\n        W_RAW = (70 - PATCH_X)//STR_X + 1   # 28\n        freq_dim = 70//2 + 1               # 36\n        W_FFT = (freq_dim - PATCH_X)//STR_X + 1  # 11\n\n        residual_scale = 1.0\n        self.enc_raw = MetaEncoder(DEPTH, EMB_D, W_RAW, residual_scale=residual_scale)\n        self.enc_fft = MetaEncoder(DEPTH, EMB_D, W_FFT, residual_scale=residual_scale)\n\n        self.fuse = nn.Sequential(\n            nn.LayerNorm(2*EMB_D),\n            nn.Linear(2*EMB_D, 4*EMB_D),\n            nn.GELU(),\n            nn.Linear(4*EMB_D, 70*70)\n        )\n        # normalize coarse before refine\n        self.coarse_norm = nn.GroupNorm(num_groups=1, num_channels=1)\n        self.refine = ShuffleUNetRefine(base=32)\n\n    def forward(self, s):\n        # s: (B,5,1000,70)\n        B = s.size(0)\n        r_tok = self.raw(s)\n        f_tok = self.fft(s)\n        zero_raw = torch.zeros(B, 1, EMB_D, device=s.device, dtype=s.dtype)\n        zero_fft = torch.zeros(B, 1, EMB_D, device=s.device, dtype=s.dtype)\n        r_all = torch.cat([zero_raw, r_tok], dim=1)\n        f_all = torch.cat([zero_fft, f_tok], dim=1)\n        r_enc = self.enc_raw(r_all)\n        f_enc = self.enc_fft(f_all)\n        r_cls = r_enc[:, 1:, :].mean(dim=1)\n        f_cls = f_enc[:, 1:, :].mean(dim=1)\n        fused = torch.cat([r_cls, f_cls], dim=1)\n        coarse = self.fuse(fused).view(B, 1, 70, 70)\n        # normalize coarse\n        coarse = self.coarse_norm(coarse)\n        return self.refine(coarse)\n# --------------------------------------------------\n\n# -------------- Training setup -----------------\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = DualBranchMeta().to(DEVICE)\n\n# 1) Compute mean normalized velocity for bias init\nprint(\"Computing mean normalized velocity for bias init...\")\nsum_vel = 0.0\ncount = 0\nwith torch.no_grad():\n    for _, vel in dl_tr:\n        sum_vel += vel.sum().item()\n        count += vel.numel()\nmean_norm = sum_vel / count\nprint(f\"  mean normalized velocity ≈ {mean_norm:.4f} (raw ≈ {mean_norm*V_SCALE:.1f} m/s)\")\n\n# 2) Initialize final-layer bias\ntry:\n    model.refine.outc.bias.data.fill_(mean_norm)\n    print(f\"Initialized final bias to {mean_norm:.4f} (normalized)\")\nexcept Exception:\n    print(\"Warning: could not initialize final bias automatically; check final layer.\")\n\n# 3) Optimizer and LR scheduler\noptimizer = torch.optim.AdamW(model.parameters(), lr=LR_BASE, weight_decay=WEIGHT_DECAY)\nsteps_per_epoch = len(dl_tr)\ntotal_steps = EPOCHS * steps_per_epoch\n\ndef lr_schedule(step):\n    if step < warmup_steps:\n        return (LR_START + (LR_BASE - LR_START) * (step + 1) / warmup_steps) / LR_BASE\n    else:\n        progress = float(step - warmup_steps) / float(max(1, total_steps - warmup_steps))\n        cosine = 0.5 * (1.0 + math.cos(math.pi * min(progress, 1.0)))\n        lr = LR_MIN + (LR_BASE - LR_MIN) * cosine\n        return lr / LR_BASE\n\nscheduler = LambdaLR(optimizer, lr_schedule)\nscheduler_plateau = ReduceLROnPlateau(optimizer,\n                                      mode='min',\n                                      factor=0.5,\n                                      patience=3,\n                                      min_lr=1e-6,\n                                      verbose=True)\n\ndef adaptive_clip(model, clip_coef=adaptive_clip_coef, eps=1e-6):\n    if clip_coef <= 0: return\n    for p in model.parameters():\n        if p.grad is None: continue\n        param_norm = torch.norm(p.detach())\n        grad_norm = torch.norm(p.grad.detach())\n        max_norm = clip_coef * (param_norm + eps)\n        if grad_norm > max_norm:\n            p.grad.mul_(max_norm / (grad_norm + eps))\n\n# 4) Training loop\n# ... [model, optimizer, scheduler, dl_tr, etc. already defined] ...\n\nACCUM_STEPS = 4   # e.g., batch_size=2 × 4 → effective batch 8\nCLIP_NORM = 300.0\nloss_scale_start = 10.0\nloss_scale_target = V_SCALE\nloss_scale_warmup_steps = 5000\n\nscaler = GradScaler()\nmae = nn.L1Loss(reduction='mean')\nglobal_step = 0\nrunning_mae = 0.0\nbest_val_mae = float('inf')\npatience_counter = 0\n\nfor ep in range(EPOCHS):\n    model.train()\n    running_mae = 0.0\n    print(f\"\\nEpoch {ep:02d}:\")\n    # Optionally freeze refine early\n    if ep < 2:\n        for p in model.refine.parameters(): p.requires_grad=False\n        print(\"  Freezing refine stage this epoch\")\n    else:\n        for p in model.refine.parameters(): p.requires_grad=True\n\n    optimizer.zero_grad()\n    accum_counter = 0\n\n    for step, (seis, vel) in enumerate(dl_tr):\n        seis, vel = seis.to(DEVICE), vel.to(DEVICE)\n        # Data augmentation as before\n        if noise_std > 0:\n            seis = seis + torch.randn_like(seis) * noise_std\n        if amp_range > 0:\n            scale = 1.0 + (torch.rand(seis.size(0),1,1,1,device=DEVICE)*2 -1)*amp_range\n            seis = seis * scale\n\n        with autocast('cuda'):\n            pred = model(seis)  # normalized output\n            # track MAE for reporting\n            with torch.no_grad():\n                batch_mae = mae(pred, vel) * V_SCALE\n                running_mae += batch_mae.item() * seis.size(0)\n\n            # dynamic loss-scale\n            if global_step < loss_scale_warmup_steps:\n                frac = float(global_step+1)/loss_scale_warmup_steps\n                loss_scale = loss_scale_start + frac*(loss_scale_target - loss_scale_start)\n            else:\n                loss_scale = loss_scale_target\n\n            loss = criterion_smoothl1(pred, vel) * loss_scale\n            # divide loss for accumulation\n            loss = loss / ACCUM_STEPS\n\n        scaler.scale(loss).backward()\n        accum_counter += 1\n\n        if accum_counter == ACCUM_STEPS:\n            # unscale & clip once per accumulation\n            scaler.unscale_(optimizer)\n            # optional adaptive clipping\n            # adaptive_clip(model)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), CLIP_NORM)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            # LR scheduler step once per optimizer step\n            scheduler.step()\n            global_step += 1\n            accum_counter = 0\n\n            # Diagnostics every 500 updates:\n            if global_step % 500 == 0:\n                current_lr = optimizer.param_groups[0]['lr']\n                print(f\"[Step {global_step:5d}] LR={current_lr:.2e}, loss_scale={loss_scale:.1f}\")\n                # Log gradient stats:\n                p_vals = [p.grad.norm().item() for p in model.parameters() if p.grad is not None]\n                if p_vals:\n                    t = torch.tensor(p_vals, device='cpu')\n                    print(f\"  Grad norms → overall L2: {t.norm(p=2).item():.3e}, mean: {t.mean().item():.3e}, std: {t.std(unbiased=False).item():.3e}, min: {t.min().item():.3e}, max: {t.max().item():.3e}\")\n                print(f\"  Recent batch MAE: {batch_mae.item():.1f} m/s\")\n\n    # If final accumulation not exactly divisible, step once more:\n    if accum_counter > 0:\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), CLIP_NORM)\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad()\n        scheduler.step()\n        global_step += 1\n\n    train_mae = running_mae / len(dl_tr.dataset)\n\n    # Validation as before:\n    model.eval()\n    val_mae_accum = 0.0\n    with torch.no_grad():\n        for seis, vel in dl_va:\n            seis, vel = seis.to(DEVICE), vel.to(DEVICE)\n            pred = model(seis)\n            val_mae_accum += mae(pred, vel).item() * V_SCALE * seis.size(0)\n    val_mae = val_mae_accum / len(dl_va.dataset)\n\n    print(f\"Epoch {ep:02d}/{EPOCHS}  Train MAE: {train_mae:6.1f} m/s   Val MAE: {val_mae:6.1f} m/s\")\n\n    # LR plateau after warm-up\n    if global_step >= warmup_steps:\n        scheduler_plateau.step(val_mae)\n\n    # Early stopping\n    if val_mae < best_val_mae:\n        best_val_mae = val_mae\n        torch.save(model.state_dict(), \"best_model.pth\")\n        patience_counter = 0\n    else:\n        patience_counter += 1\n    if patience_counter >= 5:\n        print(\"No improvement for 5 epochs, stopping early.\")\n        break\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, re, math, numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim.lr_scheduler import LambdaLR, ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset\nfrom sklearn.model_selection import train_test_split\nfrom torch.amp import autocast, GradScaler\n\nBASE_DIR = \"/kaggle/input/waveform-inversion/train_samples\"\nVAL_FRAC = 0.2\npat_data  = re.compile(r'^(data|seis)[_\\-]?')\npat_model = re.compile(r'^(model|vel)[_\\-]?')\nBASE_DIR = \"/kaggle/input/waveform-inversion/train_samples\"\nPATCH_T, PATCH_X = 32, 16\nSTR_T,   STR_X   = 4,  2  \nEMB_D    = 256\nDEPTH    = 6                       # MetaFormer blocks per branch\nDROP     = 0.1                     # token dropout in stems\n\n# LR schedule\nLR_START = 1e-6                   # for warm-up start\nLR_BASE  = 1e-4\nLR_MIN   = 2e-5\nwarmup_steps = 5000               # steps to ramp LR_START→LR_BASE\nEPOCHS   = 30\nBATCH    = 2                      # keep small for memory\nWEIGHT_DECAY = 5e-4\nVAL_FRAC = 0.2\nV_SCALE  = 4500.0                 # velocity scaling\nDEVICE   = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ntorch.cuda.manual_seed(42); np.random.seed(42)\n\n# Dynamic loss-scale ramp parameters\nloss_scale_start        = 10.0      # lower initial multiplier\nloss_scale_target       = V_SCALE\nloss_scale_warmup_steps = 5000\n\n# Data augmentation\nnoise_std = 0.01     # additive Gaussian noise on seismic input (normalized units)\namp_range = 0.05     # random amplitude scaling ±5%\n\n# Gradient clipping\nCLIP_NORM = 300.0    # global norm clipping threshold\nadaptive_clip_coef = 0.0  # 0 to disable adaptive clipping\n\n# Loss: SmoothL1 with beta in normalized units\nbeta_norm = 0.02\ncriterion_smoothl1 = nn.SmoothL1Loss(beta=beta_norm)\ndef list_pairs(fam_dir):\n    data_dir, model_dir = os.path.join(fam_dir,\"data\"), os.path.join(fam_dir,\"model\")\n    if os.path.isdir(data_dir) and os.path.isdir(model_dir):\n        d = sorted(f for f in os.listdir(data_dir)  if f.endswith(\".npy\"))\n        m = sorted(f for f in os.listdir(model_dir) if f.endswith(\".npy\"))\n        return [(os.path.join(data_dir, d[i]), os.path.join(model_dir, m[i])) for i in range(len(d))]\n    files = [f for f in os.listdir(fam_dir) if f.endswith(\".npy\")]\n    seis, vel = {}, {}\n    for f in files:\n        if f.startswith((\"data\",\"seis\")):\n            seis[pat_data.sub(\"\",f)]  = f\n        elif f.startswith((\"model\",\"vel\")):\n            vel[pat_model.sub(\"\",f)]  = f\n    return [(os.path.join(fam_dir,seis[k]), os.path.join(fam_dir,vel[k])) for k in sorted(set(seis)&set(vel))]\n# --------------------------------------------------\n\n# -------- mmap index dataset (unchanged) ----------\ndef build_index(base):\n    idx = []\n    for fam in sorted(os.listdir(base)):\n        fam_dir = os.path.join(base, fam)\n        if not os.path.isdir(fam_dir): continue\n        for sfile, vfile in list_pairs(fam_dir):\n            for i in range(500):\n                idx.append((sfile, vfile, i))\n    return idx\n\nindex_all = build_index(BASE_DIR)\ntr_ids, va_ids = train_test_split(np.arange(len(index_all)),\n                                  test_size=VAL_FRAC, random_state=42, shuffle=True)\n\nclass VelDS(Dataset):\n    def __init__(self, idx_list):\n        self.idx = idx_list\n    def __len__(self):\n        return len(self.idx)\n    def __getitem__(self, k):\n        sfile, vfile, i = self.idx[k]\n        s = np.load(sfile, mmap_mode='r')[i].copy().astype(np.float32)   # (5,1000,70)\n        v = np.load(vfile, mmap_mode='r')[i]\n        v = v.squeeze().copy().astype(np.float32) / V_SCALE\n        return torch.from_numpy(s), torch.from_numpy(v)\n\ndl_tr = DataLoader(VelDS([index_all[i] for i in tr_ids]), batch_size=BATCH,\n                   shuffle=True, num_workers=4, pin_memory=True)\ndl_va = DataLoader(VelDS([index_all[i] for i in va_ids]), batch_size=BATCH,\n                   shuffle=False, num_workers=4, pin_memory=True)\n# --------------------------------------------------\n\n# ---------------- token stems --------------------\ndef split_fft(x):\n    # x: (B,5,1000,70)\n    F = torch.fft.rfft2(x, dim=(-2,-1))  # yields (B,5,1000,36)\n    return torch.cat([F.real, F.imag, torch.log(torch.abs(F)+1e-6)], dim=1)\n\nclass StemRaw(nn.Module):\n    def __init__(self, emb_d):\n        super().__init__()\n        self.ln = nn.LayerNorm(emb_d)\n        self.conv = nn.Conv2d(5, emb_d, (PATCH_T, PATCH_X), stride=(STR_T, STR_X))\n        self.drop = nn.Dropout(DROP)\n    def forward(self, x):\n        f = self.conv(x)                      # (B, emb_d, T', X')\n        tok = f.flatten(2).transpose(1,2)     # (B, Ntok, emb_d)\n        tok = self.ln(tok)\n        return self.drop(tok)\n\nclass StemFFT(nn.Module):\n    def __init__(self, emb_d):\n        super().__init__()\n        self.ln = nn.LayerNorm(emb_d)\n        self.conv = nn.Conv2d(15, emb_d, (PATCH_T, PATCH_X), stride=(STR_T, STR_X))\n        self.drop = nn.Dropout(DROP)\n    def forward(self, x):\n        spec = split_fft(x)                   # (B,15,1000,36)\n        f = self.conv(spec)                   # (B, emb_d, T', X')\n        tok = f.flatten(2).transpose(1,2)     \n        tok = self.ln(tok)\n        return self.drop(tok)\n# --------------------------------------------------\n\n# -------------- MetaFormer (PoolMixer) ------------\nclass PoolMixer(nn.Module):\n    def __init__(self, C, W_tok):\n        super().__init__()\n        self.W = W_tok\n    def forward(self, x):\n        # x: (B, 1+N, C)\n        cls, tok = x[:, :1], x[:, 1:]\n        B, Np, C = tok.shape\n        W = self.W\n        H = Np // W\n        assert H * W == Np, f\"grid {H}×{W}!={Np}\"\n        y = tok.transpose(1, 2).reshape(B, C, H, W)\n        y = F.avg_pool2d(y, 3, 1, 1) - y\n        tok2 = tok + y.flatten(2).transpose(1, 2)\n        return torch.cat([cls, tok2], dim=1)\n\nclass MetaBlock(nn.Module):\n    def __init__(self, C, W_tok, ff_mult=4, residual_scale=1.0):\n        super().__init__()\n        self.norm1 = nn.LayerNorm(C)\n        self.mix   = PoolMixer(C, W_tok)\n        self.norm2 = nn.LayerNorm(C)\n        self.ffn   = nn.Sequential(\n            nn.Linear(C, ff_mult*C),\n            nn.GELU(),\n            nn.Linear(ff_mult*C, C)\n        )\n        self.residual_scale = residual_scale\n    def forward(self, x):\n        y1 = self.mix(self.norm1(x))\n        x = x + self.residual_scale * y1\n        y2 = self.ffn(self.norm2(x))\n        return x + self.residual_scale * y2\n\nclass MetaEncoder(nn.Module):\n    def __init__(self, depth: int, C: int, W_tok: int, residual_scale=1.0):\n        super().__init__()\n        self.blocks = nn.ModuleList(\n            [MetaBlock(C, W_tok, residual_scale=residual_scale) for _ in range(depth)]\n        )\n    def forward(self, x):\n        for blk in self.blocks:\n            x = blk(x)\n        return x\n# --------------------------------------------------\n\n# -------------- PixelShuffle + mini-U-Net ---------\nclass ShuffleUNetRefine(nn.Module):\n    def __init__(self, base=32):\n        super().__init__()\n        self.pre = nn.Conv2d(1, base*4, 1)\n        self.pre_gn = nn.GroupNorm(num_groups=1, num_channels=base*4)\n        self.shuffle = nn.PixelShuffle(2)          \n        self.enc1 = nn.Sequential(\n            nn.Conv2d(base, base, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base),\n            nn.ReLU(True),\n            nn.Conv2d(base, base, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base),\n            nn.ReLU(True),\n        )\n        self.pool1 = nn.MaxPool2d(2)                \n        self.enc2 = nn.Sequential(\n            nn.Conv2d(base, base*2, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base*2),\n            nn.ReLU(True),\n            nn.Conv2d(base*2, base*2, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base*2),\n            nn.ReLU(True),\n        )\n        self.pool2 = nn.MaxPool2d(2)                \n        self.enc3 = nn.Sequential(\n            nn.Conv2d(base*2, base*4, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base*4),\n            nn.ReLU(True),\n            nn.Conv2d(base*4, base*4, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base*4),\n            nn.ReLU(True),\n        )\n        self.up2 = nn.ConvTranspose2d(base*4, base*2, 2,2)\n        self.dec2 = nn.Sequential(\n            nn.Conv2d(base*4, base*2, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base*2),\n            nn.ReLU(True),\n            nn.Conv2d(base*2, base*2, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base*2),\n            nn.ReLU(True),\n        )\n        self.up1 = nn.ConvTranspose2d(base*2, base, 2,2)\n        self.dec1 = nn.Sequential(\n            nn.Conv2d(base*2, base, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base),\n            nn.ReLU(True),\n            nn.Conv2d(base, base, 3, padding=1),\n            nn.GroupNorm(num_groups=1, num_channels=base),\n            nn.ReLU(True),\n        )\n        self.outc = nn.Conv2d(base, 1, 1)\n\n    def forward(self, c):\n        x = self.shuffle(self.pre_gn(self.pre(c)))       \n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool1(e1))\n        e3 = self.enc3(self.pool2(e2))\n        d2 = self.up2(e3); d2 = self.dec2(torch.cat([d2,e2], dim=1))\n        d1 = self.up1(d2); d1 = self.dec1(torch.cat([d1,e1], dim=1))\n        hi = self.outc(d1)                 \n        return F.avg_pool2d(hi, 2).squeeze(1)  # (B,70,70)\n# --------------------------------------------------\n\n# -------------- Dual-branch MetaFormer model ------\nclass DualBranchMeta(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.raw = StemRaw(EMB_D)\n        self.fft = StemFFT(EMB_D)\n        # token grid sizes\n        W_RAW = (70 - PATCH_X)//STR_X + 1\n        freq_dim = 70//2 + 1\n        W_FFT = (freq_dim - PATCH_X)//STR_X + 1\n        self.enc_raw = MetaEncoder(DEPTH, EMB_D, W_RAW)\n        self.enc_fft = MetaEncoder(DEPTH, EMB_D, W_FFT)\n        # fusion to coarse map\n        self.fuse = nn.Sequential(\n            nn.LayerNorm(2*EMB_D),\n            nn.Linear(2*EMB_D, 4*EMB_D), nn.GELU(),\n            nn.Linear(4*EMB_D, 70*70)\n        )\n        self.coarse_norm = nn.GroupNorm(1,1)\n        self.refine = ShuffleUNetRefine(base=32)\n\n    def forward(self, s):\n        B = s.size(0)\n        r_tok = self.raw(s)\n        f_tok = self.fft(s)\n        zero = torch.zeros(B,1,EMB_D,device=s.device)\n        r_all = torch.cat([zero, r_tok], dim=1)\n        f_all = torch.cat([zero, f_tok], dim=1)\n        r_enc = self.enc_raw(r_all)\n        f_enc = self.enc_fft(f_all)\n        r_cls = r_enc[:,1:,:].mean(dim=1)\n        f_cls = f_enc[:,1:,:].mean(dim=1)\n        fused = torch.cat([r_cls, f_cls], dim=1)\n        coarse = self.fuse(fused).view(B,1,70,70)\n        coarse_norm = self.coarse_norm(coarse)\n        fine = self.refine(coarse_norm)\n        return coarse_norm, fine\nprint(\"Done\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pickle\ncheckpoint_path = \"training_state.pkl\"\n\n# -------- resume from checkpoint --------\nimport os\nif os.path.exists(checkpoint_path):\n    with open(checkpoint_path, 'rb') as f:\n        state = pickle.load(f)\n    model.load_state_dict(state['model_state'])\n    optimizer.load_state_dict(state['optimizer_state'])\n    scheduler.load_state_dict(state['scheduler_state'])\n    scheduler_plateau.load_state_dict(state['scheduler_plateau_state'])\n    scaler.load_state_dict(state['scaler_state'])\n    global_step = state['global_step']\n    best_val = state['best_val']\n    start_epoch = state['epoch']\n    print(f\"[Resume] Loaded checkpoint. Resuming from epoch {start_epoch}\")\nelse:\n    global_step = 0\n    best_val = float('inf')\n    start_epoch = 1","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CHECKPOINT=False\ndef register_linear_hooks(model):\n    def make_hook(name):\n        def hook(module, inp, out):\n            if torch.isnan(out).any() or torch.isinf(out).any():\n                mn, mx = out.min().item(), out.max().item()\n                print(f\"[LEAK] NaN/Inf in output of Linear '{name}': min={mn:.3e}, max={mx:.3e}\")\n        return hook\n\n    for name, module in model.named_modules():\n        if isinstance(module, nn.Linear):\n            module.register_forward_hook(make_hook(name))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/usr/bin/env python3\n\"\"\"\nDual-branch MetaFormer + Stable Training Pipeline with Phase-Aware CQT\nIncludes:\n- Phase-aware Constant-Q Transform (magnitude + phase channels)\n- CQT clipping to avoid extremes\n- MetaFormer encoder with clamped residuals and FFN\n- PixelShuffle U-Net refine with GroupNorm\n- Mixed-precision, gradient accumulation & clipping\n- LR scheduling with plateau re-warm\n- Checkpointing & resume\n\"\"\"\nimport os, re, math, pickle\nimport numpy as np\nimport librosa\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn.utils import spectral_norm\nfrom torch.optim.lr_scheduler import LambdaLR, ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.amp import autocast, GradScaler\nfrom sklearn.model_selection import train_test_split\n\n# ----------------- Hyperparameters -----------------\nBASE_DIR        = \"/kaggle/input/waveform-inversion/train_samples\"\nPATCH_T, PATCH_X= 32, 16\nSTR_T, STR_X    = 4,  2\nEMB_D           = 256\nDEPTH           = 6\nDROP            = 0.1\n\n# CQT hyperparameters\nCQT_BINS                = 36\nCQT_BINS_PER_OCTAVE     = 12\nCQT_HOP_LENGTH          = 1\n\n# LR schedule & accumulation\nLR_START        = 1e-6\nLR_BASE_RAW     = 1e-4\nLR_MIN          = 2e-5\nWARMUP_STEPS    = 500\nEPOCHS          = 30\nBATCH           = 2\nACCUM_STEPS     = 4\nLR_BASE         = LR_BASE_RAW * ACCUM_STEPS\n\nWEIGHT_DECAY    = 5e-4\nVAL_FRAC        = 0.2\nV_SCALE         = 4500.0\nDEVICE          = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# Disable augmentations for stability\nrecv_drop_p     = 0.0\nnoise_std       = 0.0\n\n# Smoothness regularizer weight\nsmooth_lmbda    = 1e-3\n\n# Gradient clipping norm\nCLIP_NORM       = 300.0\n\n# Loss\nbeta_norm       = 0.02\ncriterion_smooth= nn.SmoothL1Loss(beta=beta_norm)\n\n# --------------- Data pipeline ---------------------\npat_data  = re.compile(r'^(data|seis)[_\\-]?')\npat_model = re.compile(r'^(model|vel)[_\\-]?')\n\ndef list_pairs(fam_dir):\n    ddir, mdir = os.path.join(fam_dir, \"data\"), os.path.join(fam_dir, \"model\")\n    if os.path.isdir(ddir) and os.path.isdir(mdir):\n        files_d = sorted(f for f in os.listdir(ddir) if f.endswith('.npy'))\n        files_m = sorted(f for f in os.listdir(mdir) if f.endswith('.npy'))\n        return [(os.path.join(ddir, files_d[i]), os.path.join(mdir, files_m[i]))\n                for i in range(len(files_d))]\n    seis, vel = {}, {}\n    for f in os.listdir(fam_dir):\n        if not f.endswith('.npy'): continue\n        if f.startswith(('data','seis')): seis[pat_data.sub('', f)] = f\n        else: vel[pat_model.sub('', f)] = f\n    keys = sorted(set(seis) & set(vel))\n    return [(os.path.join(fam_dir, seis[k]), os.path.join(fam_dir, vel[k])) for k in keys]\n\nindex_all = [(d,m,i) for fam in sorted(os.listdir(BASE_DIR)) if os.path.isdir(os.path.join(BASE_DIR,fam))\n             for d,m in list_pairs(os.path.join(BASE_DIR,fam)) for i in range(500)]\ntrain_ids, val_ids = train_test_split(list(range(len(index_all))), test_size=VAL_FRAC,\n                                     random_state=42)\n\nclass VelDS(Dataset):\n    def __init__(self, ids): self.ids = ids\n    def __len__(self): return len(self.ids)\n    def __getitem__(self, idx):\n        sfile, vfile, shot = index_all[self.ids[idx]]\n        seis = np.load(sfile, mmap_mode='r')[shot].astype(np.float32)\n        vel  = np.load(vfile, mmap_mode='r')[shot].astype(np.float32) / V_SCALE\n        return torch.from_numpy(seis), torch.from_numpy(vel)\n\ntrain_loader = DataLoader(VelDS(train_ids), batch_size=BATCH, shuffle=True,\n                          num_workers=4, pin_memory=True)\nval_loader   = DataLoader(VelDS(val_ids),   batch_size=BATCH, shuffle=False,\n                          num_workers=4, pin_memory=True)\n\n# -------------- Phase-aware CQT -------------------\ndef split_cqt(x):\n    B, C, T, N = x.shape\n    flat = x.permute(0,1,3,2).reshape(-1, T).cpu().numpy()\n    sr = 1.0; nyquist = sr/2.0\n    fmin = nyquist / (2 ** (CQT_BINS / CQT_BINS_PER_OCTAVE))\n    cqt_list = []\n    for vec in flat:\n        CQT = librosa.cqt(vec, sr=sr, hop_length=CQT_HOP_LENGTH,\n                          fmin=fmin, n_bins=CQT_BINS,\n                          bins_per_octave=CQT_BINS_PER_OCTAVE)\n        cqt_list.append(CQT)\n    CQT = torch.from_numpy(np.stack(cqt_list)).to(x.device)\n    mag, phase = torch.abs(CQT), torch.angle(CQT)\n    spec = torch.cat([mag, phase], dim=1)\n    spec = spec.view(B, C, N, 2*CQT_BINS, -1).permute(0,1,3,4,2)\n    return torch.clamp(spec.reshape(B, C*2*CQT_BINS, spec.size(3), N), -10.0, 10.0)\n\n# ------------- Model components -------------------\nclass StemRaw(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.conv = nn.Conv2d(5, EMB_D, (PATCH_T, PATCH_X), (STR_T, STR_X))\n        self.norm = nn.LayerNorm(EMB_D)\n        self.drop = nn.Dropout(DROP)\n    def forward(self, x):\n        tok = self.conv(x).flatten(2).transpose(1,2)\n        tok = torch.clamp(tok, -1e2, 1e2)\n        tok = self.norm(tok)\n        return self.drop(tok)\n\nclass StemCQT(nn.Module):\n    def __init__(self):\n        super().__init__()\n        in_ch = 5 * 2 * CQT_BINS\n        self.conv = nn.Conv2d(in_ch, EMB_D, (PATCH_T, PATCH_X), (STR_T, STR_X))\n        self.norm = nn.LayerNorm(EMB_D)\n        self.drop = nn.Dropout(DROP)\n    def forward(self, x):\n        spec = split_cqt(x)\n        tok  = self.conv(spec).flatten(2).transpose(1,2)\n        tok  = torch.clamp(tok, -1e2, 1e2)\n        tok  = self.norm(tok)\n        return self.drop(tok)\n\nclass PoolMixer(nn.Module):\n    def __init__(self, C, W): super().__init__(); self.W = W\n    def forward(self, x):\n        cls, tok = x[:,:1], x[:,1:]\n        B,N,Cc = tok.shape; H = N//self.W\n        y = tok.transpose(1,2).reshape(B,Cc,H,self.W)\n        y2 = F.avg_pool2d(y,3,1,1) - y\n        return torch.cat([cls, tok + y2.flatten(2).transpose(1,2)],1)\n\nclass MetaBlock(nn.Module):\n    def __init__(self, C, W, scale=0.5):\n        super().__init__()\n        self.norm1 = nn.LayerNorm(C)\n        self.mix   = PoolMixer(C,W)\n        self.norm2 = nn.LayerNorm(C)\n        self.ffn   = nn.Sequential(nn.Linear(C,4*C), nn.GELU(), nn.Linear(4*C,C))\n        self.scale = scale\n    def forward(self, x):\n        y = self.mix(self.norm1(x)); y = torch.clamp(y,-1e2,1e2)\n        x = x + self.scale*y\n        z = self.ffn(self.norm2(x)); z=torch.clamp(z,-1e2,1e2)\n        return x + self.scale*z\n\nclass MetaEncoder(nn.Module):\n    def __init__(self, d,C,W): super().__init__(); self.blocks=nn.ModuleList([MetaBlock(C,W) for _ in range(d)])\n    def forward(self,x):\n        for b in self.blocks: x=b(x)\n        return x\n\nclass ShuffleUNetRefine(nn.Module):\n    def __init__(self, base=32):\n        super().__init__()\n        self.pre      = spectral_norm(nn.Conv2d(1,base*4,1))\n        self.norm_pre = nn.GroupNorm(1,base*4)\n        self.shuffle  = nn.PixelShuffle(2)\n        self.enc1     = spectral_norm(nn.Conv2d(base,base,3,padding=1))\n        self.norm1    = nn.GroupNorm(1,base)\n        self.enc2     = spectral_norm(nn.Conv2d(base,base*2,3,padding=1))\n        self.norm2    = nn.GroupNorm(1,base*2)\n        self.enc3     = spectral_norm(nn.Conv2d(base*2,base*4,3,padding=1))\n        self.norm3    = nn.GroupNorm(1,base*4)\n        self.up2      = nn.ConvTranspose2d(base*4,base*2,2,2)\n        self.dec2     = spectral_norm(nn.Conv2d(base*4,base*2,3,padding=1))\n        self.normd2   = nn.GroupNorm(1,base*2)\n        self.up1      = nn.ConvTranspose2d(base*2,base,2,2)\n        self.dec1     = spectral_norm(nn.Conv2d(base*2,base,3,padding=1))\n        self.normd1   = nn.GroupNorm(1,base)\n        self.outc     = nn.Conv2d(base,1,1)\n    def forward(self, c):\n        x = self.shuffle(self.norm_pre(self.pre(c)))\n        e1 = F.gelu(self.norm1(self.enc1(x)))\n        e2 = F.gelu(self.norm2(self.enc2(F.max_pool2d(e1,2))))\n        e3 = F.gelu(self.norm3(self.enc3(F.max_pool2d(e2,2))))\n        d2=F.gelu(self.normd2(self.dec2(torch.cat([self.up2(e3),e2],1))))\n        d1=F.gelu(self.normd1(self.dec1(torch.cat([self.up1(d2),e1],1))))\n        hi = self.outc(d1)\n        return F.avg_pool2d(hi,2).squeeze(1)\n\nclass DualBranchMeta(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.raw = StemRaw()\n        self.fft = StemCQT()\n        # compute token grid widths\n        W_TIME  = (1000 - PATCH_T)//STR_T + 1\n        W_SPACE = (70   - PATCH_X)//STR_X + 1\n        self.enc_raw = MetaEncoder(DEPTH, EMB_D, W_TIME)\n        self.enc_fft = MetaEncoder(DEPTH, EMB_D, W_SPACE)\n        self.fuse_norm = nn.LayerNorm(2*EMB_D)\n        self.fuse      = nn.Sequential(nn.Linear(2*EMB_D,4*EMB_D), nn.GELU(), nn.Linear(4*EMB_D,70*70))\n        self.coarse_norm = nn.GroupNorm(1,1)\n        self.refine = ShuffleUNetRefine(base=32)\n    def forward(self,s):\n        B = s.size(0)\n        r_tok = self.raw(s)\n        f_tok = self.fft(s)\n        zero = torch.zeros(B,1,EMB_D,device=s.device)\n        r_enc= self.enc_raw(torch.cat([zero,r_tok],1))\n        f_enc= self.enc_fft(torch.cat([zero,f_tok],1))\n        r_cls= r_enc[:,1:,:].mean(1)\n        f_cls= f_enc[:,1:,:].mean(1)\n        h    = self.fuse_norm(torch.cat([r_cls,f_cls],1))\n        h    = torch.clamp(h,-1e2,1e2)\n        coarse = self.fuse(h).view(B,1,70,70)\n        coarse = torch.clamp(coarse,-1e3,1e3)\n        coarse = self.coarse_norm(coarse)\n        return coarse, self.refine(coarse)\n\n# -------------- Training loop ---------------------\nmodel = DualBranchMeta().to(DEVICE)\nscaler= GradScaler()\noptimizer= torch.optim.AdamW([\n    {'params': model.raw.parameters()},\n    {'params': model.fft.parameters()},\n    {'params': model.enc_raw.parameters()},\n    {'params': model.enc_fft.parameters(), 'lr': LR_BASE_RAW*0.5},\n    {'params': model.fuse.parameters(),   'lr': LR_BASE_RAW*0.1},\n    {'params': model.coarse_norm.parameters()},\n    {'params': model.refine.parameters(), 'lr': LR_BASE_RAW/10}\n], lr=LR_BASE, weight_decay=WEIGHT_DECAY)\n\nsteps_per_epoch = len(train_loader) // ACCUM_STEPS\nscheduler = LambdaLR(optimizer, lambda step: \n    (LR_START + (LR_BASE_RAW - LR_START)\n     * min(step, WARMUP_STEPS)/WARMUP_STEPS)/LR_BASE_RAW\n    if step < WARMUP_STEPS else\n    (LR_MIN + 0.5*(LR_BASE_RAW - LR_MIN)*\n     (1 + math.cos(math.pi * min((step - WARMUP_STEPS)/(EPOCHS*steps_per_epoch - WARMUP_STEPS),1))))/LR_BASE_RAW\n)\nscheduler_plateau = ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=3, min_lr=1e-6)\nbest_val = float('inf'); global_step=0\n\n# optional resume\ncheckpoint_path = 'checkpoint.pkl'\nif CHECKPOINT:\n    state = pickle.load(open(checkpoint_path,'rb'))\n    model.load_state_dict(state['model_state'])\n    optimizer.load_state_dict(state['optimizer_state'])\n    scheduler.load_state_dict(state['scheduler_state'])\n    scheduler_plateau.load_state_dict(state['scheduler_plateau_state'])\n    scaler.load_state_dict(state['scaler_state'])\n    global_step = state['global_step']\n    best_val    = state['best_val']\n    start_epoch = state['epoch']+1\nelse:\n    start_epoch = 1\n\nfor ep in range(start_epoch, EPOCHS+1):\n    model.train()\n    running_train=0.0; count_train=0; skip_ctr=0\n    optimizer.zero_grad()\n    for step, (seis, vel) in enumerate(train_loader):\n        seis, vel = seis.to(DEVICE), vel.to(DEVICE).unsqueeze(1)\n        with autocast('cuda'):\n            coarse, pred = model(seis)\n            loss_mae = F.l1_loss(pred, vel)\n            loss_s   = smooth_lmbda * F.l1_loss(coarse, vel)\n            loss     = (loss_mae + loss_s) / ACCUM_STEPS\n        if torch.isnan(loss) or torch.isinf(loss):\n            skip_ctr += 1; optimizer.zero_grad(); continue\n        scaler.scale(loss).backward()\n        if (step+1) % ACCUM_STEPS == 0:\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), CLIP_NORM)\n            scaler.step(optimizer); scaler.update(); optimizer.zero_grad()\n            scheduler.step(); global_step+=1\n            running_train += loss_mae.item() * V_SCALE * BATCH\n            count_train   += BATCH\n    train_mae = running_train/count_train if count_train>0 else float('nan')\n    print(f\"Epoch {ep}/{EPOCHS} Train MAE: {train_mae:.1f} m/s (skipped {skip_ctr})\")\n    \n    model.eval(); val_err=0.0\n    with torch.no_grad():\n        for seis, vel in val_loader:\n            seis, vel = seis.to(DEVICE), vel.to(DEVICE).unsqueeze(1)\n            _, pred = model(seis)\n            val_err += F.l1_loss(pred, vel).item() * V_SCALE * vel.size(0)\n    val_mae = val_err / len(val_loader.dataset)\n    print(f\"Epoch {ep}/{EPOCHS} Validation MAE: {val_mae:.1f} m/s\")\n\n    if global_step >= WARMUP_STEPS:\n        scheduler_plateau.step(val_mae)\n    if val_mae < best_val:\n        best_val = val_mae\n        torch.save(model.state_dict(), 'best_model.pth')\n        print(\"Saved new best model\")\n    # checkpoint\n    state = {'epoch': ep, 'model_state': model.state_dict(),\n             'optimizer_state': optimizer.state_dict(),\n             'scheduler_state': scheduler.state_dict(),\n             'scheduler_plateau_state': scheduler_plateau.state_dict(),\n             'scaler_state': scaler.state_dict(),\n             'global_step': global_step, 'best_val': best_val}\n    pickle.dump(state, open(checkpoint_path,'wb'))\n    print(f\"Checkpoint saved for epoch {ep}\\n\")\n\nprint(\"Training complete.\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import math, os, re, random, pickle, numpy as np\nfrom pathlib import Path\nfrom tqdm import tqdm\n\nimport torch, torch.nn as nn, torch.nn.functional as F\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torch.utils.data import Dataset, DataLoader\n\n# ────── paths ──────\nBASE_DIR = \"waveform-inversion/train_samples\"         # <-- change\nCKPT_DIR = \"./ckpt_dual\"\nPath(CKPT_DIR).mkdir(parents=True, exist_ok=True)\n\n# ────── data & training hyper-params ──────\nVAL_FRAC = 0.20\nV_SCALE  = 4_500.0\nEMB_D   = 256\nBATCH_STAGE1 = 2\nBATCH_STAGE2 = 6\nPATCH_T, PATCH_X = 32, 16\nSTR_T,   STR_X   = 4,   2\n\nBASE_DIR    = \"/kaggle/input/waveform-inversion/train_samples\"\nDEPTH   = 6\n\nEPOCH1, EPOCH2 = 25, 15\nLR1,    LR2    = 2e-4, 1e-4\nWARM1,  WARM2  = 5_000, 2_000\nACC1,   ACC2   = 4, 1          # grad-accum\nDROP = 0.1\nSMOOTH_LMBDA  = 0.1            # coarse smoothness loss\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ntorch.manual_seed(42); np.random.seed(42); random.seed(42)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"_pat_d, _pat_m = re.compile(r'^(data|seis)[_\\-]?'), re.compile(r'^(model|vel)[_\\-]?')\n\ndef _list_pairs(fam_dir):\n    files = [f for f in os.listdir(fam_dir) if f.endswith('.npy')]\n    seis, vel = {}, {}\n    for f in files:\n        if f.startswith(('data', 'seis')):    seis[_pat_d.sub('', f)] = f\n        elif f.startswith(('model','vel')):   vel [_pat_m.sub('', f)] = f\n    return [(os.path.join(fam_dir, seis[k]), os.path.join(fam_dir, vel[k]))\n            for k in sorted(set(seis)&set(vel))]\n\ndef build_index(root):\n    idx=[]\n    for fam in sorted(os.listdir(root)):\n        famdir=os.path.join(root,fam)\n        if not os.path.isdir(famdir): continue\n        for s,v in _list_pairs(famdir):\n            n = np.load(s, mmap_mode='r').shape[0]      # true #shots\n            idx.extend([(s,v,i) for i in range(n)])\n    return idx\n\nclass VelDS(Dataset):\n    def __init__(self, items): self.items=items\n    def __len__(self): return len(self.items)\n    def __getitem__(self,k):\n        s,v,i = self.items[k]\n        seis = np.load(s, mmap_mode='r')[i].astype(np.float32)   # (5,1000,70)\n        vel  = np.load(v, mmap_mode='r')[i].squeeze().astype(np.float32)\n        return torch.from_numpy(seis), torch.from_numpy(vel/V_SCALE)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# -----------------------------------------\n\n# ----------- model -----------------------\ndef split_fft(x):\n    \"\"\"log-FFT over (t,x) with sign-preserving magnitude\"\"\"\n    F2   = torch.fft.rfft2(x, dim=(-2,-1))\n    mag  = torch.log(torch.abs(F2)+1e-6)\n    real = torch.sign(F2.real) * mag\n    imag = torch.sign(F2.imag) * mag\n    return torch.cat([real, imag], dim=1)   # (B, 10, 1000, 36)\n\ndef center_crop(feat: torch.Tensor, size):\n    h, w = feat.shape[-2:]\n    dh, dw = (h - size[0]) // 2, (w - size[1]) // 2\n    return feat[..., dh:dh+size[0], dw:dw+size[1]]\n\nclass StemRaw(nn.Module):\n    def __init__(self, ed):\n        super().__init__()\n        self.conv = nn.Conv2d(5, ed, (PATCH_T,PATCH_X), stride=(STR_T,STR_X))\n        self.ln   = nn.LayerNorm(ed); self.drop = nn.Dropout(DROP)\n    def forward(self, x):\n        f = self.conv(x).flatten(2).transpose(1,2)\n        return self.drop(self.ln(f))\n\nclass StemFFT(nn.Module):\n    def __init__(self, ed):\n        super().__init__()\n        self.conv = nn.Conv2d(10, ed, (PATCH_T,PATCH_X), stride=(STR_T,STR_X))\n        self.ln   = nn.LayerNorm(ed); self.drop = nn.Dropout(DROP)\n    def forward(self, x):\n        f = self.conv(split_fft(x)).flatten(2).transpose(1,2)\n        return self.drop(self.ln(f))\n\nclass PoolMixer(nn.Module):\n    def __init__(self, C, W): super().__init__(); self.W = W\n    def forward(self, x):\n        cls, tok = x[:,:1], x[:,1:]\n        B,N,C = tok.shape; H = N//self.W\n        y = tok.transpose(1,2).reshape(B,C,H,self.W)\n        y = F.avg_pool2d(y,3,1,1) - y\n        return torch.cat([cls, tok+y.flatten(2).transpose(1,2)],1)\n\nclass MetaBlock(nn.Module):\n    def __init__(self, C, W, drop=DROP):\n        super().__init__()\n        self.norm1 = nn.LayerNorm(C); self.mix = PoolMixer(C,W)\n        self.norm2 = nn.LayerNorm(C)\n        self.ffn = nn.Sequential(nn.Linear(C,4*C), nn.GELU(),\n                                 nn.Dropout(drop), nn.Linear(4*C,C))\n    def forward(self,x):\n        x = x + self.mix(self.norm1(x))\n        x = x + self.ffn(self.norm2(x))\n        return x\n\nclass MetaEncoder(nn.Module):\n    def __init__(self, d,C,W): super().__init__()\n    def __init__(self, d,C,W):\n        super().__init__()\n        self.blocks = nn.ModuleList([MetaBlock(C,W) for _ in range(d)])\n    def forward(self,x):\n        for b in self.blocks: x=b(x)\n        return x\n\nclass ShuffleUNetRefine(nn.Module):\n    def __init__(self, base: int = 32):\n        super().__init__()\n        self.drop = nn.Dropout2d(0.1)\n\n        # ---------- encoder ----------\n        self.pre  = nn.Conv2d(1, base * 4, 1)\n        self.normp = nn.LayerNorm([base * 4, 70, 70])\n        self.shuffle = nn.PixelShuffle(2)                    # 70×70 → 140×140\n\n        def CBR(cin, cout):\n            return nn.Sequential(\n                nn.Conv2d(cin, cout, 3, 1, 1, bias=False),\n                nn.GroupNorm(1, cout),                       # res-agnostic norm\n                nn.GELU()\n            )\n\n        self.enc1  = nn.Sequential(CBR(base,     base),     CBR(base,     base))\n        self.pool1 = nn.MaxPool2d(2)                         # 140 → 70\n        self.enc2  = nn.Sequential(CBR(base, base * 2), CBR(base * 2, base * 2))\n        self.pool2 = nn.MaxPool2d(2)                         # 70  → 35\n        self.enc3  = nn.Sequential(CBR(base * 2, base * 4),\n                                   CBR(base * 4, base * 4))  # 35 → 17\n\n        # ---------- decoder ----------\n        self.up2  = nn.ConvTranspose2d(base * 4, base * 2, 2, 2)   # 17→34\n        self.dec2 = nn.Sequential(CBR(base * 4, base * 2), CBR(base * 2, base * 2))\n\n        self.up1  = nn.ConvTranspose2d(base * 2, base, 2, 2)       # 35→70\n        self.dec1 = nn.Sequential(CBR(base * 2, base), CBR(base, base))\n\n        self.head = nn.Conv2d(base, 1, 1)\n\n        # --- NEW — learnable γ, zero-init so net = identity at start -----------\n        self.gamma = nn.Parameter(torch.zeros(1))\n\n        # also (re-)zero the 1×1 head so residual starts at *exactly* 0\n        nn.init.constant_(self.head.weight, 0.)\n        nn.init.constant_(self.head.bias,   0.)\n\n    # --------------------------------------------------------------------- #\n    @staticmethod\n    def _center_crop(src: torch.Tensor, size: tuple[int, int]) -> torch.Tensor:\n        h, w = src.shape[-2:]\n        th, tw = size\n        dh, dw = (h - th) // 2, (w - tw) // 2\n        return src[..., dh:dh + th, dw:dw + tw]\n\n    def forward(self, c: torch.Tensor) -> torch.Tensor:           # c:(B,1,70,70)\n        # ----------------- encoder -----------------\n        x0 = self.shuffle(self.normp(self.pre(c)))                # (B,base,140,140)\n\n        e1 = self.enc1(x0)                                        # 140×140\n        e2 = self.enc2(self.pool1(e1))                            # 70 → 35×35\n        e3 = self.enc3(self.pool2(e2))                            # 35 → 17×17\n\n        # ----------------- decoder -----------------\n        u2 = self.up2(e3)                                         # 34×34\n        u2 = self._center_crop(u2, e2.shape[-2:])                 # → 35×35\n        d2 = self.dec2(torch.cat([u2, self.drop(e2)], dim=1))     # 35×35\n\n        u1 = self.up1(d2)                                         # 70×70\n        d1 = self.dec1(torch.cat([u1, self.drop(e1)], dim=1))     # 70×70\n\n        resid = self.head(d1).squeeze(1)                          # (B,70,70)\n\n        # ----------------- gated residual -----------------\n        return self.gamma * resid\n     # (B,70,70)\n\nclass DualBranchMeta(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.raw = StemRaw(EMB_D);  self.fft = StemFFT(EMB_D)\n\n        W_t  = (70-PATCH_X)//STR_X+1\n        W_f  = (36-PATCH_X)//STR_X+1\n\n        self.enc_raw = MetaEncoder(DEPTH, EMB_D, W_t)\n        self.enc_fft = MetaEncoder(DEPTH, EMB_D, W_f)\n\n        self.fuse = nn.Sequential(nn.LayerNorm(2*EMB_D),\n                                  nn.Linear(2*EMB_D,4*EMB_D),\n                                  nn.GELU(), nn.Linear(4*EMB_D,70*70))\n        self.coarse_norm = nn.GroupNorm(1,1)\n        self.refine = ShuffleUNetRefine()\n\n    def forward(self,s):\n        B=s.size(0)\n        r_tok, f_tok = self.raw(s), self.fft(s)\n        zr = torch.zeros(B,1,EMB_D,device=s.device)\n        r_enc = self.enc_raw(torch.cat([zr,r_tok],1))\n        f_enc = self.enc_fft(torch.cat([zr,f_tok],1))\n        fused = torch.cat([r_enc[:,1:,:].mean(1), f_enc[:,1:,:].mean(1)],1)\n        coarse= self.fuse(fused).view(B,1,70,70)\n        coarse= self.coarse_norm(coarse)\n        refined= self.refine(coarse)\n        return coarse.squeeze(1), refined    ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def grad_xy(t):\n    gx = F.pad(t[:,:,1:,:]-t[:,:,:-1,:], (0,0,0,0,1,0))\n    gy = F.pad(t[:,:,:,1:]-t[:,:,:,:-1], (1,0,0,0))\n    return gx, gy\n\ndef laplacian(t):\n    return (F.pad(t[:,:,2:,:]-2*t[:,:,1:-1,:]+t[:,:,:-2,:], (0,0,1,1)) +\n            F.pad(t[:,:,:,2:]-2*t[:,:,:,1:-1]+t[:,:,:,:-2], (1,1,0,0)))\n\n# ─── small conv-norm-activation helper ─────────────────────\nimport torch.nn as nn\ndef _blk(cin, cout, k=3, p=1):\n    return nn.Sequential(\n        nn.Conv2d(cin, cout, k, padding=p, bias=False),\n        nn.GroupNorm(1, cout),\n        nn.GELU()\n    )\n\n\nclass ResidualUNet(nn.Module):\n    def __init__(self, in_ch=4, base=32):\n        super().__init__()\n        # ------- encoder -------\n        self.enc1 = _blk(in_ch,  base)\n        self.pool1= nn.MaxPool2d(2)            # 70→35\n        self.enc2 = _blk(base,  base*2)\n        self.pool2= nn.MaxPool2d(2)            # 35→17\n        self.enc3 = _blk(base*2, base*4)\n\n        # ------- decoder -------\n        self.up2  = nn.ConvTranspose2d(base*4, base*2, 2, 2)\n        self.dec2 = _blk(base*4, base*2)\n        self.up1  = nn.ConvTranspose2d(base*2, base,   2, 2)\n        self.dec1 = _blk(base*2, base)\n\n        self.head = nn.Conv2d(base, 1, 1)\n\n    @staticmethod\n    def _crop(src, size):          # <-- now static & always present\n        h, w = src.shape[-2:]\n        dh, dw = (h-size[0])//2, (w-size[1])//2\n        return src[..., dh:dh+size[0], dw:dw+size[1]]\n\n    def forward(self, x4):\n        e1 = self.enc1(x4)\n        e2 = self.enc2(self.pool1(e1))\n        e3 = self.enc3(self.pool2(e2))\n\n        u2 = self.up2(e3)\n        d2 = self.dec2(torch.cat([u2, self._crop(e2, u2.shape[-2:])], 1))\n\n        u1 = self.up1(d2)\n        d1 = self.dec1(torch.cat([u1, self._crop(e1, u1.shape[-2:])], 1))\n\n        out = self.head(d1)                      # 1×68×68\n        out = F.pad(out, (1,1,1,1), mode='reflect')  # → 70×70\n        return out.squeeze(1)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def cosine_warm(step, warm, total):\n    if step < warm: return (step+1)/warm\n    q = (step-warm)/(total-warm)\n    return 0.5*(1+math.cos(math.pi*min(q,1)))\n\ndef make_sched(opt, lr0, warm, total):\n    return torch.optim.lr_scheduler.LambdaLR(\n        opt, lambda s: cosine_warm(s,warm,total))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"items = build_index(BASE_DIR)\nperm  = np.random.permutation(len(items))\ncut   = int(len(perm)*(1-VAL_FRAC))\ntrain_items, val_items = [items[i] for i in perm[:cut]], [items[i] for i in perm[cut:]]\n\ntrain_set, val_set = VelDS(train_items), VelDS(val_items)\ndl_tr1 = DataLoader(train_set, BATCH_STAGE1, shuffle=True,  num_workers=4, pin_memory=True, drop_last=True)\ndl_va1 = DataLoader(val_set,   BATCH_STAGE1, shuffle=False, num_workers=2, pin_memory=True)\ncoarse_net = DualBranchMeta().to(DEVICE)\n\n\n\nopt1   = torch.optim.AdamW(coarse_net.parameters(), lr=LR1, weight_decay=5e-4)\nsteps1 = EPOCH1 * math.ceil(len(dl_tr1)/ACC1)\nsched1 = make_sched(opt1, LR1, WARM1, steps1)\nscaler = GradScaler()\nbest_val = 9e9","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\nopt1   = torch.optim.AdamW(coarse_net.parameters(), lr=LR1, weight_decay=5e-4)\nsteps1 = EPOCH1 * math.ceil(len(dl_tr1)/ACC1)\nsched1 = make_sched(opt1, LR1, WARM1, steps1)\nscaler = GradScaler()\nbest_val = 9e9\n\nfor epoch in range(1, EPOCH1+1):\n    coarse_net.train(); run=cnt=0\n    pbar = tqdm(enumerate(dl_tr1), total=len(dl_tr1), ncols=90, desc=f\"[1] {epoch}/{EPOCH1}\")\n    opt1.zero_grad(); acc=0\n    for i,(s,v) in pbar:\n        s,v = s.to(DEVICE), v.to(DEVICE)\n        with autocast():\n            coarse,pred = coarse_net(s)\n            tgt = v.unsqueeze(1)\n            loss = (F.l1_loss(pred.unsqueeze(1),tgt) +\n                    SMOOTH_LMBDA*F.l1_loss(coarse,tgt)) / ACC1\n        scaler.scale(loss).backward(); acc+=1\n        if acc==ACC1:\n            scaler.unscale_(opt1)\n            torch.nn.utils.clip_grad_norm_(coarse_net.parameters(), 300)\n            scaler.step(opt1); scaler.update()\n            opt1.zero_grad(); acc=0; sched1.step()\n        run += loss.item()*V_SCALE*s.size(0)*ACC1; cnt += s.size(0)\n    # -------- validation --------\n    coarse_net.eval(); val=0\n    with torch.no_grad():\n        for s,v in dl_va1:\n            s,v=s.to(DEVICE),v.to(DEVICE)\n            _,pred = coarse_net(s)\n            val += F.l1_loss(pred, v.item())*V_SCALE*s.size(0)\n    tr_mae, va_mae = run/cnt, val/len(val_set)\n    print(f\"  ↪ MAE train {tr_mae:6.1f} | val {va_mae:6.1f}\")\n    if va_mae<best_val:\n        best_val=va_mae\n        torch.save(coarse_net.state_dict(), f\"{CKPT_DIR}/coarse_best.pt\")\n        print(\"    ★ saved best coarse\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===========================================================\n#  stage1_baseline.ipynb  –  dual-branch + UNet (Kaggle-ready)\n# ===========================================================\n\n# ------------------- paths & misc --------------------------\nBASE_DIR  = \"/kaggle/input/waveform-inversion/train_samples\"   # ← adjust\nCKPT_DIR  = \"/kaggle/working/ckpt_stage1\"\n!mkdir -p $CKPT_DIR                                             # Kaggle shell\n\nimport os, re, math, pickle, random, numpy as np\nfrom pathlib import Path\nfrom tqdm import tqdm\n\nimport torch, torch.nn as nn, torch.nn.functional as F\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import LambdaLR, ReduceLROnPlateau\n\n# reproducibility -------------------------------------------------------------\ntorch.manual_seed(42); np.random.seed(42); random.seed(42)\n\nDEVICE      = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nBATCH       = 2\nVAL_FRAC    = 0.2\nEPOCHS      = 50\nACC_STEPS   = 4\n\n# model hyper-params ----------------------------------------------------------\nPATCH_T, PATCH_X = 32, 16\nSTR_T,   STR_X   = 4, 2\nEMB_D, DEPTH     = 256, 6\nDROP             = 0.1\nSMOOTH_LMBDA     = 0.1\n\n# optimiser & LR -------------------------------------------------------------\nLR_BASE, LR_MIN  = 8e-4, 2e-5\nUNET_RATIO       = 0.10\nWARM_STEPS       = 10_000\nWEIGHT_DEC       = 5e-4\nCLIP_GLOBAL      = 100.\nV_SCALE          = 4_500.0                # km/s → ~1\n\n# -------------- dataset helpers --------------------------------------------\n_pat_d = re.compile(r'^(data|seis)[_\\-]?')\n_pat_m = re.compile(r'^(model|vel)[_\\-]?')\n\ndef _list_pairs(fam_dir):\n    files = [f for f in os.listdir(fam_dir) if f.endswith(\".npy\")]\n    seis, vel = {}, {}\n    for f in files:\n        if f.startswith((\"data\", \"seis\")):\n            seis[_pat_d.sub(\"\", f)] = f\n        elif f.startswith((\"model\", \"vel\")):\n            vel[_pat_m.sub(\"\", f)] = f\n    return [(os.path.join(fam_dir, seis[k]),\n             os.path.join(fam_dir, vel[k])) for k in sorted(set(seis)&set(vel))]\n\ndef build_index(base):\n    idx=[]\n    for fam in sorted(os.listdir(base)):\n        famdir=os.path.join(base,fam)\n        if not os.path.isdir(famdir): continue\n        for s,v in _list_pairs(famdir):\n            for i in range(500):          # 500 shots/file\n                idx.append((s,v,i))\n    return idx\n\nclass VelDS(Dataset):\n    def __init__(self, tuples):  self.items=tuples\n    def __len__(self):           return len(self.items)\n    def __getitem__(self,k):\n        s,v,i=self.items[k]\n        seis = np.load(s,mmap_mode='r')[i].astype(np.float32)    # (5,1000,70)\n        vel  = np.load(v,mmap_mode='r')[i].squeeze().astype(np.float32)\n        return torch.from_numpy(seis), torch.from_numpy(vel/V_SCALE)\n\n# ------------ model parts ---------------------------------------------------\ndef split_fft(x):\n    F2 = torch.fft.rfft2(x, dim=(-2,-1))\n    mag = torch.log(torch.abs(F2)+1e-6)\n    real = torch.sign(F2.real)*mag\n    imag = torch.sign(F2.imag)*mag\n    return torch.cat([real,imag],1)\n\nclass StemRaw(nn.Module):\n    def __init__(self,d):\n        super().__init__()\n        self.conv=nn.Conv2d(5,d,(PATCH_T,PATCH_X),stride=(STR_T,STR_X))\n        self.ln=nn.LayerNorm(d); self.drop=nn.Dropout(DROP)\n    def forward(self,x):\n        f=self.conv(x).flatten(2).transpose(1,2)\n        return self.drop(self.ln(f))\n\nclass StemFFT(nn.Module):\n    def __init__(self,d):\n        super().__init__()\n        self.conv=nn.Conv2d(10,d,(PATCH_T,PATCH_X),stride=(STR_T,STR_X))\n        self.ln=nn.LayerNorm(d); self.drop=nn.Dropout(DROP)\n    def forward(self,x):\n        f=self.conv(split_fft(x)).flatten(2).transpose(1,2)\n        return self.drop(self.ln(f))\n\nclass PoolMix(nn.Module):\n    def __init__(self,W): super().__init__(); self.W=W\n    def forward(self,x):\n        cls,tok=x[:,:1],x[:,1:]\n        B,N,C=tok.shape; H=N//self.W\n        y=tok.transpose(1,2).reshape(B,C,H,self.W)\n        y=F.avg_pool2d(y,3,1,1)-y\n        return torch.cat([cls,tok+y.flatten(2).transpose(1,2)],1)\n\nclass MetaBlock(nn.Module):\n    def __init__(self,C,W):\n        super().__init__()\n        self.norm1=nn.LayerNorm(C); self.mix=PoolMix(W)\n        self.norm2=nn.LayerNorm(C)\n        self.ffn=nn.Sequential(nn.Linear(C,4*C),nn.GELU(),nn.Dropout(DROP),\n                               nn.Linear(4*C,C))\n    def forward(self,x):\n        x=x+self.mix(self.norm1(x))\n        x=x+self.ffn(self.norm2(x))\n        return x\n\nclass MetaEncoder(nn.Module):\n    def __init__(self,d,C,W):\n        super().__init__()\n        self.blocks=nn.ModuleList([MetaBlock(C,W) for _ in range(d)])\n    def forward(self,x):\n        for b in self.blocks: x=b(x)\n        return x\n\n# --- Pixel-Shuffle U-Net refine --------------------------------------------\ndef _blk(cin,cout):\n    return nn.Sequential(nn.Conv2d(cin,cout,3,1,1,bias=False),\n                         nn.GroupNorm(1,cout), nn.GELU())\n\nclass ShuffleUNetRefine(nn.Module):\n    def __init__(self,base=32):\n        super().__init__()\n        self.pre = nn.Conv2d(1, base*4, 1)\n        self.normp = nn.LayerNorm([base*4,70,70])\n        self.shuffle=nn.PixelShuffle(2)          # → (B,base,140,140)\n\n        self.enc1=_blk(base,base)\n        self.pool1=nn.MaxPool2d(2)\n        self.enc2=_blk(base,base*2)\n        self.pool2=nn.MaxPool2d(2)\n        self.enc3=_blk(base*2,base*4)\n\n        self.up2 = nn.ConvTranspose2d(base*4,base*2,2,2)\n        self.dec2=_blk(base*4,base*2)\n        self.up1 = nn.ConvTranspose2d(base*2,base,2,2)\n        self.dec1=_blk(base*2,base)\n\n        self.outc=nn.Conv2d(base,1,1)\n\n    def forward(self,c):\n        x=self.shuffle(self.normp(self.pre(c)))\n        e1=self.enc1(x)\n        e2=self.enc2(self.pool1(e1))\n        e3=self.enc3(self.pool2(e2))\n        d2=self.dec2(torch.cat([self.up2(e3),e2],1))\n        d1=self.dec1(torch.cat([self.up1(d2),e1],1))\n        out=F.avg_pool2d(self.outc(d1),2)\n        return out.squeeze(1)           # (B,70,70)\n\n# ---------------- full model -----------------------------------------------\nclass DualBranchMeta(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.raw=StemRaw(EMB_D); self.fft=StemFFT(EMB_D)\n        Wt=(70-PATCH_X)//STR_X+1\n        Wf=(36-PATCH_X)//STR_X+1\n        self.enc_r=MetaEncoder(DEPTH,EMB_D,Wt)\n        self.enc_f=MetaEncoder(DEPTH,EMB_D,Wf)\n        self.fuse=nn.Sequential(nn.LayerNorm(2*EMB_D),\n                                nn.Linear(2*EMB_D,4*EMB_D),\n                                nn.GELU(),nn.Linear(4*EMB_D,70*70))\n        self.coarse_norm=nn.GroupNorm(1,1)\n        self.refine=ShuffleUNetRefine()\n\n    def forward(self,s):\n        B=s.size(0)\n        r_tok,f_tok=self.raw(s),self.fft(s)\n        zr=torch.zeros(B,1,EMB_D,device=s.device)\n        r=self.enc_r(torch.cat([zr,r_tok],1))\n        f=self.enc_f(torch.cat([zr,f_tok],1))\n        fused=torch.cat([r[:,1:].mean(1),f[:,1:].mean(1)],1)\n        coarse=self.fuse(fused).view(B,1,70,70)\n        coarse=self.coarse_norm(coarse)\n        pred=self.refine(coarse)\n        return coarse.squeeze(1),pred\n\n# ----------------- data loaders -------------------------------------------\nall_idx=build_index(BASE_DIR)\nperm=np.random.permutation(len(all_idx))\nsplit=int(len(perm)*(1-VAL_FRAC))\ntr,va=perm[:split],perm[split:]\n\ntrain_loader=DataLoader(VelDS([all_idx[i] for i in tr]),\n                        batch_size=BATCH,shuffle=True,num_workers=4,pin_memory=True)\nval_loader  =DataLoader(VelDS([all_idx[i] for i in va]),\n                        batch_size=BATCH,shuffle=False,num_workers=4,pin_memory=True)\n\n# ----------------- training setup -----------------------------------------\nmodel=DualBranchMeta().to(DEVICE)\n\ndecay,no_wd=[],[]\nref_ids={id(p) for p in model.refine.parameters()}\nfor n,p in model.named_parameters():\n    if id(p) in ref_ids: continue\n    (no_wd if p.ndim==1 or n.endswith(\".bias\") else decay).append(p)\n\nopt=torch.optim.AdamW(\n    [{\"params\":decay,\"lr\":LR_BASE},\n     {\"params\":no_wd,\"lr\":LR_BASE,\"weight_decay\":0.},\n     {\"params\":model.refine.parameters(),\"lr\":LR_BASE*UNET_RATIO}],\n    betas=(0.9,0.95), weight_decay=WEIGHT_DEC)\n\ntotal_steps=EPOCHS*math.ceil(len(train_loader)/ACC_STEPS)\ndef lr_main(step):\n    if step<WARM_STEPS: return (step+1)/WARM_STEPS\n    prog=(step-WARM_STEPS)/(total_steps-WARM_STEPS)\n    return 0.8 * max(LR_MIN/LR_BASE,0.5*(1+math.cos(0.9*math.pi*min(prog,1))))\ndef lr_unet(step): return lr_main(step)*UNET_RATIO\nsched=LambdaLR(opt,[lr_main,lr_main,lr_unet])\nplateau=ReduceLROnPlateau(opt,'min',factor=0.5,patience=3,min_lr=1e-6)\n\nloss_fn=nn.L1Loss(); scaler=GradScaler()\nbest=float('inf')\n\n# ----------------- training loop ------------------------------------------\nfor ep in range(1,EPOCHS+1):\n    model.train(); run=cnt=0; opt.zero_grad()\n    pbar=tqdm(enumerate(train_loader),total=len(train_loader),ncols=90,desc=f\"Epoch {ep}\")\n    for i,(seis,vel) in pbar:\n        seis,vel=seis.to(DEVICE),vel.to(DEVICE)\n        with autocast():\n            coarse,pred=model(seis)\n            v=vel\n            l_mae=loss_fn(pred,v)\n            l_s  =SMOOTH_LMBDA*loss_fn(coarse,v)\n            loss=(l_mae+l_s)/ACC_STEPS\n        scaler.scale(loss).backward()\n        if (i+1)%ACC_STEPS==0:\n            scaler.unscale_(opt)\n            torch.nn.utils.clip_grad_norm_(model.parameters(),CLIP_GLOBAL)\n            scaler.step(opt); scaler.update(); opt.zero_grad(); sched.step()\n        run+=l_mae.item()*V_SCALE*seis.size(0); cnt+=seis.size(0)\n        pbar.set_postfix(train_MAE=run/cnt)\n\n    # validation ------------------------------------------------------------\n    model.eval(); val=0.\n    with torch.no_grad():\n        for seis,vel in val_loader:\n            seis,vel=seis.to(DEVICE),vel.to(DEVICE)\n            _,pred=model(seis)\n            val+=loss_fn(pred,vel).item()*V_SCALE*seis.size(0)\n    val/=len(val_loader.dataset); tr=run/cnt\n    plateau.step(val)\n\n    print(f\"→ epoch {ep:2d} | train {tr:6.1f} | val {val:6.1f} m/s \"\n          f\"| lr {opt.param_groups[0]['lr']:.2e}\")\n\n    if val<best:\n        best=val; torch.save(model.state_dict(),f\"{CKPT_DIR}/best.pth\")\n        print(\"  ✔ saved best\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-22T22:12:46.150186Z","iopub.execute_input":"2025-06-22T22:12:46.150708Z","iopub.status.idle":"2025-06-22T22:12:58.195968Z","shell.execute_reply.started":"2025-06-22T22:12:46.150681Z","shell.execute_reply":"2025-06-22T22:12:58.194846Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------- paths & basic config --------------------\nBASE_DIR  = \"/kaggle/input/waveform-inversion/train_samples\"   # ← adjust if needed\nCKPT_DIR1 = \"/kaggle/working/ckpt_stage1\"\nCKPT_DIR2 = \"/kaggle/working/ckpt_stage2\"\nimport os, re, math, random, pickle, numpy as np, torch, torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import LambdaLR\nfrom pathlib import Path\nfrom tqdm import tqdm\n\nfor d in (CKPT_DIR1, CKPT_DIR2): Path(d).mkdir(parents=True, exist_ok=True)\n\nDEVICE    = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ntorch.manual_seed(42); np.random.seed(42); random.seed(42)\n\n# ---------------- dataset ---------------------------------------------------\n_pat_d = re.compile(r'^(data|seis)[_\\-]?')\n_pat_m = re.compile(r'^(model|vel)[_\\-]?')\n\ndef _pairs(dir_):\n    f = [x for x in os.listdir(dir_) if x.endswith('.npy')]\n    seis, vel = {}, {}\n    for x in f:\n        if x.startswith((\"data\", \"seis\")):  seis[_pat_d.sub('', x)] = x\n        else:                              vel [_pat_m.sub('', x)] = x\n    return [(os.path.join(dir_, seis[k]), os.path.join(dir_, vel[k]))\n            for k in sorted(set(seis) & set(vel))]\n\ndef build_index(root):\n    out = []\n    for fam in sorted(os.listdir(root)):\n        fd = os.path.join(root, fam)\n        if not os.path.isdir(fd): continue\n        for s, v in _pairs(fd):\n            n = np.load(s, mmap_mode='r').shape[0]    # real #shots\n            out += [(s, v, i) for i in range(n)]\n    return out\n\nV_SCALE = 4_500.0\nclass VelDS(Dataset):\n    def __init__(self, items): self.items = items\n    def __len__(self): return len(self.items)\n    def __getitem__(self, k):\n        s,v,i = self.items[k]\n        seis = np.load(s, mmap_mode='r')[i].astype(np.float32)    # (5,1000,70)\n        vel  = np.load(v, mmap_mode='r')[i].squeeze().astype(np.float32)\n        return torch.from_numpy(seis), torch.from_numpy(vel / V_SCALE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-24T00:18:40.809402Z","iopub.execute_input":"2025-06-24T00:18:40.809668Z","iopub.status.idle":"2025-06-24T00:18:40.821053Z","shell.execute_reply.started":"2025-06-24T00:18:40.809647Z","shell.execute_reply":"2025-06-24T00:18:40.820267Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------- MetaFormer + coarse UNet -------------------\nPATCH_T, PATCH_X = 32,16; STR_T,STR_X = 4,2\nEMB_D, DEPTH, DROPOUT = 256, 6, 0.1\n\ndef split_fft(x: torch.Tensor) -> torch.Tensor:\n    \"\"\"log-magnitude FFT split into real / imag, sign-preserved.\"\"\"\n    F2  = torch.fft.rfft2(x, dim=(-2, -1))           # (B, 5, 1000, 36)\n    mag = torch.log(torch.abs(F2) + 1e-6)\n    real = torch.sign(F2.real) * mag\n    imag = torch.sign(F2.imag) * mag\n    return torch.cat([real, imag], dim=1)            # (B, 10, 1000, 36)\n\n\nclass _Stem(nn.Module):\n    def __init__(self, in_ch: int, d: int):\n        super().__init__()\n        self.conv = nn.Conv2d(in_ch, d,\n                              kernel_size=(PATCH_T, PATCH_X),\n                              stride=(STR_T,  STR_X))\n        self.ln   = nn.LayerNorm(d)\n        self.drop = nn.Dropout(DROPOUT)\n\n    def forward(self, x):                 # x : (B, C, 1000, 70)\n        t = self.conv(x).flatten(2).transpose(1, 2)   # (B, N, d)\n        return self.drop(self.ln(t))\n\n\nclass PoolMix(nn.Module):\n    def __init__(self, W): super().__init__(); self.W=W\n    def forward(self,x):\n        cls,tok=x[:,:1],x[:,1:]; B,N,C=tok.shape; H=N//self.W\n        y=tok.transpose(1,2).reshape(B,C,H,self.W)\n        y=F.avg_pool2d(y,3,1,1)-y\n        return torch.cat([cls, tok+y.flatten(2).transpose(1,2)],1)\n\nclass MetaBlock(nn.Module):\n    def __init__(self,C,W):\n        super().__init__()\n        self.n1=nn.LayerNorm(C); self.mix=PoolMix(W)\n        self.n2=nn.LayerNorm(C)\n        self.ff=nn.Sequential(nn.Linear(C,4*C), nn.GELU(), nn.Dropout(DROPOUT),\n                              nn.Linear(4*C,C))\n    def forward(self,x): x=x+self.mix(self.n1(x)); return x+self.ff(self.n2(x))\n\nclass Encoder(nn.Module):\n    def __init__(self,d,C,W):\n        super().__init__()\n        self.blks=nn.ModuleList([MetaBlock(C,W) for _ in range(d)])\n    def forward(self,x):\n        for b in self.blks: x=b(x)\n        return x\n\n# --- tiny helper to crop ----------------------------------------------------\ndef _crop(t, size):\n    h,w=t.shape[-2:]; dh=(h-size[0])//2; dw=(w-size[1])//2\n    return t[..., dh:dh+size[0], dw:dw+size[1]]\n\n# --- Shuffle-UNet (coarse) --------------------------------------------------\ndef _blk(cin, cout):\n    return nn.Sequential(\n        nn.Conv2d(cin, cout, 3, 1, 1, bias=False),\n        nn.InstanceNorm2d(cout, affine=True, eps=1e-3),  # <= here\n        nn.GELU()\n)\n\nclass CoarseUNet(nn.Module):\n    def __init__(self, base=32):\n        super().__init__()\n        self.pre = nn.Conv2d(1, base*4, 1)\n        self.normp = nn.LayerNorm([base*4,70,70])\n        self.ps    = nn.PixelShuffle(2)               # → base×140×140\n        self.enc1  = _blk(base,   base)\n        self.pool1 = nn.MaxPool2d(2)                  # 140→70\n        self.enc2  = _blk(base,   base*2)\n        self.pool2 = nn.MaxPool2d(2)                  # 70→35\n        self.enc3  = _blk(base*2, base*4)\n\n        self.up2   = nn.ConvTranspose2d(base*4,base*2,2,2)     # 35→70\n        self.dec2  = _blk(base*4, base*2)\n        self.up1   = nn.ConvTranspose2d(base*2,base,2,2)       # 70→140\n        self.dec1  = _blk(base*2, base)\n        self.outc  = nn.Conv2d(base,1,1)\n\n    def forward(self,c):\n        x  = self.ps(self.normp(self.pre(c)))\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool1(e1))\n        e3 = self.enc3(self.pool2(e2))\n        d2 = self.dec2(torch.cat([self.up2(e3), e2],1))\n        d1 = self.dec1(torch.cat([self.up1(d2), e1],1))\n        out= F.avg_pool2d(self.outc(d1),2)            # back → 70×70\n        return out.squeeze(1)\n\n# --- full stage-1 model -----------------------------------------------------\nclass DualBranch(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.raw  = _Stem(5 , EMB_D)\n        self.fft  = _Stem(10, EMB_D)\n        Wr = (70-PATCH_X)//STR_X+1\n        Wf = (36-PATCH_X)//STR_X+1\n        self.enc_r = Encoder(DEPTH,EMB_D,Wr)\n        self.enc_f = Encoder(DEPTH,EMB_D,Wf)\n        self.fuse  = nn.Sequential(nn.LayerNorm(2*EMB_D),\n                                   nn.Linear(2*EMB_D,4*EMB_D),\n                                   nn.GELU(), nn.Linear(4*EMB_D,70*70))\n        self.coarse_norm = nn.GroupNorm(1,1)\n        self.refine = CoarseUNet()\n\n    def forward(self,s):\n        B = s.size(0)\n        r_tok = self.raw(s)                  # raw branch\n        f_tok = self.fft(split_fft(s))\n        z = torch.zeros(B,1,EMB_D,device=s.device)\n        r = self.enc_r(torch.cat([z,r_tok],1))\n        f = self.enc_f(torch.cat([z,f_tok],1))\n        fused = torch.cat([r[:,1:].mean(1), f[:,1:].mean(1)],1)\n        coarse = self.coarse_norm(self.fuse(fused).view(B,1,70,70))\n        coarse = coarse.squeeze(1)\n        pred   = self.refine(coarse.unsqueeze(1))\n        return coarse, pred                           # both (B,70,70)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-24T00:18:37.454602Z","iopub.execute_input":"2025-06-24T00:18:37.454909Z","iopub.status.idle":"2025-06-24T00:18:37.473848Z","shell.execute_reply.started":"2025-06-24T00:18:37.454888Z","shell.execute_reply":"2025-06-24T00:18:37.473150Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================  Stage-1 TRAINING  ======================\nEPOCHS1     = 60\nBATCH1      = 2\nACCUM1      = 4\nVAL_FRAC1   = 0.20\nLR_BASE1    = 7e-4\nLR_MIN1     = 3e-5\nUNET_RATIO1 = 0.10\nWARM_STEPS1 = 10_000\nWEIGHT_DEC1 = 5e-4\nCLIP1       = 100.\n\n# ----- split index ------------------------------------------------\nfull_idx = build_index(BASE_DIR)\nperm     = np.random.permutation(len(full_idx))\ncut      = int(len(perm)*(1-VAL_FRAC1))\ntr_idx, va_idx = perm[:cut], perm[cut:]\n\ntrain_loader = DataLoader(VelDS([full_idx[i] for i in tr_idx]),\n                          batch_size=BATCH1, shuffle=True,\n                          num_workers=4, pin_memory=True, drop_last=True)\nval_loader   = DataLoader(VelDS([full_idx[i] for i in va_idx]),\n                          batch_size=BATCH1, shuffle=False,\n                          num_workers=4, pin_memory=True)\n\n# ----- model, optimiser, scheduler --------------------------------\ndef augment(seis, p_drop=0.0, amp_jitter=False):\n    if amp_jitter:\n        amp = 1.0 + (torch.rand_like(seis[:, :1, :1, :1]) * 0.10 - 0.05)\n        seis = seis * amp\n    if p_drop > 0:\n        mask = (torch.rand(seis.shape[0], 5, 1, 1, device=seis.device) > p_drop).float()\n        seis = seis * mask\n    return seis\nmodel = DualBranch().to(DEVICE)\n\ndecay, no_wd = [], []\nref_ids = {id(p) for p in model.refine.parameters()}\nfor n, p in model.named_parameters():\n    if id(p) in ref_ids: continue          # refine gets its own LR\n    (no_wd if p.ndim==1 or n.endswith(\".bias\") else decay).append(p)\n\nopt = torch.optim.AdamW(\n    [{\"params\":decay,                         \"lr\":LR_BASE1},\n     {\"params\":no_wd,                         \"lr\":LR_BASE1, \"weight_decay\":0.},\n     {\"params\":model.refine.parameters(),     \"lr\":LR_BASE1*UNET_RATIO1}],\n    betas=(0.9,0.95), weight_decay=WEIGHT_DEC1)\n\ntotal_steps = EPOCHS1 * math.ceil(len(train_loader)/ACCUM1)\nBASE_LR = 4e-4\ndef lr_main(step):\n    if step < 4000:                 # 4 k-step warm-up\n        return (step + 1)/4000\n    prog = (step-4000)/(total_steps-4000)\n    return 0.5 * (1 + math.cos(math.pi * min(prog, 1)))  # cosine\nlr_unet = lambda s: lr_main(s) * UNET_RATIO1\nsched   = LambdaLR(opt, lr_lambda=[lr_main, lr_main, lr_unet])\n\nloss_fn = nn.L1Loss()\nscaler  = GradScaler()\nbest_val = float('inf')\n\n# ----- training loop ----------------------------------------------\nfor ep in range(1, EPOCHS1+1):\n    model.train(); run=cnt=0\n    prog = tqdm(enumerate(train_loader), total=len(train_loader),\n                desc=f\"[Stage-1] Epoch {ep}/{EPOCHS1}\", ncols=95)\n\n    opt.zero_grad(set_to_none=True)\n    for step,(seis,vel) in prog:\n        seis,vel = seis.to(DEVICE), vel.to(DEVICE)\n        seis = augment(\n            seis, \n            p_drop = 0.10 if model.training else 0.10,   # same for val if you keep it\n            amp_jitter = model.training                  # jitter only for train\n        )\n        with autocast(dtype=torch.float16):\n            coarse, pred = model(seis)\n            l_main   = loss_fn(pred,   vel)\n            l_aux    = loss_fn(coarse, vel) * 0.1\n            loss     = (l_main + l_aux) / ACCUM1\n\n        scaler.scale(loss).backward()\n\n        if (step+1)%ACCUM1 == 0:\n            scaler.unscale_(opt)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), CLIP1)\n            scaler.step(opt); scaler.update()\n            opt.zero_grad(set_to_none=True); sched.step()\n\n        run += l_main.item()*V_SCALE*seis.size(0); cnt += seis.size(0)\n        prog.set_postfix(trainMAE=f\"{run/cnt:6.1f}\",\n                         lr=f\"{opt.param_groups[0]['lr']:.1e}\")\n\n    # ---- validation ------------------------------------------------\n    model.eval(); val_sum=0.\n    with torch.no_grad():\n        for seis,vel in val_loader:\n            seis,vel = seis.to(DEVICE), vel.to(DEVICE)\n            _,pred   = model(seis)\n            val_sum += loss_fn(pred, vel).item()*V_SCALE*seis.size(0)\n    val_mae = val_sum / len(val_loader.dataset)\n    print(f\"  ↪ val {val_mae:6.1f} m/s\")\n\n    # ---- save best -------------------------------------------------\n    if val_mae < best_val:\n        best_val = val_mae\n        torch.save(model.state_dict(), f\"{CKPT_DIR1}/best.pth\")\n        print(\"    ✔ saved new best\")\n\nprint(\"Stage-1 training done; best MAE =\", best_val)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-24T00:19:09.556915Z","iopub.execute_input":"2025-06-24T00:19:09.557581Z","iopub.status.idle":"2025-06-24T03:13:14.752345Z","shell.execute_reply.started":"2025-06-24T00:19:09.557559Z","shell.execute_reply":"2025-06-24T03:13:14.750933Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2, math\nimport torch\nimport torch.nn.functional as F\n\n# ─── your existing fd8_misfit ────────────────────────────────────────────────\ndef fd8_misfit(coarse_v, seis_obs, src_ij, rec_ijs,\n               dx=10.0, dt=0.001, n_step=8, alpha=0.3):\n    # coarse_v: (70,70) in km/s\n    # seis_obs: (R, T)\n    # src_ij:   (2,)\n    # rec_ijs:  list of (i,j)\n    v = coarse_v * 1000.0                  # km/s → m/s\n    dt2_v2 = (dt * v)**2\n    p_prev = np.zeros_like(v, np.float32)\n    p_curr = np.zeros_like(v, np.float32)\n    mis_rec = np.zeros(len(rec_ijs), np.float32)\n\n    c1, c2 = 1.0, -1.0/12.0\n    for t in range(n_step):\n        # source Ricker\n        f_t = (1.0 - 2.0*(math.pi*8*(t*dt-0.04))**2) \\\n              * math.exp(-(math.pi*8*(t*dt-0.04))**2)\n        i0,j0 = src_ij\n        p_curr[i0,j0] += f_t\n\n        lap = (\n            -20.0*p_curr\n            + c1*(np.roll(p_curr,1,0)+np.roll(p_curr,-1,0)\n                 + np.roll(p_curr,1,1)+np.roll(p_curr,-1,1))\n            + c2*(np.roll(p_curr,2,0)+np.roll(p_curr,-2,0)\n                 + np.roll(p_curr,2,1)+np.roll(p_curr,-2,1))\n        ) / (dx*dx)\n\n        p_next = 2*p_curr - p_prev + dt2_v2 * lap\n        p_prev, p_curr = p_curr, p_next\n\n        for r,(ir,jr) in enumerate(rec_ijs):\n            mis_rec[r] += abs(p_curr[ir,jr] - seis_obs[r,t])\n\n    mis_map = np.zeros_like(v, np.float32)\n    for r,(ir,jr) in enumerate(rec_ijs):\n        mis_map[ir,jr] = mis_rec[r]\n\n    mis_map = cv2.GaussianBlur(mis_map, (5,5), 0)\n    mis_map = alpha * mis_map / (mis_map.max()+1e-6)\n    return mis_map  # (70,70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-24T03:13:16.967916Z","iopub.execute_input":"2025-06-24T03:13:16.968641Z","iopub.status.idle":"2025-06-24T03:13:16.977049Z","shell.execute_reply.started":"2025-06-24T03:13:16.968616Z","shell.execute_reply":"2025-06-24T03:13:16.976185Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── run once; writes ≈ 500×#families .npy files ─────────────────────────────\nimport os, re, numpy as np, skimage.draw as skdr\nfrom pathlib import Path\nfrom tqdm import tqdm\n\nFAMDIR = \"/kaggle/working\"\n\n# ----- helper: convert shot-id to (row, col) grid coordinates ---------------\nsrc_xy = lambda shot_id: (0, shot_id % 70)          #  ⬅️  CHANGE if needed\nrec_xys = lambda: [(0, j) for j in range(70)]        # 70 surface receivers\n\ndef build_illum_mask(src_ij, rec_ijs, H=70, W=70):\n    mask = np.zeros((H, W), np.float32)\n    si, sj = src_ij\n    for ri, rj in rec_ijs:\n        rr, cc = skdr.line(si, sj, ri, rj)   # Bresenham\n        mask[rr, cc] += 1.0\n    mask /= mask.max() + 1e-6\n    return mask\n\nfor fam in sorted(os.listdir(\"/kaggle/input/waveform-inversion/train_samples\")):\n    famdir = Path(FAMDIR) / fam\n    if not famdir.is_dir(): continue\n\n    out_dir = famdir / \"illum_masks\"\n    out_dir.mkdir(exist_ok=True)\n\n    # assume 500 shots per *.npy file as in the baseline code\n    for shot_id in tqdm(range(500), desc=fam):\n        \n        fn = out_dir / f\"illum_{shot_id:03d}.npy\"\n        if fn.exists(): continue\n\n        mask = build_illum_mask(src_xy(shot_id), rec_xys())\n        np.save(fn, mask)\n        print(f\"Saved: {fn}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-24T03:13:20.175405Z","iopub.execute_input":"2025-06-24T03:13:20.176024Z","iopub.status.idle":"2025-06-24T03:13:20.184408Z","shell.execute_reply.started":"2025-06-24T03:13:20.175996Z","shell.execute_reply":"2025-06-24T03:13:20.183795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T14:19:27.558204Z","iopub.execute_input":"2025-06-23T14:19:27.558525Z","iopub.status.idle":"2025-06-23T14:19:28.274177Z","shell.execute_reply.started":"2025-06-23T14:19:27.558500Z","shell.execute_reply":"2025-06-23T14:19:28.273190Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T14:34:58.234272Z","iopub.execute_input":"2025-06-23T14:34:58.234880Z","iopub.status.idle":"2025-06-23T14:34:58.688610Z","shell.execute_reply.started":"2025-06-23T14:34:58.234860Z","shell.execute_reply":"2025-06-23T14:34:58.687086Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, math\nimport numpy as np\nimport cv2\nfrom pathlib import Path\nfrom tqdm import tqdm\nfrom collections import defaultdict\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import LambdaLR\n\n# ─── Helpers ────────────────────────────────────────────────────────────────\n\ndefault_geom = { \"src\": (35,35),\n                 \"recs\": [(10,10),(10,60),(60,10),(60,60)] }\n# and in your dataset:\nshot_geometry = defaultdict(lambda: default_geom)\n\ndef center_crop(src: torch.Tensor, size: tuple[int,int]):\n    h,w = src.shape[-2:]\n    dh, dw = (h - size[0])//2, (w - size[1])//2\n    return src[..., dh:dh+size[0], dw:dw+size[1]]\n\ndef grad_xy(t: torch.Tensor):\n    gy = t[:,:,1:,:] - t[:,:,:-1,:]\n    gy = F.pad(gy, (0,0,0,1))\n    gx = t[:,:,:,1:] - t[:,:,:,:-1]\n    gx = F.pad(gx, (0,1,0,0))\n    return gx, gy\n\ndef laplacian(t: torch.Tensor):\n    gx, gy = grad_xy(t)\n    gxx,_ = grad_xy(gx)\n    _, gyy = grad_xy(gy)\n    return gxx + gyy\n\ndef fd8_misfit(coarse_v, seis_obs, src_ij, rec_ijs,\n               dx=10.0, dt=0.001, n_step=8, alpha=0.3):\n    H,W = coarse_v.shape\n    v = coarse_v * 1_000.\n    dt2_v2 = (dt * v)**2\n    p_prev = np.zeros((H,W), np.float32)\n    p_curr = np.zeros((H,W), np.float32)\n    mis_rec = np.zeros(len(rec_ijs), np.float32)\n    c1, c2 = 1.0, -1.0/12.0\n\n    for t in range(n_step):\n        # Ricker source\n        f_t = (1 - 2*(math.pi*8*(t*dt-0.04))**2) * math.exp(-(math.pi*8*(t*dt-0.04))**2)\n        i0,j0 = src_ij; p_curr[i0,j0] += f_t\n\n        lap = (\n            -20*p_curr\n            + c1*(np.roll(p_curr,1,0)+np.roll(p_curr,-1,0)\n                 +np.roll(p_curr,1,1)+np.roll(p_curr,-1,1))\n            + c2*(np.roll(p_curr,2,0)+np.roll(p_curr,-2,0)\n                 +np.roll(p_curr,2,1)+np.roll(p_curr,-2,1))\n        )/(dx*dx)\n\n        p_next = 2*p_curr - p_prev + dt2_v2*lap\n        p_prev, p_curr = p_curr, p_next\n\n        for r,(ir,jr) in enumerate(rec_ijs):\n            mis_rec[r] += abs(p_curr[ir,jr] - seis_obs[r, t, jr])\n\n    mis_map = np.zeros_like(coarse_v)\n    for r,(ir,jr) in enumerate(rec_ijs):\n        mis_map[ir,jr] = mis_rec[r]\n    mis_map = cv2.GaussianBlur(mis_map, (5,5), 0)\n    mis_map *= alpha / (mis_map.max()+1e-6)\n    return mis_map\n\n# ─── Dataset (now unpacks 3-tuples) ─────────────────────────────────────────\nclass VelDS2(Dataset):\n    def __init__(self, items, shot_geometry, mask_dir=None):\n        \"\"\"\n        items: list of (sfile, vfile, shot_id)\n        shot_geometry: dict[fam] -> {'src':(i,j), 'recs':[(i,j),...]}\n        \"\"\"\n        self.items = items\n        self.shot_geometry = shot_geometry\n        self.mask_dir = mask_dir\n\n    def __len__(self):\n        return len(self.items)\n\n    def __getitem__(self, idx):\n        sfile, vfile, shot_id = self.items[idx]\n        fam = Path(sfile).parent.name  # derive family from folder name\n\n        # 1) Load and slice numpy arrays\n        seis_np = np.load(sfile, mmap_mode='r')[shot_id].astype(np.float32)  # (R,T,W)\n        vel_np  = np.load(vfile, mmap_mode='r')[shot_id].astype(np.float32)  # (H,W) or (1,H,W)\n\n        # 2) Ensure vel_np is 2D\n        if vel_np.ndim == 3 and vel_np.shape[0] == 1:\n            vel_np = vel_np[0]\n        elif vel_np.ndim != 2:\n            raise ValueError(f\"Unexpected velocity shape: {vel_np.shape}\")\n\n        # 3) Build illumination mask\n        src_ij  = self.shot_geometry[fam]['src']\n        rec_ijs = self.shot_geometry[fam]['recs']\n        if self.mask_dir:\n            mask_file = os.path.join(self.mask_dir, f\"{fam}_{shot_id}.npy\")\n            illum_np  = np.load(mask_file).astype(np.float32)\n        else:\n            from skimage.draw import line\n            m = np.zeros_like(vel_np, dtype=np.float32)\n            si, sj = src_ij\n            for (ri, rj) in rec_ijs:\n                rr, cc = line(si, sj, ri, rj)\n                m[rr, cc] += 1\n            illum_np = m / (m.max() + 1e-6)\n\n        # 4) Convert to torch\n        seis  = torch.from_numpy(seis_np)\n        vel   = torch.from_numpy(vel_np)\n        illum = torch.from_numpy(illum_np)\n\n        return seis, vel, src_ij, rec_ijs, illum\n\n# ─── Load & Freeze Stage-1 ──────────────────────────────────────────────────\ncoarse_net = DualBranch().to(DEVICE)\ncoarse_net.load_state_dict(torch.load(f\"{CKPT_DIR1}/best.pth\", map_location=DEVICE),\n                           strict=False)\ncoarse_net.eval()\nfor p in coarse_net.parameters():\n    p.requires_grad_(False)\n\n# ─── Residual U-Net ─────────────────────────────────────────────────────────\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass RefineUNet(nn.Module):\n    def __init__(self, in_ch=6, base=64):\n        super().__init__()\n        def blk(ci, co):\n            return nn.Sequential(\n                nn.Conv2d(ci, co, 3, padding=1, bias=False),\n                nn.GroupNorm(1, co),\n                nn.GELU()\n            )\n\n        # Encoder\n        self.enc1  = blk(in_ch,   base)      # 70×70\n        self.pool1 = nn.MaxPool2d(2)         # 35×35\n        self.enc2  = blk(base,    base*2)    # 35×35\n        self.pool2 = nn.MaxPool2d(2)         # 17×17\n        self.enc3  = blk(base*2,  base*4)    # 17×17\n\n        # Decoder\n        self.up2   = nn.ConvTranspose2d(base*4, base*2, 2,2)  # → 34×34\n        self.dec2  = blk(base*4,  base*2)\n        self.up1   = nn.ConvTranspose2d(base*2, base,   2,2)  # → 68×68\n        self.dec1  = blk(base*2,  base)\n        self.head  = nn.Conv2d(base,   1,      3, padding=1)  # 68×68 → 68×68\n\n        # Residual scaling parameter, start at 0.1\n        self.gamma = nn.Parameter(torch.full((1,), 0.3))\n\n        # Initialize weights:\n        # - head conv → zeros\n        # - other convs → Kaiming normal (using 'relu' gain)\n        # - norms → identity\n        for module in self.modules():\n            if isinstance(module, nn.Conv2d):\n                if module is self.head:\n                    nn.init.zeros_(module.weight)\n                else:\n                    nn.init.kaiming_normal_(module.weight, nonlinearity='relu')\n            elif isinstance(module, (nn.GroupNorm, nn.BatchNorm2d, nn.LayerNorm)):\n                nn.init.ones_(module.weight)\n                nn.init.zeros_(module.bias)\n\n    def forward(self, x):\n        # Encoder\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool1(e1))\n        e3 = self.enc3(self.pool2(e2))\n\n        # Decoder stage 1\n        u2  = self.up2(e3)\n        e2c = center_crop(e2, u2.shape[-2:])\n        d2  = self.dec2(torch.cat([u2, e2c], dim=1))\n\n        # Decoder stage 2\n        u1  = self.up1(d2)\n        e1c = center_crop(e1, u1.shape[-2:])\n        d1  = self.dec1(torch.cat([u1, e1c], dim=1))\n\n        # Residual head\n        delta = self.head(d1)                   # (B,1,68,68)\n        delta = F.pad(delta, (1,1,1,1))         # → (B,1,70,70)\n        return delta\n\n\n# ─── Build DataLoaders ─────────────────────────────────────────────────────\nidx_all  = build_index(BASE_DIR)            # returns list[(sfile,vfile,shot_id)]\nperm     = np.random.permutation(len(idx_all))\ncut       = int(0.8 * len(idx_all))\ntr_ds     = VelDS2([idx_all[i] for i in perm[:cut]], shot_geometry)\nva_ds     = VelDS2([idx_all[i] for i in perm[cut:]], shot_geometry)\n\ntr_loader = DataLoader(tr_ds, batch_size=2, shuffle=True,  num_workers=4, pin_memory=True)\nva_loader = DataLoader(va_ds, batch_size=2, shuffle=False, num_workers=4, pin_memory=True)\n\n# ─── Training Setup ────────────────────────────────────────────────────────\nimport math\n\n# ─── Training Setup ────────────────────────────────────────────────────────\nres_net = RefineUNet(in_ch=6, base=32).to(DEVICE)\n\n# two param groups: one for gamma, one for all the rest\nopt_res = torch.optim.AdamW([\n    # Group 0: gamma, high LR, no weight decay\n    {\n      'params': [res_net.gamma],\n      'lr': 1e-3,\n      'weight_decay': 0.0\n    },\n    # Group 1: all other params, standard LR & decay\n    {\n      'params': [p for p in res_net.parameters() if p is not res_net.gamma],\n      'lr': 5e-4,\n      'weight_decay': 1e-4\n    }\n])\n# Warmup + Cosine Decay Scheduler\nnum_epochs = 25\nsteps_per_epoch = len(tr_loader)\ntotal_steps = num_epochs * steps_per_epoch\nwarmup_steps = 4000           # 10% of training for warmup\n\ndef lr_lambda(step):\n    if step < warmup_steps:\n        # linear warmup from 0 → 1\n        return float(step) / float(max(1, warmup_steps))\n    # cosine decay from 1 → 0\n    progress = float(step - warmup_steps) / float(max(1, total_steps - warmup_steps))\n    return 0.5 * (1.0 + math.cos(math.pi * progress))\n\nsched_res = LambdaLR(opt_res, lr_lambda=lr_lambda)\nloss_fn   = nn.L1Loss()\nscaler    = GradScaler()\nV_SCALE   = 4_500.0\n\nbest_val = float('inf')\ntrain_mae = 0\nfor ep in range(1, 40):\n    # — train —\n    res_net.train()\n    train_sum = n_train = 0\n    i = 0\n    for seis, vel, srcs, recs, illum in tqdm(tr_loader, desc=f\"[res] Ep{ep:2d}, [train MAE]: {train_mae}\"):\n        seis, vel, illum = seis.to(DEVICE), vel.to(DEVICE), illum.to(DEVICE)\n        vel /= V_SCALE\n        if i % 500 == 0:\n            lrs = [pg['lr'] for pg in opt_res.param_groups]\n            print(f\"γ-LR = {lrs[0]:.2e}, conv-LR = {lrs[1]:.2e}\")\n        i += 1\n        with torch.no_grad():\n            coarse, _ = coarse_net(seis)     # (B,70,70)\n\n        # physics misfit\n        phys_list = []\n        c_np, s_np = coarse.cpu().numpy(), seis.cpu().numpy()\n        for b in range(coarse.size(0)):\n            pm = fd8_misfit(c_np[b], s_np[b], srcs[b], recs[b])\n            phys_list.append(torch.from_numpy(pm))\n        phys = torch.stack(phys_list,0).unsqueeze(1).to(DEVICE)\n\n        # grads & lap\n        c4 = coarse.unsqueeze(1)\n        gx, gy = grad_xy(c4)\n        lap = laplacian(c4)\n\n        inp = torch.cat([c4, gx, gy, lap, phys, illum.unsqueeze(1)],1)\n        with autocast():\n            delta = res_net(inp)               # (B,1,70,70)\n            pred  = coarse + res_net.gamma * delta.squeeze(1)\n            loss  = loss_fn(pred, vel)\n\n        scaler.scale(loss).backward()\n        scaler.step(opt_res); scaler.update()\n        opt_res.zero_grad(); sched_res.step()\n\n        train_sum += loss.item() * V_SCALE * vel.size(0)\n        n_train   += vel.size(0)\n\n    train_mae = train_sum / n_train\n\n    # — validate —\n    res_net.eval()\n    val_sum = n_val = 0\n    with torch.no_grad():\n        for seis, vel, srcs, recs, illum in va_loader:\n            seis, vel, illum = seis.to(DEVICE), vel.to(DEVICE), illum.to(DEVICE)\n            coarse, _ = coarse_net(seis)\n            vel /= V_SCALE\n            phys_list = []\n            c_np, s_np = coarse.cpu().numpy(), seis.cpu().numpy()\n            for b in range(coarse.size(0)):\n                pm = fd8_misfit(c_np[b], s_np[b], srcs[b], recs[b])\n                phys_list.append(torch.from_numpy(pm))\n            phys = torch.stack(phys_list,0).unsqueeze(1).to(DEVICE)\n\n            c4 = coarse.unsqueeze(1)\n            gx, gy = grad_xy(c4)\n            lap = laplacian(c4)\n            inp = torch.cat([c4, gx, gy, lap, phys, illum.unsqueeze(1)],1)\n\n            delta = res_net(inp)\n            pred  = coarse + res_net.gamma * delta.squeeze(1)\n            val_sum += loss_fn(pred, vel).item() * V_SCALE * vel.size(0)\n            n_val   += vel.size(0)\n\n    val_mae = val_sum / n_val\n    print(f\"[res] Ep{ep:2d} | train {train_mae:6.1f} | val {val_mae:6.1f} | γ={res_net.gamma.item():.3f}\")\n\n    if val_mae < best_val:\n        best_val = val_mae\n        torch.save(res_net.state_dict(), f\"{CKPT_DIR2}/refine_best.pth\")\n        print(f\"→ Saved new best  (val {best_val:6.1f})\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-24T04:58:53.186663Z","iopub.execute_input":"2025-06-24T04:58:53.187273Z"}},"outputs":[],"execution_count":null}]}