{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport os\nfrom pathlib import Path\nfrom torch.utils.data import Dataset, DataLoader\nfrom pytorch_msssim import SSIM\nimport time\nimport torch.optim as optim\nimport random\nfrom torch.cuda.amp import GradScaler, autocast\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torch.utils.tensorboard import SummaryWriter\n\n\n# ========================\n# 0. Dataset Utility Functions\n# ========================\ndef inputs_files_to_output_files(input_files):\n    return [\n        Path(str(f).replace('seis', 'vel').replace('data', 'model'))\n        for f in input_files\n    ]\n\n\ndef get_train_files(data_path):\n    all_inputs = [\n        f\n        for f in\n        Path(data_path).rglob('*.npy')\n        if ('seis' in f.stem) or ('data' in f.stem)\n    ]\n    all_outputs = inputs_files_to_output_files(all_inputs)\n    assert all(f.exists() for f in all_outputs)\n    return all_inputs, all_outputs\n\n\nclass SeismicDataset(Dataset):\n    def __init__(self, inputs_files, output_files, n_examples_per_file=500):\n        assert len(inputs_files) == len(output_files)\n        self.inputs_files = inputs_files\n        self.output_files = output_files\n        self.n_examples_per_file = n_examples_per_file\n\n    def __len__(self):\n        return len(self.inputs_files) * self.n_examples_per_file\n\n    def __getitem__(self, idx):\n        file_idx = idx // self.n_examples_per_file\n        sample_idx = idx % self.n_examples_per_file\n\n        # 1. Load data as writeable copy\n        X = np.load(self.inputs_files[file_idx], mmap_mode='r')\n        y = np.load(self.output_files[file_idx], mmap_mode='r')\n\n        # 2. Numerical stability handling\n        # if np.any(np.isnan(X)) or np.any(np.isinf(X)):\n        #     X = np.nan_to_num(X, nan=0.0, posinf=1e6, neginf=-1e6)\n\n        # # 3. Data augmentation (fixing shape issues)\n        # if random.random() > 0.7:\n        #     safe_std = min(0.05 * X.std(), 1000)  # Limit noise amplitude\n        #     X += np.random.normal(0, safe_std, X.shape).astype(np.float32)\n\n        # if random.random() > 0.8:\n        #     drop_mask = np.zeros(X.shape[2])  # Time dimension\n        #     drop_mask[:int(0.3 * X.shape[2])] = 1\n        #     np.random.shuffle(drop_mask)\n        #     X *= drop_mask[None, None, :, None]  # Critical fix: add receiver dimension\n\n        # 4. Return data\n        data = X[sample_idx]  # [n_shots, nt, ng]\n        data = data.transpose(0, 2, 1)  # [n_shots, ng, nt]\n        return data, y[sample_idx]\n\n\n# ========================\n# 1. Global Configuration\n# ========================\nclass Config:\n    # Data dimensions (adapted for OpenFWI dataset)\n    n_shots = 5  # Number of shots\n    nt = 1000  # Time samples\n    ng = 70  # Number of receivers\n    nx, nz = 70, 70  # Velocity model dimensions\n    vp_min, vp_max = 1500, 4500  # Velocity range\n\n    # MRFT multi-scale parameters [8,6](@ref)\n    # Reduce FFT parameters to fit actual data\n    mrft_scales = [32, 64]  # Must be ≤ actual time dimension length\n    mrft_hop_ratios = [4, 8]\n\n    # Training parameters\n    batch_size = 16\n    num_epochs = 30\n    learning_rate = 5e-4\n    alpha, beta, gamma, delta = 0.5, 0.15, 0.15, 0.2  # Weight adjustments\n    phys_weight = 0.3  # Physics consistency loss weight\n    valid_frac = 16  # Validation sampling interval\n    train_frac = 2  # Training sample fraction\n    warmup_epochs = 5  # Velocity decoder pre-training epochs\n\n    # Mixed precision\n    use_amp = True\n\n    # Paths\n    log_dir = \"logs_optimized\"\n    model_save_path = \"optimized_model.pth\"\n    data_path = \"FWI\"  # Dataset path\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n    proj_channels = 64  # Projection channels\n\n\nconfig = Config()\n\n# Set random seeds\ntorch.manual_seed(100)\nnp.random.seed(100)\nrandom.seed(100)\n\n# ========================\n# 1. MRFT Multi-Resolution Spectrogram Module\n# ========================\nclass MultiResolutionSpectrogram(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.scales = config.mrft_scales\n        self.hop_ratios = config.mrft_hop_ratios\n        for wlen in self.scales:\n            window = torch.hann_window(wlen)\n            self.register_buffer(f'window_{wlen}', window)\n\n    def forward(self, shots):\n        B, n_shots, ng, T = shots.shape\n        shots_flat = shots.permute(0,1,3,2).reshape(-1, T)\n        all_specs = []\n        target_size = None\n        for i, wlen in enumerate(self.scales):\n            hop = wlen // self.hop_ratios[i]\n            window = getattr(self, f'window_{wlen}')\n            spec = torch.stft(\n                shots_flat, n_fft=wlen, hop_length=hop,\n                win_length=wlen, window=window,\n                return_complex=True, center=False\n            )\n            amp = spec.abs().clamp(min=1e-6)\n            log_spec = torch.log(amp)\n            log_spec = log_spec.view(B, n_shots*ng, *log_spec.shape[-2:])\n            if target_size is None:\n                Freq, Time = spec.shape[-2], (T - wlen)//hop + 1\n                target_size = (Freq, Time)\n            log_spec = F.interpolate(log_spec, size=target_size,\n                                     mode='bilinear', align_corners=False)\n            all_specs.append(log_spec)\n        return torch.cat(all_specs, dim=1)\n\n# ========================\n# 2. Physics Layer (Reflectivity)\n# ========================\nclass PhysicalLayer(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.mrft = MultiResolutionSpectrogram()\n\n    def forward(self, vp):\n        # vp: [B,1,H,W]\n        vp_curr = vp[..., :-1]     # [B,1,H,W-1]\n        vp_next = vp[..., 1:]      # [B,1,H,W-1]\n        refl = (vp_next - vp_curr) / (vp_next + vp_curr + 1e-6)  # [B,1,H,W-1]\n        # Remove channel dim and treat H as shots dimension\n        refl = refl.squeeze(1)     # [B, H, W-1]\n        # Reshape to [B, n_shots=1, ng=H, nt=W-1]\n        shots = refl.unsqueeze(1)  # [B,1,H,W-1]\n        return self.mrft(shots)\n\n# ========================\n# 3. Main Network (Core modification: MRFT replaces STFT + 1x1 projection)\n# ========================\nclass PhysicsDrivenNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.mrft_layer = MultiResolutionSpectrogram()\n        in_ch = len(config.mrft_scales) * config.n_shots * config.ng\n        proj_ch = config.proj_channels  # e.g. 64\n        self.input_proj = nn.Conv2d(in_ch, proj_ch, kernel_size=1)\n        self.shared_unet = SharedUNet(\n            in_ch=proj_ch, base=32, depth=4\n        )\n        self.velocity_decoder = nn.Sequential(\n            ResBlock(1, 64), ResBlock(64,128), ResBlock(128,256),\n            ResBlock(256,128), ResBlock(128,64), nn.Conv2d(64,1,3,padding=1), nn.Sigmoid()\n        )\n        self.physical_layer = PhysicalLayer()\n\n    def forward(self, shots):\n        if shots.dim() == 3:\n            shots = shots.unsqueeze(0)\n\n        B, n_shots, D2, D3 = shots.shape\n        # Accept either [..., ng, nt] or [..., nt, ng]\n        if D2 == config.ng and D3 == config.nt:\n            # already [B, n_shots, ng, nt]\n            pass\n        elif D2 == config.nt and D3 == config.ng:\n            # swap to [ng, nt]\n            shots = shots.permute(0, 1, 3, 2)\n        else:\n            raise ValueError(f\"Expected shots shape [B, n_shots, ng, nt] or [B, n_shots, nt, ng], got {shots.shape}\")\n\n        # Compute multi-scale spectrogram\n        log_spec = self.mrft_layer(shots)\n\n        # Channel projection\n        feat_in = self.input_proj(log_spec)\n\n        # Shared feature extraction\n        logR_pred, feat = self.shared_unet(feat_in)\n\n        # Velocity prediction\n        vp_norm = self.velocity_decoder(logR_pred)\n        vp_norm = F.interpolate(\n            vp_norm,\n            size=(config.nx, config.nz),\n            mode='bilinear',\n            align_corners=False\n        )\n\n        # Denormalize\n        vp_denorm = vp_norm * (config.vp_max - config.vp_min) + config.vp_min\n\n        # Physical consistency\n        logR_phys = self.physical_layer(vp_denorm)\n\n        return {\n            'vp_pred': vp_denorm,\n            'logR_pred': logR_pred,\n            'logR_phys': logR_phys\n        }\n\n# ========================\n# 4. Attention and Residual Blocks\n# ========================\nclass NonLocalBlock(nn.Module):\n    \"\"\"Non-local attention module[1](@ref) for long-range dependency modeling\"\"\"\n\n    def __init__(self, in_ch, reduction=2):\n        super().__init__()\n        self.in_ch = in_ch\n        self.reduction = reduction\n        self.mid_ch = max(1, in_ch // reduction)  # Ensure at least 1 channel\n\n        # Channel compression\n        self.conv_in = nn.Conv2d(in_ch, self.mid_ch, 1)\n\n        # Three branches\n        self.theta = nn.Conv2d(self.mid_ch, self.mid_ch // 8, 1)\n        self.phi = nn.Conv2d(self.mid_ch, self.mid_ch // 8, 1)\n        self.g = nn.Conv2d(self.mid_ch, self.mid_ch // 8, 1)\n\n        # Output layers\n        self.out_conv = nn.Sequential(\n            nn.Conv2d(self.mid_ch // 8, in_ch, 1),\n            nn.BatchNorm2d(in_ch)\n        )\n        self.gamma = nn.Parameter(torch.zeros(1))\n\n    def forward(self, x):\n        identity = x\n        B, C, H, W = x.shape\n\n        # Compress channels\n        x_red = F.relu(self.conv_in(x))\n\n        # Compute three branches\n        theta = self.theta(x_red).view(B, -1, H * W).permute(0, 2, 1)  # [B, HW, C']\n        phi = self.phi(x_red).view(B, -1, H * W)  # [B, C', HW]\n        g = self.g(x_red).view(B, -1, H * W)  # [B, C', HW]\n\n        # Attention map\n        attn = torch.bmm(theta, phi)  # [B, HW, HW]\n        attn = F.softmax(attn, dim=-1)\n\n        # Weighted aggregation\n        out = torch.bmm(g, attn.transpose(1, 2))  # [B, C', HW]\n        out = out.view(B, -1, H, W)  # [B, C', H, W]\n\n        # Residual connection\n        out = self.out_conv(out)\n        return identity + self.gamma * out\n\n\nclass SCSE(nn.Module):\n    \"\"\"SCSE attention module: Learns feature importance along spatial and channel dimensions\"\"\"\n\n    def __init__(self, in_ch, reduction=16):\n        super().__init__()\n        # Channel attention branch (cSE)\n        self.cSE = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(in_ch, max(1, in_ch // reduction), 1),  # Ensure at least 1 channel\n            nn.ReLU(inplace=True),\n            nn.Conv2d(max(1, in_ch // reduction), in_ch, 1),\n            nn.Sigmoid()\n        )\n        # Spatial attention branch (sSE)\n        self.sSE = nn.Sequential(\n            nn.Conv2d(in_ch, 1, kernel_size=1, stride=1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        # Channel attention weights\n        cSE_weight = self.cSE(x)\n        # Spatial attention weights\n        sSE_weight = self.sSE(x)\n        # Combine both attention mechanisms\n        return x * cSE_weight + x * sSE_weight\n\n\nclass ResBlock(nn.Module):\n    def __init__(self, in_ch, out_ch, use_nl=False):\n        super().__init__()\n        mid_ch = out_ch\n\n        self.conv1 = nn.Conv2d(in_ch, mid_ch, 3, padding=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(mid_ch)\n\n        self.conv2 = nn.Conv2d(mid_ch, mid_ch, 3, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(mid_ch)\n\n        # Use SCSE instead of channel attention\n        self.scse = SCSE(mid_ch)\n\n        # Non-local attention[1](@ref)\n        self.nl = NonLocalBlock(mid_ch) if use_nl else None\n\n        # Shortcut connection\n        if in_ch != out_ch:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_ch, out_ch, 1, bias=False),\n                nn.BatchNorm2d(out_ch)\n            )\n        else:\n            self.shortcut = nn.Identity()\n\n    def forward(self, x):\n        identity = self.shortcut(x)\n\n        out = F.leaky_relu(self.bn1(self.conv1(x)), 0.2)\n        out = self.bn2(self.conv2(out))\n        # Apply SCSE attention\n        out = self.scse(out)\n\n        if self.nl:\n            out = self.nl(out)\n\n        out += identity\n        return F.leaky_relu(out, 0.2)\n\n\n# ========================\n# 5. Shared U-Net Feature Extraction (Added SCSE in bottleneck)\n# ========================\nclass SharedUNet(nn.Module):\n    def __init__(self, in_ch, base=64, depth=4):\n        super().__init__()\n        self.depth = depth\n\n        # 1) Initial convolution\n        self.init_conv = nn.Sequential(\n            nn.Conv2d(in_ch, base, 3, padding=1, bias=False),\n            nn.BatchNorm2d(base),\n            nn.LeakyReLU(0.2)\n        )\n\n        # 2) Encoder\n        self.encoders = nn.ModuleList()\n        self.downsamplers = nn.ModuleList()\n        self.channels = []\n        current_ch = base\n\n        for i in range(depth):\n            next_ch = min(current_ch * 2, 512)  # Limit to 512 channels\n            use_nl = (i == depth - 1)  # Use NL block only in last encoder\n\n            self.encoders.append(ResBlock(current_ch, next_ch, use_nl))\n            self.downsamplers.append(nn.Conv2d(next_ch, next_ch, 3, 2, 1))\n            self.channels.append(next_ch)\n            current_ch = next_ch\n\n        # 3) Bottleneck - Add SCSE here\n        self.bottleneck = nn.Sequential(\n            ResBlock(current_ch, current_ch, True),\n            SCSE(current_ch),  # Add SCSE module\n            ResBlock(current_ch, current_ch, True)\n        )\n\n        # 4) Decoder\n        self.upsamplers = nn.ModuleList()\n        self.decoders = nn.ModuleList()\n\n        for i in range(depth):\n            skip_ch = self.channels[-(i + 1)]\n            in_ch = current_ch\n\n            self.upsamplers.append(nn.ConvTranspose2d(in_ch, in_ch // 2, 2, 2))\n\n            self.decoders.append(\n                ResBlock((in_ch // 2) + skip_ch, in_ch // 2)\n            )\n            current_ch = in_ch // 2\n\n        # 5) Output\n        self.output_conv = nn.Conv2d(current_ch, 2, 1)\n\n    def forward(self, x):\n        x = self.init_conv(x)\n        skips = []\n\n        # Encoder\n        for enc, down in zip(self.encoders, self.downsamplers):\n            x = enc(x)\n            skips.append(x)\n            x = down(x)\n\n        # Bottleneck\n        x = self.bottleneck(x)\n\n        # Decoder\n        for up, dec in zip(self.upsamplers, self.decoders):\n            x = up(x)\n            skip = skips.pop()\n\n            # Adjust spatial dimensions if needed\n            if x.size(2) != skip.size(2) or x.size(3) != skip.size(3):\n                x = F.interpolate(x, skip.size()[2:], mode='bilinear', align_corners=False)\n\n            x = torch.cat([x, skip], dim=1)\n            x = dec(x)\n\n        logR_pred, feat = torch.split(self.output_conv(x), 1, dim=1)\n        return logR_pred, feat\n\n# ========================\n# 6. Hybrid Loss (Optimized)\n# ========================\nclass PhysicsLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.mse = nn.MSELoss()\n        self.mae = nn.L1Loss()\n        self.l1 = nn.L1Loss()\n        self.ssim = SSIM(data_range=4500.0, win_size=7, size_average=True, channel=1)\n        # Depthwise Sobel filters\n        kx = torch.tensor([[[[1,0,-1],[2,0,-2],[1,0,-1]]]], dtype=torch.float32)\n        ky = torch.tensor([[[[1,2,1],[0,0,0],[-1,-2,-1]]]], dtype=torch.float32)\n        self.register_buffer('sobel_x', kx)\n        self.register_buffer('sobel_y', ky)\n\n    def compute_edges(self, x):\n        # x: [B,C,H,W]\n        C = x.size(1)\n        wx = self.sobel_x.to(x.dtype).to(x.device).repeat(C,1,1,1)\n        wy = self.sobel_y.to(x.dtype).to(x.device).repeat(C,1,1,1)\n        gx = F.conv2d(x, wx, padding=1, groups=C)\n        gy = F.conv2d(x, wy, padding=1, groups=C)\n        # Compute gradient magnitude\n        return torch.sqrt(gx.pow(2) + gy.pow(2))\n\n    def forward(self, pred, target, logR_pred, logR_phys):\n        # pred/target: [B,1,H,W]\n        pred = pred.squeeze(2) if pred.dim()==5 else pred\n        target = target.squeeze(2) if target.dim()==5 else target\n        recon = self.mse(pred, target)\n        mae_l = self.mae(pred, target)\n        edge_p = self.compute_edges(pred)\n        edge_t = self.compute_edges(target)\n        edge_l = self.l1(edge_p, edge_t)\n        # Compute SSIM in full precision to avoid dtype mismatch\n        # Temporarily disable AMP for SSIM and cast to float\n        with autocast(enabled=False):\n            ssim_val = self.ssim(pred.float(), target.float())\n        ssim_val = ssim_val.to(pred.dtype)\n        ssim_l = 1 - ssim_val\n        # Reduce to single channel then resize to match logR_pred\n        phys_map = logR_phys.mean(dim=1, keepdim=True).detach()\n        # Interpolate physical spectrogram to match prediction size\n        phys_map = F.interpolate(\n            phys_map,\n            size=logR_pred.shape[-2:],\n            mode='bilinear', align_corners=False\n        )\n        phys_l = F.mse_loss(logR_pred, phys_map)\n        total = (config.alpha*recon + config.delta*mae_l +\n                 config.beta*edge_l + config.gamma*ssim_l +\n                 config.phys_weight*phys_l)\n        return total, {\n            'recon': recon.item(),\n            'mae': mae_l.item(),\n            'edge': edge_l.item(),\n            'ssim': ssim_l.item(),\n            'phys': phys_l.item()\n        }\n\n\n# ========================\n# 7. Training Process (Fixed metrics key mismatch)\n# ========================\ndef train():\n    device = config.device\n    model = PhysicsDrivenNet().to(device)\n    optimizer = optim.AdamW(model.parameters(), lr=config.learning_rate, weight_decay=1e-3)\n    scheduler = CosineAnnealingLR(optimizer, T_max=config.num_epochs//5, eta_min=1e-5)\n    criterion = PhysicsLoss().to(device)\n    scaler = GradScaler() if config.use_amp else None\n    writer = SummaryWriter(config.log_dir)\n\n    all_inputs, all_outputs = get_train_files(config.data_path)\n    valid_inputs = [all_inputs[i] for i in range(0,len(all_inputs), config.valid_frac)]\n    train_inputs = [f for f in all_inputs if f not in valid_inputs]\n    train_inputs = train_inputs[::config.train_frac] if config.train_frac>1 else train_inputs\n    train_outputs = inputs_files_to_output_files(train_inputs)\n    valid_outputs = inputs_files_to_output_files(valid_inputs)\n\n    train_loader = torch.utils.data.DataLoader(\n        SeismicDataset(train_inputs, train_outputs,500),\n        batch_size=config.batch_size, shuffle=True,\n        num_workers=4, pin_memory=True)\n    valid_loader = torch.utils.data.DataLoader(\n        SeismicDataset(valid_inputs, valid_outputs,500),\n        batch_size=config.batch_size, shuffle=False,\n        num_workers=2, pin_memory=True)\n\n    best_val = float('inf')\n    for epoch in range(config.num_epochs):\n        # Staging (selective training)\n        if epoch < config.warmup_epochs:\n            for n,p in model.named_parameters(): p.requires_grad = 'velocity_decoder' in n\n        else:\n            for p in model.parameters(): p.requires_grad = True\n        model.train()\n        train_loss = 0\n        train_metrics = {k:0 for k in ['recon','mae','edge','ssim','phys']}\n        for batch_idx,(seis,vp) in enumerate(train_loader):\n            seis, vp = seis.to(device), vp.unsqueeze(1).to(device)\n            optimizer.zero_grad()\n            with autocast(enabled=config.use_amp):\n                out = model(seis)\n                loss, mets = criterion(out['vp_pred'], vp, out['logR_pred'], out['logR_phys'])\n            if scaler:\n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(),2)\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(),2)\n                optimizer.step()\n            train_loss += loss.item()\n            for k in train_metrics: train_metrics[k] += mets[k]\n        # Average metrics\n        ntr = len(train_loader)\n        train_loss /= ntr\n        for k in train_metrics: train_metrics[k] /= ntr\n        scheduler.step()\n\n        # Validation\n        model.eval()\n        val_loss = 0\n        val_metrics = {k:0 for k in ['recon','mae','edge','ssim','phys']}\n        with torch.no_grad():\n            for seis,vp in valid_loader:\n                seis, vp = seis.to(device), vp.unsqueeze(1).to(device)\n                out = model(seis)\n                loss, mets = criterion(out['vp_pred'], vp, out['logR_pred'], out['logR_phys'])\n                val_loss += loss.item()\n                for k in val_metrics: val_metrics[k] += mets[k]\n        nv = len(valid_loader) or 1\n        val_loss /= nv\n        for k in val_metrics: val_metrics[k] /= nv\n\n        # Logging\n        writer.add_scalars('Loss', {'Train':train_loss,'Val':val_loss}, epoch)\n        for k in train_metrics:\n            writer.add_scalar(f\"Metrics/Train_{k}\", train_metrics[k], epoch)\n            writer.add_scalar(f\"Metrics/Val_{k}\", val_metrics[k], epoch)\n        writer.add_scalar('LR', optimizer.param_groups[0]['lr'], epoch)\n\n        # Save best checkpoint\n        if val_loss < best_val:\n            best_val = val_loss\n            torch.save({'model':model.state_dict(), 'opt':optimizer.state_dict()}, config.model_save_path)\n    writer.close()\n\nif __name__ == '__main__':\n    train()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}