{"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"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# waveform_dataset_fft.py\nimport numpy as np\nimport scipy.signal as sig\nfrom scipy.ndimage import laplace\nimport torch\nfrom torch.utils.data import Dataset\nfrom pathlib import Path\nfrom typing import List, Tuple\n\n\nclass WaveformDatasetFFT(Dataset):\n    \"\"\"\n    __getitem__ returns:\n        x : (S, 10, T, R)  – per source, ten feature maps\n        y : (H,  W)        – velocity map (metres s⁻¹)\n    Channels: raw, Re, Im, log|Z|, d/dt, d/dx, ∇², envelope, sinΦ, cosΦ\n    \"\"\"\n\n    def __init__(self, root: str,\n                 fft_norm: str = \"ortho\",\n                 log_eps:   float = 1e-9):\n        self.fft_norm, self.log_eps = fft_norm, log_eps\n        root = Path(root)\n        if not root.exists():\n            raise FileNotFoundError(root)\n\n        # ---------------- catalogue every sample -----------------------\n        self.index: List[Tuple[Path, Path, int]] = []   # (wave, vel, local_idx)\n\n        for fam in [p for p in root.iterdir() if p.is_dir()]:\n            data_dir, model_dir = fam / \"data\", fam / \"model\"\n\n            # Vel / Style families\n            if data_dir.is_dir() and model_dir.is_dir():\n                for w in sorted(data_dir.glob(\"*.npy\")):\n                    v = model_dir / w.name.replace(\"data\", \"model\")\n                    if not v.exists():\n                        raise FileNotFoundError(v)\n                    n = np.load(w, mmap_mode=\"r\").shape[0]\n                    for i in range(n):\n                        self.index.append((w, v, i))\n                continue\n\n            # Fault families\n            for w in sorted(fam.glob(\"seis*.npy\")):\n                v = w.with_name(w.name.replace(\"seis\", \"vel\"))\n                if not v.exists():\n                    raise FileNotFoundError(v)\n                n = np.load(w, mmap_mode=\"r\").shape[0]\n                for i in range(n):\n                    self.index.append((w, v, i))\n\n        if not self.index:\n            raise RuntimeError(f\"No samples found under {root}\")\n\n    # ------------------------------------------------------------------\n    def _feature_stack(self, gather: np.ndarray) -> torch.Tensor:\n        \"\"\"\n        gather : (S, T, R) → tensor (S, 10, T, R)\n        \"\"\"\n        S, T, R = gather.shape\n        feats = []\n        for g in gather:                                  # loop over sources\n            # FFT-based channels\n            Z     = np.fft.fft2(g, norm=self.fft_norm)\n            reZ   = np.real(Z)\n            imZ   = np.imag(Z)\n            lmag  = np.log(np.abs(Z) + self.log_eps)\n\n            # Derivatives / local operators\n            d_t   = sig.convolve2d(g,  [[-1],[0],[1]], mode=\"same\")   # ∂/∂t\n            d_x   = sig.convolve2d(g,  [[-1,0,1]],      mode=\"same\")  # ∂/∂x\n            lap   = laplace(g, mode=\"nearest\")                        # ∇²\n\n            # Analytic-signal envelope & phase\n            hilb  = sig.hilbert(g, axis=0)\n            env   = np.abs(hilb)\n            phase = np.angle(hilb)\n            sinp, cosp = np.sin(phase), np.cos(phase)\n\n            feats.append(\n                np.stack([g, reZ, imZ, lmag,\n                          d_t, d_x, lap,\n                          env, sinp, cosp], axis=0)\n            )                                             # (10,T,R)\n\n        return torch.from_numpy(np.stack(feats)).float()  # (S,10,T,R)\n\n    # ------------------------------------------------------------------\n    def __len__(self) -> int:\n        return len(self.index)\n\n    def __getitem__(self, idx: int):\n        w_path, v_path, i = self.index[idx]\n\n        gather = np.load(w_path, mmap_mode=\"r\")[i]          # (S,T,R)\n        vel    = np.load(v_path,  mmap_mode=\"r\")[i].copy()  # (H,W) or (1,H,W)\n        vel    = np.squeeze(vel)                            # ensure (H,W)\n\n        x = self._feature_stack(gather)                     # (S,10,T,R)\n        y = torch.from_numpy(vel).float()\n\n        return x, y\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-11T05:53:39.391382Z","iopub.execute_input":"2025-06-11T05:53:39.391602Z","iopub.status.idle":"2025-06-11T05:53:45.129317Z","shell.execute_reply.started":"2025-06-11T05:53:39.391575Z","shell.execute_reply":"2025-06-11T05:53:45.128547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model_fft_fusion.py\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\n# ───────────────── encodes one source's 10-channel image ──────────────\nclass SmallSourceEncoder(nn.Module):\n    def __init__(self, in_ch: int = 10, base: int = 32):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(in_ch, base,      3, padding=1), nn.ReLU(),\n            nn.Conv2d(base,  base * 2,  3, stride=2, padding=1), nn.ReLU(),\n            nn.Conv2d(base*2, base * 4, 3, stride=2, padding=1), nn.ReLU(),\n        )\n\n    def forward(self, x):            # (B,10,T,R)\n        return self.net(x)           # (B,base*4,T/4,R/4)\n\n\n# ───────────────── simple up-conv decoder ─────────────────────────────\nclass FusionDecoder(nn.Module):\n    def __init__(self, in_ch: int, out_ch: int = 4, base: int = 64):\n        super().__init__()\n        self.up1 = nn.ConvTranspose2d(in_ch, base*4, 2, stride=2)\n        self.up2 = nn.ConvTranspose2d(base*4, base*2, 2, stride=2)\n        self.conv_final = nn.Conv2d(base*2, out_ch, 3, padding=1)\n\n    def forward(self, x):\n        x = F.relu(self.up1(x))\n        x = F.relu(self.up2(x))\n        return self.conv_final(x)    # (B,out_ch,H′,W′)\n\n\n# ───────────────── full model ─────────────────────────────────────────\nclass FFTNetLateFusion(nn.Module):\n    \"\"\"\n    Input  : (B,S,10,T,R)\n    Output : (B,70,70)\n    \"\"\"\n    def __init__(self, base: int = 32):\n        super().__init__()\n        self.enc = SmallSourceEncoder(in_ch=10, base=base)\n        self.dec = FusionDecoder(in_ch=base*4, out_ch=4, base=base)\n\n    def forward(self, x):\n        B, S, C, T, R = x.shape\n\n        # Encode each source separately (shared weights)\n        x = x.view(-1, C, T, R)           # (B*S,10,T,R)\n        feat = self.enc(x)                # (B*S,C′,T/4,R/4)\n\n        # Fuse across sources by mean (replace with attention later)\n        _, C2, H2, W2 = feat.shape\n        feat = feat.view(B, S, C2, H2, W2).mean(dim=1)  # (B,C′,H2,W2)\n\n        # Decode & project to 70×70 velocity plane\n        out = self.dec(feat)              # (B,4,H′,W′)\n        out = F.adaptive_avg_pool2d(out, (70, 70))  # (B,4,70,70)\n        out = out.mean(dim=1)             # (B,70,70)\n\n        return out\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-11T05:59:32.860699Z","iopub.execute_input":"2025-06-11T05:59:32.861198Z","iopub.status.idle":"2025-06-11T05:59:32.871257Z","shell.execute_reply.started":"2025-06-11T05:59:32.861175Z","shell.execute_reply":"2025-06-11T05:59:32.870549Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_fft_mae.py\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\n\n# ── CONFIG ────────────────────────────────────────────────────────────\nROOT    = \"/kaggle/input/waveform-inversion/train_samples\"   # edit as needed\nBATCH   = 4\nEPOCHS  = 30\nLR      = 2e-4            # slightly lower for richer input\nWORKERS = 4\n\n# ── DATA ──────────────────────────────────────────────────────────────\nds = WaveformDatasetFFT(ROOT)\ndl = DataLoader(ds, batch_size=BATCH, shuffle=True,\n                num_workers=WORKERS, pin_memory=True)\n\n# ── MODEL ─────────────────────────────────────────────────────────────\nnet  = FFTNetLateFusion(base=32).cuda()\nopt  = optim.AdamW(net.parameters(), lr=LR)\nloss_fn = nn.L1Loss()\n\n# ── TRAIN ─────────────────────────────────────────────────────────────\nfor epoch in tqdm(range(1, EPOCHS + 1)):\n    net.train()\n    running = 0.0\n    for x, y in tqdm(dl, desc=f\"Epoch: {epoch}\"):\n        x = x.cuda(non_blocking=True)\n        y = y.cuda(non_blocking=True)           # (B,70,70)\n\n        pred = net(x)                           # (B,70,70)\n\n        # competition slice – keep odd columns\n        pred_sub = pred[:, :, 1::2]\n        y_sub    = y[:,  :, 1::2]\n\n        loss = loss_fn(pred_sub, y_sub)\n\n        opt.zero_grad()\n        loss.backward()\n        opt.step()\n\n        running += loss.item() * x.size(0)\n\n    mae = running / len(ds)\n    print(f\"Epoch {epoch:02d} · MAE {mae:8.3f} m s⁻¹\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-11T06:00:30.857829Z","iopub.execute_input":"2025-06-11T06:00:30.858104Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"pred\", pred.shape)  # torch.Size([B, 70, 70])\nprint(\"y   \", y.shape)     # should print torch.Size([B, 1, 70, 70]) or [B,70,70]\n\npred_sub = pred[:, :, 1::2]\ny        = y.squeeze(1) if y.dim() == 4 else y\ny_sub    = y[:, :, 1::2]\n\nprint(\"sub shapes\", pred_sub.shape, y_sub.shape)  # both (B,70,35)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}