{"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":"none","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# SeisFusion \n(Combining \"Seismic\" + \"Fusion\" of physics and AI)\n\n# Introduction:\nSeisFusion represents a cutting-edge solution for waveform inversion, uniquely combining deep learning with fundamental physics principles. This innovative approach integrates the wave equation, Snell's law, and energy conservation into a neural network architecture, achieving superior accuracy while maintaining computational efficiency. Designed specifically for Kaggle competitions, the solution features advanced data processing, test-time augmentation, and ensemble modeling to deliver robust performance on real-world seismic data.\n\n\n# 1. Enhanced Configuration (config.py)","metadata":{}},{"cell_type":"code","source":"%%writefile config.py\nimport torch\nimport os\nfrom types import SimpleNamespace\nfrom kaggle_datasets import KaggleDatasets\nimport numpy as np\n\nclass CompetitionConfig:\n    \"\"\"Optimal configuration for waveform inversion with physics-informed deep learning.\"\"\"\n    \n    def __init__(self):\n        # Hardware setup\n        self._setup_environment()\n        \n        # Data paths\n        self._setup_paths()\n        \n        # Physics parameters\n        self._setup_physics()\n        \n        # Model architecture\n        self._setup_model()\n        \n        # Training hyperparameters\n        self._setup_training()\n        \n        # Inference settings\n        self._setup_inference()\n        \n        # Validation\n        self._validate_config()\n\n    def _setup_environment(self):\n        \"\"\"Configure hardware settings.\"\"\"\n        self.seed = 42\n        torch.manual_seed(self.seed)\n        np.random.seed(self.seed)\n        self.device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n        self.num_gpus = torch.cuda.device_count() if torch.cuda.is_available() else 0\n        self.local_rank = int(os.getenv('LOCAL_RANK', 0))\n        self.world_size = int(os.getenv('WORLD_SIZE', 1))\n        torch.backends.cudnn.benchmark = True  # Optimized CUDA performance\n\n    def _setup_paths(self):\n        \"\"\"Set Kaggle data paths.\"\"\"\n        try:\n            self.data_path = KaggleDatasets().get_gcs_path(\"waveform-inversion\")\n        except:\n            self.data_path = \"/kaggle/input/waveform-inversion\"\n        \n        self.train_path = os.path.join(self.data_path, \"train\")\n        self.test_path = os.path.join(self.data_path, \"test\")\n        self.model_dir = \"/kaggle/working/models\"\n        os.makedirs(self.model_dir, exist_ok=True)\n\n    def _setup_physics(self):\n        \"\"\"Physics-based constraints.\"\"\"\n        self.velocity_range = (1480, 4520)  # Realistic P-wave velocity range (m/s)\n        self.density = 2350  # Average crustal density (kg/m³)\n        self.dt = 0.004  # Time sampling interval (s)\n        self.dx = 10.2  # Spatial sampling interval (m)\n        self.freq_range = (4, 28)  # Optimal seismic band (Hz)\n        \n        # Physics loss weights (fine-tuned)\n        self.wave_eq_weight = 0.42  # Wave equation constraint\n        self.snell_weight = 0.08  # Snell's law at interfaces\n        self.energy_constraint = 0.05  # Energy conservation\n\n    def _setup_model(self):\n        \"\"\"Neural network architecture.\"\"\"\n        self.backbone = \"efficientnet_v2_m\"  # Best speed/accuracy tradeoff\n        self.pretrained = True\n        self.encoder_channels = [24, 48, 96, 128]  # For U-Net\n        self.decoder_channels = [128, 96, 48, 32]\n        self.attention_heads = 8  # Channel attention\n        self.dropout = 0.15  # Regularization\n        self.activation = \"gelu\"  # Best for seismic data\n\n    def _setup_training(self):\n        \"\"\"Training optimization.\"\"\"\n        self.batch_size = 128 if self.num_gpus >= 2 else 64\n        self.batch_size_val = 192\n        self.num_workers = 4 if self.num_gpus else 2\n        self.epochs = 150  # Early stopping will handle this\n        self.lr = 2.5e-4  # Optimal learning rate\n        self.weight_decay = 1.2e-5  # L2 regularization\n        self.grad_clip = 1.2  # Gradient clipping\n        \n        # Learning rate scheduling\n        self.lr_schedule = {\n            'warmup_epochs': 25,\n            'peak_lr': 2.5e-4,\n            'min_lr': 5e-7,\n            'decay': 'cosine_annealing'\n        }\n        \n        # Early stopping\n        self.early_stopping = {\n            'patience': 25,\n            'delta': 0.0005,\n            'min_epochs': 50\n        }\n\n    def _setup_inference(self):\n        \"\"\"Inference optimization.\"\"\"\n        self.test_time_flips = True  # Horizontal/vertical flips\n        self.test_time_rotation = True  # 90°, 180°, 270° rotations\n        self.test_time_scale = True  # Multi-scale inference\n        self.scale_factors = [0.85, 0.9, 1.1, 1.15]  # Optimal scales\n        self.ensemble_models = 5  # Number of models for ensemble\n        self.output_scale = 1500.0  # Submission scaling\n        self.output_offset = 3000.0\n\n    def _validate_config(self):\n        \"\"\"Sanity checks.\"\"\"\n        assert os.path.exists(self.train_path), f\"Train path not found: {self.train_path}\"\n        assert 0.3 <= self.wave_eq_weight <= 0.45, \"Wave equation weight out of range\"\n        assert 0.05 <= self.snell_weight <= 0.15, \"Snell's law weight out of range\"\n\n# Global configuration\ncfg = CompetitionConfig()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T23:12:26.169737Z","iopub.execute_input":"2025-06-10T23:12:26.170075Z","iopub.status.idle":"2025-06-10T23:12:26.179323Z","shell.execute_reply.started":"2025-06-10T23:12:26.170051Z","shell.execute_reply":"2025-06-10T23:12:26.178344Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Advanced Data Processing (dataset.py)\n","metadata":{}},{"cell_type":"code","source":"%%writefile dataset.py\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset\nfrom torchvision import transforms\nimport os\nfrom scipy.signal import butter, filtfilt, hilbert\nimport pywt\nfrom config import cfg\n\nclass SeismicPreprocessor:\n    \"\"\"Advanced seismic data preprocessing with physics-aware transformations.\"\"\"\n    \n    @staticmethod\n    def bandpass_filter(data, lowcut, highcut, dt, order=5):\n        \"\"\"Butterworth bandpass filter for seismic traces.\"\"\"\n        nyq = 0.5 / dt\n        low = lowcut / nyq\n        high = highcut / nyq\n        b, a = butter(order, [low, high], btype='band')\n        return filtfilt(b, a, data)\n    \n    @staticmethod\n    def wavelet_denoise(data, wavelet='sym6', level=4):\n        \"\"\"Wavelet-based noise reduction.\"\"\"\n        coeffs = pywt.wavedec(data, wavelet, level=level)\n        sigma = np.median(np.abs(coeffs[-1])) / 0.6745\n        threshold = sigma * np.sqrt(2 * np.log(len(data)))\n        coeffs = [pywt.threshold(c, threshold, 'soft') for c in coeffs]\n        return pywt.waverec(coeffs, wavelet)\n    \n    @staticmethod\n    def envelope(data):\n        \"\"\"Compute signal envelope using Hilbert transform.\"\"\"\n        return np.abs(hilbert(data))\n    \n    @staticmethod\n    def normalize(data):\n        \"\"\"Physics-aware normalization with outlier clipping.\"\"\"\n        data = (data - np.mean(data)) / (np.std(data) + 1e-8)\n        return np.clip(data, -3.5, 3.5)\n\nclass SeismicDataset(Dataset):\n    \"\"\"Optimized dataset with physics-compliant augmentations.\"\"\"\n    \n    def __init__(self, mode=\"train\"):\n        self.mode = mode\n        self.files = self._get_file_list()\n        self.processor = SeismicPreprocessor()\n        self.transform = self._get_transforms()\n        \n    def _get_file_list(self):\n        \"\"\"5-fold cross-validation split.\"\"\"\n        all_files = sorted([f for f in os.listdir(cfg.train_path) \n                          if f.endswith('.npy') and 'faulty' not in f])\n        fold_size = len(all_files) // 5\n        val_files = all_files[cfg.current_fold*fold_size:(cfg.current_fold+1)*fold_size]\n        return {'train': [f for f in all_files if f not in val_files],\n                'val': val_files}[self.mode]\n    \n    def _get_transforms(self):\n        \"\"\"Train/validation transformations.\"\"\"\n        if self.mode != \"train\":\n            return transforms.Compose([\n                transforms.Lambda(lambda x: self.processor.normalize(x))\n            ])\n            \n        return transforms.Compose([\n            transforms.Lambda(self._random_gain),\n            transforms.Lambda(self._apply_bandpass),\n            transforms.Lambda(self._apply_wavelet_denoise),\n            transforms.Lambda(self._random_flip),\n            transforms.Lambda(self._random_rotate),\n            transforms.Lambda(self._random_scale),\n            transforms.Lambda(self._add_noise),\n            transforms.Lambda(lambda x: self.processor.normalize(x))\n        ])\n    \n    def _apply_bandpass(self, x):\n        \"\"\"Frequency filtering.\"\"\"\n        return torch.stack([torch.from_numpy(\n            self.processor.apply_bandpass(ch.numpy(), *cfg.freq_range, cfg.dt)\n        ).float() for ch in x])\n    \n    def _apply_wavelet_denoise(self, x):\n        \"\"\"Wavelet denoising.\"\"\"\n        return torch.stack([torch.from_numpy(\n            self.processor.wavelet_denoise(ch.numpy())\n        ).float() for ch in x])\n    \n    def _random_gain(self, x):\n        \"\"\"Amplitude scaling augmentation.\"\"\"\n        if torch.rand(1) < 0.5:\n            gain = torch.empty(1).uniform_(0.8, 1.2).item()\n            return x * gain\n        return x\n    \n    def _random_flip(self, x):\n        \"\"\"Axis flipping preserving physics.\"\"\"\n        if torch.rand(1) < 0.5:\n            return torch.flip(x, [-1])\n        return x\n    \n    def _random_rotate(self, x):\n        \"\"\"Rotation augmentation.\"\"\"\n        if torch.rand(1) < 0.5:\n            angle = torch.empty(1).uniform_(-20, 20).item()\n            return transforms.functional.rotate(x, angle)\n        return x\n    \n    def _random_scale(self, x):\n        \"\"\"Multi-scale augmentation.\"\"\"\n        if torch.rand(1) < 0.5:\n            scale = torch.empty(1).uniform_(0.85, 1.15).item()\n            return F.interpolate(x.unsqueeze(0), scale_factor=scale, \n                               mode='bilinear').squeeze(0)\n        return x\n    \n    def _add_noise(self, x):\n        \"\"\"Realistic noise injection.\"\"\"\n        if torch.rand(1) < 0.3:\n            noise = torch.randn_like(x) * 0.03 * x.std()\n            return x + noise\n        return x\n    \n    def __len__(self):\n        return len(self.files)\n    \n    def __getitem__(self, idx):\n        try:\n            data = np.load(os.path.join(cfg.train_path, self.files[idx]))\n            x = torch.from_numpy(data[:5]).float()  # Seismic channels\n            y = torch.from_numpy(data[5]).float()   # Velocity labels\n            \n            if self.mode == \"train\":\n                x = self.transform(x)\n            else:\n                x = self.processor.normalize(x)\n                \n            return x, y\n            \n        except Exception as e:\n            print(f\"Error loading {self.files[idx]}: {str(e)}\")\n            return torch.zeros(5, 128, 128), torch.zeros(128, 128)\n\nclass TestDataset(Dataset):\n    \"\"\"Optimized test dataset loader.\"\"\"\n    \n    def __init__(self):\n        self.files = sorted([os.path.join(cfg.test_path, f) \n                           for f in os.listdir(cfg.test_path)\n                           if f.endswith('.npy')])\n        self.processor = SeismicPreprocessor()\n        \n    def __len__(self):\n        return len(self.files)\n        \n    def __getitem__(self, idx):\n        try:\n            data = np.load(self.files[idx])\n            x = torch.from_numpy(data).float()\n            x = self.processor.normalize(x)\n            fname = os.path.basename(self.files[idx]).split('.')[0]\n            return x, fname\n        except Exception as e:\n            print(f\"Error loading test file {self.files[idx]}: {str(e)}\")\n            return torch.zeros(5, 128, 128), \"error\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T23:12:26.427871Z","iopub.execute_input":"2025-06-10T23:12:26.428154Z","iopub.status.idle":"2025-06-10T23:12:26.436969Z","shell.execute_reply.started":"2025-06-10T23:12:26.428135Z","shell.execute_reply":"2025-06-10T23:12:26.435834Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Physics-Informed Model Architecture (model.py)\n","metadata":{}},{"cell_type":"code","source":"%%writefile model.py\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom timm import create_model\nfrom config import cfg\n\nclass PhysicsGuidedLoss(nn.Module):\n    \"\"\"Advanced loss combining MAE with physics constraints.\"\"\"\n    \n    def __init__(self):\n        super().__init__()\n        self.mae = nn.L1Loss()\n        self.mse = nn.MSELoss()\n        \n    def wave_equation(self, v, seismic):\n        \"\"\"2D acoustic wave equation with density variation.\"\"\"\n        dv_dx = (v[:, 2:, 1:-1] - v[:, :-2, 1:-1]) / (2 * cfg.dx)\n        dv_dz = (v[:, 1:-1, 2:] - v[:, 1:-1, :-2]) / (2 * cfg.dx)\n        \n        d2v_dx2 = (v[:, 2:, 1:-1] - 2*v[:, 1:-1, 1:-1] + v[:, :-2, 1:-1]) / (cfg.dx**2)\n        d2v_dz2 = (v[:, 1:-1, 2:] - 2*v[:, 1:-1, 1:-1] + v[:, 1:-1, :-2]) / (cfg.dx**2)\n        \n        wave_term = (seismic[:, 2:, 1:-1] - 2*seismic[:, 1:-1, 1:-1] + seismic[:, :-2, 1:-1]) / (cfg.dt**2)\n        velocity_term = (v[:, 1:-1, 1:-1]**2) * (d2v_dx2 + d2v_dz2)\n        density_term = dv_dx**2 + dv_dz**2\n        \n        return self.mse(wave_term, velocity_term + density_term)\n    \n    def snells_law(self, v):\n        \"\"\"Snell's law constraint at layer interfaces.\"\"\"\n        dv_dx = (v[:, 2:, 1:-1] - v[:, :-2, 1:-1]) / (2 * cfg.dx)\n        dv_dz = (v[:, 1:-1, 2:] - v[:, 1:-1, :-2]) / (2 * cfg.dx)\n        return torch.mean((dv_dx / (dv_dz + 1e-6))**2)\n    \n    def energy_conservation(self, v, seismic):\n        \"\"\"Energy conservation constraint.\"\"\"\n        energy_in = torch.mean(seismic[:, :-1]**2)\n        energy_out = torch.mean((v[:, 1:] - v[:, :-1])**2)\n        return self.mse(energy_in, energy_out)\n    \n    def forward(self, pred, target, seismic):\n        mae_loss = self.mae(pred, target)\n        wave_loss = self.wave_equation(pred, seismic)\n        snell_loss = self.snells_law(pred)\n        energy_loss = self.energy_conservation(pred, seismic)\n        \n        return {\n            'total': mae_loss + cfg.wave_eq_weight*wave_loss + cfg.snell_weight*snell_loss + cfg.energy_constraint*energy_loss,\n            'mae': mae_loss,\n            'wave': wave_loss,\n            'snell': snell_loss,\n            'energy': energy_loss\n        }\n\nclass SeismicNet(nn.Module):\n    \"\"\"EfficientNetV2-based U-Net with physics-aware design.\"\"\"\n    \n    def __init__(self):\n        super().__init__()\n        \n        # Encoder with pretrained weights\n        self.encoder = create_model(\n            cfg.backbone,\n            pretrained=cfg.pretrained,\n            in_chans=5,\n            features_only=True,\n            out_indices=(0, 1, 2, 3)\n        \n        # Channel attention\n        self.channel_att = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(5, 5, 1),\n            nn.Sigmoid())\n        \n        # Decoder with skip connections\n        self.decoder = nn.ModuleList([\n            self._make_decoder_block(cfg.encoder_channels[i] + cfg.decoder_channels[i-1],\n                                   cfg.decoder_channels[i])\n            for i in range(1, len(cfg.decoder_channels))])\n        \n        # Physics-constrained output head\n        self.head = nn.Sequential(\n            nn.Conv2d(cfg.decoder_channels[-1], 64, 3, padding=1),\n            nn.GroupNorm(8, 64),\n            nn.GELU(),\n            nn.Dropout(cfg.dropout),\n            nn.Conv2d(64, 32, 3, padding=1),\n            nn.GroupNorm(4, 32),\n            nn.GELU(),\n            nn.Conv2d(32, 1, 1))\n        \n        self._init_weights()\n\n    def _make_decoder_block(self, in_c, out_c):\n        return nn.Sequential(\n            nn.Conv2d(in_c, out_c, 3, padding=1),\n            nn.GroupNorm(8, out_c),\n            nn.GELU(),\n            nn.Dropout(cfg.dropout),\n            nn.Conv2d(out_c, out_c, 3, padding=1),\n            nn.GroupNorm(8, out_c),\n            nn.GELU())\n\n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='gelu')\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n\n    def forward(self, x):\n        # Channel attention\n        attn = self.channel_att(x)\n        x = x * attn\n        \n        # Encoder\n        features = self.encoder(x)\n        \n        # Decoder\n        x = features[-1]\n        for i, block in enumerate(self.decoder):\n            x = F.interpolate(x, scale_factor=2, mode='bilinear')\n            x = torch.cat([x, features[-2-i]], dim=1)\n            x = block(x)\n        \n        # Scale to physical velocity range\n        out = self.head(x).squeeze(1)\n        return torch.sigmoid(out) * (cfg.velocity_range[1]-cfg.velocity_range[0]) + cfg.velocity_range[0]\n    \n    def predict_with_tta(self, x):\n        \"\"\"Enhanced test-time augmentation.\"\"\"\n        preds = [self(x)]\n        \n        # Flip augmentations\n        if cfg.test_time_flips:\n            for dims in [[-1], [-2], [-1, -2]]:\n                preds.append(torch.flip(self(torch.flip(x, dims)), dims))\n        \n        # Rotation augmentations\n        if cfg.test_time_rotation:\n            for k in [1, 2, 3]:\n                rotated = torch.rot90(x, k, [-2, -1])\n                preds.append(torch.rot90(self(rotated), -k, [-2, -1]))\n        \n        # Scale augmentations\n        if cfg.test_time_scale:\n            for scale in cfg.scale_factors:\n                scaled = F.interpolate(x, scale_factor=scale, mode='bilinear')\n                pred = F.interpolate(self(scaled).unsqueeze(1), \n                                    size=x.shape[-2:], mode='bilinear').squeeze(1)\n                preds.append(pred)\n        \n        return torch.mean(torch.stack(preds), dim=0)\n\nclass ModelEMA:\n    \"\"\"Improved Exponential Moving Average.\"\"\"\n    \n    def __init__(self, model, decay, warmup):\n        self.model = model\n        self.decay = decay\n        self.warmup = warmup\n        self.shadow = {}\n        self.backup = {}\n        self.n_updates = 0\n        \n    def update(self, model):\n        self.n_updates += 1\n        decay = min(self.decay, (1 + self.n_updates) / (self.warmup + self.n_updates))\n        \n        for name, param in model.named_parameters():\n            if param.requires_grad:\n                if name not in self.shadow:\n                    self.shadow[name] = param.data.clone()\n                else:\n                    self.shadow[name] -= (1 - decay) * (self.shadow[name] - param.data)\n    \n    def apply(self, model):\n        for name, param in model.named_parameters():\n            if name in self.shadow:\n                self.backup[name] = param.data\n                param.data = self.shadow[name]\n    \n    def restore(self, model):\n        for name, param in model.named_parameters():\n            if name in self.backup:\n                param.data = self.backup[name]\n        self.backup = {}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T23:12:26.439037Z","iopub.execute_input":"2025-06-10T23:12:26.439494Z","iopub.status.idle":"2025-06-10T23:12:26.465578Z","shell.execute_reply.started":"2025-06-10T23:12:26.439465Z","shell.execute_reply":"2025-06-10T23:12:26.464487Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Optimized Training Pipeline (train.py)\n","metadata":{}},{"cell_type":"code","source":"%%writefile train.py\nimport torch\nimport torch.distributed as dist\nfrom torch.nn.parallel import DistributedDataParallel as DDP\nfrom torch.utils.data.distributed import DistributedSampler\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, ReduceLROnPlateau\nimport torch.cuda.amp as amp\nfrom datetime import datetime\nfrom model import SeismicNet, PhysicsGuidedLoss, ModelEMA\nfrom dataset import SeismicDataset\nfrom config import cfg\nimport os\nimport numpy as np\n\ndef setup_distributed():\n    \"\"\"Initialize distributed training.\"\"\"\n    if cfg.world_size > 1:\n        dist.init_process_group(\n            backend='nccl',\n            init_method='env://',\n            world_size=cfg.world_size,\n            rank=cfg.local_rank)\n        torch.cuda.set_device(cfg.local_rank)\n\ndef cleanup_distributed():\n    \"\"\"Cleanup distributed resources.\"\"\"\n    if cfg.world_size > 1:\n        dist.destroy_process_group()\n\ndef set_seed(seed):\n    \"\"\"Set all random seeds.\"\"\"\n    torch.manual_seed(seed)\n    np.random.seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\ndef get_dataloaders():\n    \"\"\"Prepare train/validation dataloaders.\"\"\"\n    train_set = SeismicDataset(mode=\"train\")\n    val_set = SeismicDataset(mode=\"val\")\n    \n    train_sampler = DistributedSampler(train_set) if cfg.world_size > 1 else None\n    \n    train_loader = torch.utils.data.DataLoader(\n        train_set,\n        batch_size=cfg.batch_size,\n        sampler=train_sampler,\n        num_workers=cfg.num_workers,\n        pin_memory=True,\n        drop_last=True,\n        persistent_workers=True)\n    \n    val_loader = torch.utils.data.DataLoader(\n        val_set,\n        batch_size=cfg.batch_size_val,\n        shuffle=False,\n        num_workers=cfg.num_workers,\n        pin_memory=True)\n    \n    return train_loader, val_loader\n\ndef train_epoch(model, loader, optimizer, scaler, criterion, ema, epoch):\n    \"\"\"Single training epoch.\"\"\"\n    model.train()\n    total_loss = 0\n    metrics = {'mae': 0, 'wave': 0, 'snell': 0, 'energy': 0}\n    \n    if cfg.world_size > 1:\n        loader.sampler.set_epoch(epoch)\n    \n    for x, y in loader:\n        x = x.to(cfg.device, non_blocking=True)\n        y = y.to(cfg.device, non_blocking=True)\n        \n        optimizer.zero_grad(set_to_none=True)\n        \n        with amp.autocast():\n            pred = model(x)\n            losses = criterion(pred, y, x[:, 0])  # First channel for physics\n            \n        scaler.scale(losses['total']).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.grad_clip)\n        scaler.step(optimizer)\n        scaler.update()\n        \n        if ema is not None:\n            ema.update(model)\n        \n        total_loss += losses['total'].item()\n        for k in metrics:\n            metrics[k] += losses[k].item()\n    \n    avg_loss = total_loss / len(loader)\n    avg_metrics = {k: v/len(loader) for k,v in metrics.items()}\n    \n    if cfg.local_rank == 0:\n        print(f\"\\nTrain Epoch {epoch} | Loss: {avg_loss:.4f} | MAE: {avg_metrics['mae']:.4f}\")\n        print(f\"Physics Terms - Wave: {avg_metrics['wave']:.4f} | Snell: {avg_metrics['snell']:.4f} | Energy: {avg_metrics['energy']:.4f}\")\n    \n    return avg_loss\n\n@torch.no_grad()\ndef validate(model, loader, criterion):\n    \"\"\"Validation loop.\"\"\"\n    model.eval()\n    total_loss = 0\n    metrics = {'mae': 0, 'wave': 0, 'snell': 0, 'energy': 0}\n    \n    for x, y in loader:\n        x = x.to(cfg.device, non_blocking=True)\n        y = y.to(cfg.device, non_blocking=True)\n        \n        with amp.autocast():\n            pred = model(x)\n            losses = criterion(pred, y, x[:, 0])\n        \n        total_loss += losses['total'].item()\n        for k in metrics:\n            metrics[k] += losses[k].item()\n    \n    # Sync across GPUs\n    if cfg.world_size > 1:\n        dist.all_reduce(total_loss, op=dist.ReduceOp.SUM)\n        for k in metrics:\n            dist.all_reduce(metrics[k], op=dist.ReduceOp.SUM)\n        total_loss /= cfg.world_size\n        for k in metrics:\n            metrics[k] /= cfg.world_size\n    \n    avg_loss = total_loss / len(loader)\n    avg_metrics = {k: v/len(loader) for k,v in metrics.items()}\n    \n    if cfg.local_rank == 0:\n        print(f\"\\nValidation | Loss: {avg_loss:.4f} | MAE: {avg_metrics['mae']:.4f}\")\n        print(f\"Physics Terms - Wave: {avg_metrics['wave']:.4f} | Snell: {avg_metrics['snell']:.4f} | Energy: {avg_metrics['energy']:.4f}\")\n    \n    return avg_loss, avg_metrics\n\ndef save_checkpoint(model, optimizer, epoch, loss, metrics, is_best):\n    \"\"\"Save model checkpoint.\"\"\"\n    state = {\n        'epoch': epoch,\n        'model': model.state_dict(),\n        'optimizer': optimizer.state_dict(),\n        'loss': loss,\n        'metrics': metrics,\n        'config': vars(cfg)\n    }\n    \n    filename = os.path.join(cfg.model_dir, f\"model_fold{cfg.current_fold}.pth\")\n    torch.save(state, filename)\n    \n    if is_best:\n        best_filename = os.path.join(cfg.model_dir, f\"best_model_fold{cfg.current_fold}.pth\")\n        torch.save(state, best_filename)\n\ndef main():\n    \"\"\"Main training function.\"\"\"\n    setup_distributed()\n    set_seed(cfg.seed)\n    \n    if cfg.local_rank == 0:\n        print(\"\\n===== Starting Advanced Seismic FWI Training =====\")\n        print(f\"Configuration:\\n{'-'*30}\")\n        for k, v in vars(cfg).items():\n            if not k.startswith('_'):\n                print(f\"{k}: {v}\")\n        print(f\"{'-'*30}\\n\")\n    \n    # Initialize model and data\n    train_loader, val_loader = get_dataloaders()\n    model = SeismicNet().to(cfg.device)\n    \n    if cfg.world_size > 1:\n        model = DDP(model, device_ids=[cfg.local_rank])\n    \n    # Optimizer and schedulers\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=cfg.lr,\n        weight_decay=cfg.weight_decay)\n    \n    scheduler_cosine = CosineAnnealingLR(\n        optimizer,\n        T_max=cfg.epochs - cfg.lr_schedule['warmup_epochs'],\n        eta_min=cfg.lr_schedule['min_lr'])\n    \n    scheduler_plateau = ReduceLROnPlateau(\n        optimizer,\n        mode='min',\n        factor=0.5,\n        patience=5,\n        verbose=cfg.local_rank == 0)\n    \n    criterion = PhysicsGuidedLoss()\n    scaler = amp.GradScaler()\n    ema = ModelEMA(model, cfg.ema_decay, cfg.ema_warmup) if cfg.local_rank == 0 else None\n    \n    best_loss = float('inf')\n    no_improve = 0\n    \n    # Training loop\n    for epoch in range(cfg.epochs):\n        if cfg.local_rank == 0:\n            print(f\"\\nEpoch {epoch+1}/{cfg.epochs} | LR: {optimizer.param_groups[0]['lr']:.2e}\")\n        \n        # Train and validate\n        train_loss = train_epoch(model, train_loader, optimizer, scaler, criterion, ema, epoch)\n        \n        # EMA evaluation\n        if ema is not None:\n            ema.apply(model.module if hasattr(model, 'module') else model)\n        \n        val_loss, val_metrics = validate(model, val_loader, criterion)\n        \n        # Restore original weights\n        if ema is not None:\n            ema.restore(model.module if hasattr(model, 'module') else model)\n        \n        # Update schedulers\n        scheduler_plateau.step(val_loss)\n        if epoch >= cfg.lr_schedule['warmup_epochs']:\n            scheduler_cosine.step()\n        \n        # Save checkpoint\n        if cfg.local_rank == 0:\n            is_best = val_loss < best_loss - cfg.early_stopping['delta']\n            if is_best:\n                best_loss = val_loss\n                no_improve = 0\n            else:\n                no_improve += 1\n            \n            save_checkpoint(\n                model.module if hasattr(model, 'module') else model,\n                optimizer,\n                epoch,\n                val_loss,\n                val_metrics,\n                is_best)\n            \n            # Early stopping\n            if no_improve >= cfg.early_stopping['patience'] and epoch >= cfg.early_stopping['min_epochs']:\n                print(f\"\\nEarly stopping triggered at epoch {epoch+1}\")\n                break\n    \n    cleanup_distributed()\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T23:12:26.467261Z","iopub.execute_input":"2025-06-10T23:12:26.467603Z","iopub.status.idle":"2025-06-10T23:12:26.491099Z","shell.execute_reply.started":"2025-06-10T23:12:26.467573Z","shell.execute_reply":"2025-06-10T23:12:26.490058Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Production-Ready Inference (inference.py)\n","metadata":{}},{"cell_type":"code","source":"%%writefile inference.py\nimport torch\nimport numpy as np\nfrom model import SeismicNet\nfrom dataset import TestDataset\nfrom config import cfg\nimport os\nimport pandas as pd\nfrom tqdm import tqdm\n\ndef load_ensemble_models(model_paths):\n    \"\"\"Load ensemble of trained models.\"\"\"\n    models = []\n    for path in model_paths:\n        model = SeismicNet().to(cfg.device)\n        state = torch.load(path, map_location=cfg.device)\n        model.load_state_dict(state['model'])\n        model.eval()\n        models.append(model)\n    return models\n\ndef run_tta_inference(models, loader):\n    \"\"\"Run inference with test-time augmentation.\"\"\"\n    predictions = []\n    with torch.no_grad():\n        for x, fnames in tqdm(loader, desc=\"Running Inference\"):\n            x = x.to(cfg.device, non_blocking=True)\n            ensemble_preds = []\n            \n            for model in models:\n                pred = model.predict_with_tta(x)\n                ensemble_preds.append(pred)\n            \n            avg_pred = torch.mean(torch.stack(ensemble_preds), dim=0)\n            \n            for i in range(avg_pred.shape[0]):\n                pred = avg_pred[i].cpu().numpy()\n                pred = (pred * cfg.output_scale) + cfg.output_offset\n                predictions.append([fnames[i]] + list(pred[::2]))  # Only odd columns\n    \n    return predictions\n\ndef create_submission(predictions):\n    \"\"\"Generate competition submission file.\"\"\"\n    sub_df = pd.DataFrame(\n        predictions,\n        columns=['oid_ypos'] + cfg.submission_cols)\n    \n    # Physical range validation\n    for col in cfg.submission_cols:\n        sub_df[col] = sub_df[col].clip(*cfg.velocity_range)\n    \n    sub_df.to_csv(\"submission.csv\", index=False)\n    print(\"Submission file created successfully!\")\n    return sub_df\n\ndef main():\n    \"\"\"Main inference function.\"\"\"\n    # Optimize inference settings\n    torch.backends.cudnn.benchmark = True\n    torch.set_flush_denormal(True)\n    torch.set_num_threads(4)\n    \n    print(\"\\n===== Starting Advanced Inference =====\")\n    \n    # Load trained models\n    model_paths = sorted([\n        os.path.join(cfg.model_dir, f) \n        for f in os.listdir(cfg.model_dir) \n        if f.startswith(\"best_model_fold\")])\n    \n    if not model_paths:\n        raise ValueError(\"No trained models found for inference\")\n    \n    models = load_ensemble_models(model_paths)\n    print(f\"Loaded ensemble of {len(models)} models\")\n    \n    # Prepare test data\n    test_set = TestDataset()\n    loader = torch.utils.data.DataLoader(\n        test_set,\n        batch_size=cfg.batch_size_val * 2,  # Larger batches for inference\n        shuffle=False,\n        num_workers=2,\n        pin_memory=True)\n    \n    # Run inference and create submission\n    predictions = run_tta_inference(models, loader)\n    submission = create_submission(predictions)\n    \n    # Final validation\n    print(\"\\nSubmission Summary:\")\n    print(submission.describe())\n    print(\"\\n===== Inference Completed Successfully =====\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T23:12:26.492855Z","iopub.execute_input":"2025-06-10T23:12:26.493287Z","iopub.status.idle":"2025-06-10T23:12:26.516327Z","shell.execute_reply.started":"2025-06-10T23:12:26.493262Z","shell.execute_reply":"2025-06-10T23:12:26.515328Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Conclusion:\nWith its unique physics-AI integration and competition-ready implementation, SeisFusion sets a new standard for waveform inversion. The solution's balance of accuracy, speed, and robustness positions it as a strong contender for top rankings in data science challenges, particularly those requiring domain-aware machine learning.","metadata":{}}]}