{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"},{"sourceId":11568812,"sourceType":"datasetVersion","datasetId":7253205},{"sourceId":11569667,"sourceType":"datasetVersion","datasetId":7253605},{"sourceId":11569755,"sourceType":"datasetVersion","datasetId":7253661}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Improved UNet Model with TV Regularization and DDP\n\nThis notebook builds upon the work of Egor Trushin's notebook:  \n[GWI UNet with float16 Dataset](https://www.kaggle.com/code/egortrushin/gwi-unet-with-float16-dataset)  \nwith several key improvements:\n\n## Key Enhancements\n\n1. **Total Variation (TV) Regularization Loss**  \n   - Added TV regularization with weight = `0.01` (`tv_weight=0.01`)  \n   - Significantly improves validation scores on test data  \n   - Particularly effective for enhancing sharp features in velocity models (e.g., faults)  \n   - Suggests test data may contain more sharp features than openFWI data\n\n2. **Training Configuration**  \n   - Uses batch size = `256` (`batch_size=256`)  \n   - Trained for `max_epochs=10`  and optimizer learning rate `lr=0.0005`\n   - Utilizes both T4 GPUs available in Kaggle via **Distributed Data Parallel (DDP)**  \n   - Debugging friendly `_train.py`\n   - Maintains float16 dataset compatibility for efficient training\n   ","metadata":{}},{"cell_type":"code","source":"%%writefile config.yaml\n\ndata_path: /kaggle/input/waveform-inversion\ntest_path: /kaggle/input/open-wfi-test/test\nmodel: \n    name: UNet\n    unet_params:\n        init_features: 32\n        depth: 5\ntv_weight: 0.01  \nread_weights: null\nbatch_size: 256\nprint_freq: 100\nmax_epochs: 10\nes_epochs: 3\nseed: 42\nvalid_frac: 16\ntrain_frac: 2\noptimizer:\n    lr: 0.0005\n    weight_decay: 0.001\nscheduler:\n    params:\n        factor: 0.316227766\n        patience: 1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T15:27:34.134718Z","iopub.execute_input":"2025-05-22T15:27:34.134965Z","iopub.status.idle":"2025-05-22T15:27:34.142834Z","shell.execute_reply.started":"2025-05-22T15:27:34.134945Z","shell.execute_reply":"2025-05-22T15:27:34.142035Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile _dataset.py\n\n# Data\n\nimport os\nimport numpy as np\nfrom pathlib import Path\nfrom torch.utils.data import Dataset, DataLoader\nimport torch\nimport pickle\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\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\n    all_outputs = inputs_files_to_output_files(all_inputs)\n\n    assert all(f.exists() for f in all_outputs)\n\n    return all_inputs, all_outputs\n\n\n# Data loader\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        # Calculate file offset and sample offset within file\n        file_idx = idx // self.n_examples_per_file\n        sample_idx = idx % self.n_examples_per_file\n\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        try:\n            return X[sample_idx].copy(), y[sample_idx].copy()\n        finally:\n            del X, y\n\n        \nclass TestDataset(Dataset):\n    def __init__(self, test_files):\n        self.test_files = test_files\n\n\n    def __len__(self):\n        return len(self.test_files)\n\n\n    def __getitem__(self, i):\n        test_file = self.test_files[i]\n        \n        try:\n            data = np.load(test_file, mmap_mode='r')  # Use memory mapping\n            return torch.from_numpy(data.copy()).float(), test_file.stem\n        finally:\n            del data  # Explicit cleanup","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T15:27:36.930033Z","iopub.execute_input":"2025-05-22T15:27:36.930635Z","iopub.status.idle":"2025-05-22T15:27:36.936562Z","shell.execute_reply.started":"2025-05-22T15:27:36.930607Z","shell.execute_reply":"2025-05-22T15:27:36.935909Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile _model.py\n# Model with TV Regularization\n\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch\n\nclass TVRegularizer(nn.Module):\n    \"\"\"Differentiable TV regularization layer\"\"\"\n    def __init__(self):\n        super().__init__()\n        \n    def forward(self, x):\n        batch_size = x.size(0)\n        # Calculate horizontal and vertical differences\n        h_diff = x[:,:,1:,:] - x[:,:,:-1,:]\n        v_diff = x[:,:,:,1:] - x[:,:,:,:-1]\n        \n        # Calculate anisotropic TV (sum of absolute differences)\n        h_tv = torch.abs(h_diff).sum()\n        v_tv = torch.abs(v_diff).sum()\n        \n        return (h_tv + v_tv) / batch_size\n\nclass ResidualDoubleConv(nn.Module):\n    \"\"\"(Convolution => [BN] => ReLU) * 2 + Residual Connection\"\"\"\n    def __init__(self, in_channels, out_channels, mid_channels=None):\n        super().__init__()\n        if not mid_channels:\n            mid_channels = out_channels\n\n        # First convolution layer\n        self.conv1 = nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(mid_channels)\n        self.relu = nn.ReLU(inplace=True)\n\n        # Second convolution layer\n        self.conv2 = nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n\n        # Shortcut connection to handle potential channel mismatch\n        if in_channels == out_channels:\n            self.shortcut = nn.Identity()\n        else:\n            # Projection shortcut: 1x1 conv + BN to match output channels\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n\n    def forward(self, x):\n        identity = x  # Store the input for the residual connection\n\n        # First conv block\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        # Second conv block (without final ReLU yet)\n        out = self.conv2(out)\n        out = self.bn2(out)\n\n        # Apply shortcut to the identity path\n        identity_mapped = self.shortcut(identity)\n\n        # Add the residual connection\n        out += identity_mapped\n\n        # Apply final ReLU\n        out = self.relu(out)\n        return out\n\n\nclass Up(nn.Module):\n    \"\"\"Upscaling then ResidualDoubleConv\"\"\"\n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super().__init__()\n        self.bilinear = bilinear\n\n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode=\"bilinear\", align_corners=False)\n            # Input to ResidualDoubleConv = channels from upsampled layer below + channels from skip connection\n            # Output of ResidualDoubleConv = desired output channels for this decoder stage\n            self.conv = ResidualDoubleConv(in_channels + out_channels, out_channels) # Use ResidualDoubleConv\n\n        else: # Using ConvTranspose2d\n            # ConvTranspose halves the channels: in_channels -> in_channels // 2\n            self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)\n            # Input channels to ResidualDoubleConv\n            conv_in_channels = in_channels // 2 # Channels after ConvTranspose\n            skip_channels = out_channels       # Channels from skip connection\n            total_in_channels = conv_in_channels + skip_channels\n            self.conv = ResidualDoubleConv(total_in_channels, out_channels) # Use ResidualDoubleConv\n\n    def forward(self, x1, x2):\n        # x1 is the feature map from the layer below (needs upsampling)\n        # x2 is the skip connection from the corresponding encoder layer\n        x1 = self.up(x1)\n\n        # Pad x1 if its dimensions don't match x2 after upsampling\n        # Input is CHW\n        diffY = x2.size(2) - x1.size(2)\n        diffX = x2.size(3) - x1.size(3)\n\n        # Pad format: (padding_left, padding_right, padding_top, padding_bottom)\n        x1 = F.pad(\n            x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]\n        )\n\n        # Concatenate along the channel dimension\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\nclass OutConv(nn.Module):\n    \"\"\"1x1 Convolution for the output layer\"\"\"\n    # [Previous implementation remains exactly the same]\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)\n\n    def forward(self, x):\n        return self.conv(x)\n\nclass UNet(nn.Module):\n    \"\"\"U-Net with TV Regularization\"\"\"\n\n    def __init__(\n        self,\n        n_channels=5,\n        n_classes=1,\n        init_features=32,\n        depth=5, # number of pooling layers\n        bilinear=True,\n    ):\n        super().__init__()\n        self.n_channels = n_channels\n        self.n_classes = n_classes\n        self.bilinear = bilinear\n        self.depth = depth\n        self.tv_reg = TVRegularizer()  # Initialize TV regularizer\n        \n        # [Rest of the initialization remains the same]\n        self.initial_pool = nn.AvgPool2d(kernel_size=(14, 1), stride=(14, 1))\n\n        # --- Encoder ---\n        self.encoder_convs = nn.ModuleList() # Store conv blocks\n        self.encoder_pools = nn.ModuleList() # Store pool layers\n\n        # Initial conv block (no pooling before it)\n        # Use ResidualDoubleConv for the initial convolution block\n        self.inc = ResidualDoubleConv(n_channels, init_features)\n        self.encoder_convs.append(self.inc)\n\n        current_features = init_features\n        for _ in range(depth):\n            # Define convolution block for this stage\n            conv = ResidualDoubleConv(current_features, current_features * 2)\n            # Define pooling layer for this stage\n            pool = nn.MaxPool2d(2)\n            self.encoder_convs.append(conv)\n            self.encoder_pools.append(pool)\n            current_features *= 2\n\n        # --- Bottleneck ---\n        # Use ResidualDoubleConv for the bottleneck\n        self.bottleneck = ResidualDoubleConv(current_features, current_features)\n\n        # --- Decoder ---\n        self.decoder_blocks = nn.ModuleList()\n        # Input features start from bottleneck output features\n        # Output features at each stage are halved\n        for _ in range(depth):\n            # Up block uses ResidualDoubleConv internally and handles channels\n            up_block = Up(current_features, current_features // 2, bilinear)\n            self.decoder_blocks.append(up_block)\n            current_features //= 2 # Halve features for next Up block input\n\n        # --- Output Layer ---\n        # Input features are the output features of the last Up block\n        self.outc = OutConv(current_features, n_classes)\n\n    def _pad_or_crop(self, x, target_h=70, target_w=70):\n        \"\"\"Pads or crops input tensor x to target height and width.\"\"\"\n        # [Previous implementation remains exactly the same]\n        _, _, h, w = x.shape\n        # Pad Height if needed\n        if h < target_h:\n            pad_top = (target_h - h) // 2\n            pad_bottom = target_h - h - pad_top\n            x = F.pad(x, (0, 0, pad_top, pad_bottom))  # Pad height only\n            h = target_h\n        # Pad Width if needed\n        if w < target_w:\n            pad_left = (target_w - w) // 2\n            pad_right = target_w - w - pad_left\n            x = F.pad(x, (pad_left, pad_right, 0, 0))  # Pad width only\n            w = target_w\n        # Crop Height if needed\n        if h > target_h:\n            crop_top = (h - target_h) // 2\n            # Use slicing to crop\n            x = x[:, :, crop_top : crop_top + target_h, :]\n            h = target_h\n        # Crop Width if needed\n        if w > target_w:\n            crop_left = (w - target_w) // 2\n            x = x[:, :, :, crop_left : crop_left + target_w]\n            w = target_w\n        return x\n\n    def forward(self, x):\n        # [Forward pass remains the same until the end]\n        \n        # Initial pooling and resizing\n        x_pooled = self.initial_pool(x)\n        x_resized = self._pad_or_crop(x_pooled, target_h=70, target_w=70)\n\n        # --- Encoder Path ---\n        skip_connections = []\n        xi = x_resized\n\n        # Apply initial conv (inc)\n        xi = self.encoder_convs[0](xi)\n        skip_connections.append(xi) # Store output of inc\n\n        # Apply subsequent encoder convs and pools\n        # self.depth is the number of pooling layers\n        for i in range(self.depth):\n            # Apply conv block for this stage\n            xi = self.encoder_convs[i+1](xi)\n            # Store skip connection *before* pooling\n            skip_connections.append(xi)\n            # Apply pooling layer for this stage\n            xi = self.encoder_pools[i](xi)\n\n        # Apply bottleneck conv\n        xi = self.bottleneck(xi)\n\n        # --- Decoder Path ---\n        xu = xi # Start with bottleneck output\n        # Iterate through decoder blocks and corresponding skip connections in reverse\n        for i, block in enumerate(self.decoder_blocks):\n            # Determine the correct skip connection index from the end\n            # Example: depth=5. Skips stored: [inc, enc1, enc2, enc3, enc4] (indices 0-4)\n            # Decoder 0 (Up(1024, 512)) needs skip 4 (enc4)\n            # Decoder 1 (Up(512, 256)) needs skip 3 (enc3) ...\n            # Decoder 4 (Up(64, 32)) needs skip 0 (inc)\n            skip_index = self.depth - 1 - i\n            skip = skip_connections[skip_index]\n            xu = block(xu, skip) # Up block combines xu (from below) and skip\n\n        # --- Final Output ---\n        logits = self.outc(xu)\n        # Apply scaling and offset specific to the problem's target range\n        output = logits * 1000.0 + 1500.0\n\n        # Calculate TV regularization\n        tv_loss = self.tv_reg(logits)\n        \n        return output, tv_loss  # Now returns both output and TV loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T15:27:45.025682Z","iopub.execute_input":"2025-05-22T15:27:45.026100Z","iopub.status.idle":"2025-05-22T15:27:45.034400Z","shell.execute_reply.started":"2025-05-22T15:27:45.026077Z","shell.execute_reply":"2025-05-22T15:27:45.033635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile _utils.py\n# Utils\n\nimport datetime\nimport random\nimport torch\nimport numpy as np\n\ndef format_time(elapsed):\n    \"\"\"Take a time in seconds and return a string hh:mm:ss.\"\"\"\n    elapsed_rounded = int(round((elapsed)))\n    return str(datetime.timedelta(seconds=elapsed_rounded))\n\ndef seed_everything(\n    seed_value: int\n) -> None:\n    \"\"\"\n    Controlling a unified seed value for Python, NumPy, and PyTorch (CPU, GPU).\n\n    Parameters:\n    ----------\n    seed_value : int\n        The unified random seed value.\n    \"\"\"\n    random.seed(seed_value) # Python\n    np.random.seed(seed_value) # cpu vars\n    torch.manual_seed(seed_value) # cpu  vars    \n    if torch.cuda.is_available(): \n        torch.cuda.manual_seed(seed_value)\n        torch.cuda.manual_seed_all(seed_value) # gpu vars\n    if torch.backends.cudnn.is_available:\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T15:27:53.660156Z","iopub.execute_input":"2025-05-22T15:27:53.660434Z","iopub.status.idle":"2025-05-22T15:27:53.665152Z","shell.execute_reply.started":"2025-05-22T15:27:53.660414Z","shell.execute_reply":"2025-05-22T15:27:53.664551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile _train.py\n\"\"\"\nDistributed Training Script with TV Regularization\n\"\"\"\n\"\"\"\nDistributed Training Script for Seismic Data using PyTorch DDP\n\nFeatures:\n- Distributed Data Parallel (DDP) training across multiple GPUs\n- Automatic Mixed Precision (AMP) training\n- Learning rate scheduling with plateau detection\n- Early stopping based on validation loss\n- Comprehensive logging and model checkpointing\n\"\"\"\n# [Previous imports remain the same]\nimport os\nimport sys\nimport yaml\nimport time\nimport torch\nimport numpy as np\nimport torch.distributed as dist\nimport torch.multiprocessing as mp\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, DistributedSampler\nfrom torch.nn.parallel import DistributedDataParallel as DDP\nfrom pprint import pprint\n\n# Local imports\nfrom _dataset import inputs_files_to_output_files, get_train_files, SeismicDataset \nfrom _model import UNet\nfrom _utils import format_time, seed_everything\n\n# Constants\nDEFAULT_PORT = 29500\nMASTER_ADDR = '127.0.0.1'\n\ndef load_config():\n    \"\"\"Load and validate training configuration.\"\"\"\n    with open(\"config.yaml\", \"r\") as f:\n        config = yaml.safe_load(f)\n    \n    if config[\"data_path\"] is None:\n        config[\"data_path\"] = os.environ[\"TMPDIR\"]\n    \n    print(\"\\nConfiguration:\")\n    pprint(config)\n    return config\n\ndef prepare_datasets(config):\n    \"\"\"Prepare training and validation datasets.\n    \n    Returns:\n        tuple: (train_dataset, valid_dataset)\n    \"\"\"\n    print(\"\\nPreparing datasets ...\")\n    all_inputs, all_outputs = [], []\n    for data_dir in [\"/kaggle/input/open-wfi-1/openfwi_float16_1\", \n                    \"/kaggle/input/open-wfi-2/openfwi_float16_2\"]:\n        inputs, outputs = get_train_files(data_dir)\n        all_inputs.extend(inputs)\n        all_outputs.extend(outputs)\n    \n    # Split into train/validation\n    valid_idx = range(0, len(all_inputs), config[\"valid_frac\"])\n    valid_inputs = [all_inputs[i] for i in valid_idx]\n    train_inputs = [f for f in all_inputs if f not in valid_inputs]\n    \n    if config[\"train_frac\"] > 1:\n        train_inputs = [train_inputs[i] for i in range(0, len(train_inputs), config[\"train_frac\"])]\n    \n    # Get corresponding output files\n    train_outputs = inputs_files_to_output_files(train_inputs)\n    valid_outputs = inputs_files_to_output_files(valid_inputs)\n\n    print(f\"Total files: {len(all_inputs)}\")\n    print(f\"Training files: {len(train_inputs)}\")\n    print(f\"Validation files: {len(valid_inputs)}\")\n    \n    return (train_inputs, train_outputs,\n            valid_inputs, valid_outputs)\n\ndef setup_distributed(rank, world_size):\n    \"\"\"Initialize distributed training environment.\"\"\"\n    os.environ['MASTER_ADDR'] = MASTER_ADDR\n    os.environ['MASTER_PORT'] = str(DEFAULT_PORT)\n    \n    # NCCL optimization for Kaggle\n    os.environ['NCCL_SOCKET_IFNAME'] = 'lo'\n    os.environ['NCCL_NSOCKS_PERTHREAD'] = '4'\n    os.environ['NCCL_SOCKET_NTHREADS'] = '2'\n    os.environ['NCCL_DEBUG'] = 'INFO'\n    # Add these for more stable runs\n    os.environ['NCCL_SOCKET_TIMEOUT'] = '600'  # 10 minute timeout\n    os.environ['NCCL_IB_DISABLE'] = '1' \n    os.environ['NCCL_P2P_DISABLE'] = '1'\n    os.environ['NCCL_BUFFSIZE'] = '2097152'  # 2MB buffer size\n    \n    dist.init_process_group(\n        backend='nccl',\n        init_method='env://',\n        rank=rank,\n        world_size=world_size\n    )\n    torch.cuda.set_device(rank)\n\nclass Trainer:\n    \"\"\"Handles the training process with DDP and TV regularization.\"\"\"\n    \n    def __init__(self, rank, world_size, config):\n        self.rank = rank\n        self.world_size = world_size\n        self.config = config\n        self.setup()\n    \n    def setup(self):\n        \"\"\"Initialize training components.\"\"\"\n        (train_inputs, train_outputs, valid_inputs, valid_outputs) = prepare_datasets(self.config)\n        \n        train_sampler = DistributedSampler(\n            SeismicDataset(train_inputs, \n                           train_outputs,\n                           n_examples_per_file=500,\n                           ),\n            num_replicas=self.world_size,\n            rank=self.rank,\n            shuffle=True\n        )\n        \n        self.train_loader = self._create_data_loader(train_sampler, is_train=True)\n        self.valid_loader = self._create_data_loader(\n            DistributedSampler(\n                SeismicDataset(valid_inputs, \n                               valid_outputs,\n                               n_examples_per_file=500,\n                               ),\n                num_replicas=self.world_size,\n                rank=self.rank,\n                shuffle=False\n            ),\n            is_train=False\n        )\n        \n        # Model with TV regularization\n        self.model = DDP(\n            torch.nn.SyncBatchNorm.convert_sync_batchnorm(\n                UNet(**self.config[\"model\"][\"unet_params\"]).to(self.rank)\n            ),\n            device_ids=[self.rank]\n        )\n        \n        \n        # Training components\n        self.criterion = nn.L1Loss()\n        self.optimizer = torch.optim.AdamW(self.model.parameters(), **self.config[\"optimizer\"])\n        self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n            self.optimizer, 'min', **self.config[\"scheduler\"][\"params\"])\n        self.scaler = torch.cuda.amp.GradScaler()\n        \n        # Training state\n        self.best_val_loss = float('inf')\n        self.epochs_wo_improvement = 0\n        self.start_time = time.time()\n    \n    def _create_data_loader(self, sampler, is_train=True):\n        \"\"\"Create a distributed data loader.\"\"\"\n        return DataLoader(\n            sampler.dataset,\n            batch_size=self.config[\"batch_size\"],\n            sampler=sampler,\n            pin_memory=True,\n            drop_last=is_train,\n            num_workers=4\n        )\n    \n    def train_epoch(self, epoch):\n        \"\"\"Execute one training epoch with TV regularization.\"\"\"\n        self.model.train()\n        self.train_loader.sampler.set_epoch(epoch)\n        train_losses = []\n        \n        for step, (inputs, targets) in enumerate(self.train_loader):\n            inputs, targets = inputs.to(self.rank), targets.to(self.rank)\n            \n            with torch.autocast(device_type='cuda', dtype=torch.float16):\n                outputs, tv_loss = self.model(inputs) # Now returns both output and TV loss\n                data_loss = self.criterion(outputs, targets)\n                total_loss = data_loss + self.config['tv_weight'] * tv_loss  # Combine data loss and TV regularization\n            \n            self._update_model(total_loss) # changed\n            train_losses.append(self._reduce_loss(data_loss))  # Only track data loss for logging\n            \n            if self._should_log(step, len(self.train_loader)):\n                self._log_training(epoch, step, train_losses)\n        \n        return np.mean(train_losses) if train_losses else 0\n    \n    def validate(self):\n        \"\"\"Run validation (without TV regularization).\"\"\"\n        self.model.eval()\n        valid_losses = []\n        \n        with torch.inference_mode():\n            for inputs, targets in self.valid_loader:\n                inputs, targets = inputs.to(self.rank), targets.to(self.rank)\n                \n                with torch.autocast(device_type='cuda', dtype=torch.float16):\n                    outputs, _ = self.model(inputs)  # Ignore TV loss during validation\n                    loss = self.criterion(outputs, targets)\n                \n                valid_losses.append(loss.float().item())\n        \n        return self._gather_losses(valid_losses)\n    \n    def _update_model(self, loss):\n        \"\"\"Update model parameters.\"\"\"\n        self.optimizer.zero_grad()\n        self.scaler.scale(loss).backward()\n        self.scaler.step(self.optimizer)\n        self.scaler.update()\n    \n    def _reduce_loss(self, loss):\n        \"\"\"Reduce loss across all processes.\"\"\"\n        torch.distributed.reduce(loss, dst=0)\n        return loss.item() / self.world_size\n    \n    def _gather_losses(self, losses):\n        \"\"\"Gather losses from all processes.\"\"\"\n        \"\"\"Ensure consistent dtype when gathering losses across processes\"\"\"\n        mean_loss = torch.tensor(np.mean(losses), \n                dtype=torch.float32,  # Explicitly set dtype\n                device=self.rank)\n        gathered = [torch.zeros_like(mean_loss) for _ in range(self.world_size)]\n        dist.all_gather(gathered, mean_loss)\n        return torch.stack(gathered).mean().item()\n    \n    def _should_log(self, step, total_steps):\n        \"\"\"Determine if we should log at this step.\"\"\"\n        return (self.rank == 0 and \n               (step % self.config[\"print_freq\"] == self.config[\"print_freq\"] - 1 or \n                step == total_steps - 1))\n    \n    def _log_training(self, epoch, step, losses):\n        \"\"\"Log training progress.\"\"\"\n        trn_loss = np.mean(losses)\n        elapsed = format_time(time.time() - self.start_time)\n        mem_used = torch.cuda.memory_allocated(self.rank) / 1024**3\n        lr = self.optimizer.param_groups[-1]['lr']\n        \n        print(f\"Epoch: {epoch:02d} Step {step+1}/{len(self.train_loader)} \"\n              f\"Trn Loss: {trn_loss:.2f} LR: {lr:.2e} \"\n              f\"GPU Mem: {mem_used:.2f}GB Time: {elapsed}\", flush=True)\n    \n    def run(self):\n        \"\"\"Main training loop.\"\"\"\n        for epoch in range(1, self.config[\"max_epochs\"] + 1):\n            train_loss = self.train_epoch(epoch)\n            val_loss = self.validate()\n            \n            if self.rank == 0:\n                self._log_epoch(epoch, train_loss, val_loss)\n                self._checkpoint_model(val_loss)\n                \n                if self._should_stop(val_loss):\n                    break\n            \n            dist.barrier()\n    \n    def _log_epoch(self, epoch, train_loss, val_loss):\n        \"\"\"Log epoch-level metrics.\"\"\"\n        elapsed = format_time(time.time() - self.start_time)\n        mem_used = torch.cuda.memory_allocated(self.rank) / 1024**3\n        \n        print(f\"\\nEpoch: {epoch:02d} \"\n              f\"Trn Loss: {train_loss:.2f} Val Loss: {val_loss:.2f} \"\n              f\"GPU Mem: {mem_used:.2f}GB Time: {elapsed}\", flush=True)\n        \n    \n    def _checkpoint_model(self, val_loss):\n        \"\"\"Save model if validation loss improves.\"\"\"\n        if val_loss < self.best_val_loss:\n            self.best_val_loss = val_loss\n            self.epochs_wo_improvement = 0\n            torch.save(self.model.module.state_dict(), \"best_model.pth\")\n            print(f\"\\nNew best val_loss: {val_loss:.2f}\\n\", flush=True)\n        else:\n            self.epochs_wo_improvement += 1\n            print(f\"\\nEpochs without improvement: {self.epochs_wo_improvement}\\n\", flush=True)\n    \n    def _should_stop(self, val_loss):\n        \"\"\"Check if training should stop.\"\"\"\n        self.scheduler.step(val_loss)\n        return self.epochs_wo_improvement >= self.config[\"es_epochs\"]\n\n\ndef main(rank, world_size, config):\n    \"\"\"Main training function for each process.\"\"\"\n    try:\n        setup_distributed(rank, world_size)\n        trainer = Trainer(rank, world_size, config)\n        trainer.run()\n    finally:\n        dist.destroy_process_group()\n\nif __name__ == \"__main__\":\n    # Load configuration\n    config = load_config()\n    seed_everything(config[\"seed\"])\n    \n    # Verify GPU availability\n    if torch.cuda.device_count() < 2:\n        raise RuntimeError(\"This script requires at least 2 GPUs\")\n    \n    # Launch training\n    if 'RANK' in os.environ:  # Launched with torchrun\n        world_size = int(os.environ['WORLD_SIZE'])\n        rank = int(os.environ['RANK'])\n        torch.cuda.set_device(int(os.environ['LOCAL_RANK']))\n        main(rank, world_size, config)\n    else:\n        print(\"Launch with: torchrun --nproc_per_node=2 --nnodes=1 \"\n              \"--master_addr=127.0.0.1 --master_port=29500 _train.py\")\n        sys.exit(1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T15:27:56.685479Z","iopub.execute_input":"2025-05-22T15:27:56.686005Z","iopub.status.idle":"2025-05-22T15:27:56.695304Z","shell.execute_reply.started":"2025-05-22T15:27:56.685983Z","shell.execute_reply":"2025-05-22T15:27:56.694661Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"RUN_TRAIN = True\nif RUN_TRAIN:\n    print(\"Starting training..\")\n    !torchrun --nproc_per_node=2 --nnodes=1 --master_addr=127.0.0.1 --master_port=29500 _train.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T15:28:03.430946Z","iopub.execute_input":"2025-05-22T15:28:03.431617Z","execution_failed":"2025-05-22T15:34:01.745Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !tail -f submission.csv  # Verify writes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T10:47:19.97505Z","iopub.execute_input":"2025-05-19T10:47:19.976191Z","iopub.status.idle":"2025-05-19T10:47:37.063331Z","shell.execute_reply.started":"2025-05-19T10:47:19.976156Z","shell.execute_reply":"2025-05-19T10:47:37.062391Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Inference\n\nimport csv\nfrom pathlib import Path\nimport os\nimport sys\nimport yaml\nimport time\nimport torch\nimport numpy as np\nimport torch.distributed as dist\nimport torch.multiprocessing as mp\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, DistributedSampler\nfrom torch.nn.parallel import DistributedDataParallel as DDP\nfrom pprint import pprint\n\n# Local imports\nfrom _dataset import inputs_files_to_output_files, get_train_files, SeismicDataset, TestDataset\nfrom _model import UNet\nfrom _utils import format_time, seed_everything\n\nt0 = time.time()\n\n# Load config\nwith open(\"config.yaml\") as f:\n    config = yaml.safe_load(f)\n\ntest_files = list(Path(\"/kaggle/input/open-wfi-test/test\").glob(\"*.npy\"))\nx_cols = [f\"x_{i}\" for i in range(1, 70, 2)]\nfieldnames = [\"oid_ypos\"] + x_cols\nds = TestDataset(test_files)\ndl = DataLoader(ds, batch_size=config[\"batch_size\"], num_workers=2, pin_memory=False)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = UNet(**config[\"model\"][\"unet_params\"]).to(device)\nif RUN_TRAIN:\n    model.load_state_dict(torch.load(\"best_model.pth\", weights_only=True))\nelse:\n    # model.load_state_dict(torch.load(\"/kaggle/input/fwi_unet_19may/pytorch/default/2/best_model.pth\", weights_only=True))\n    print(\"model is missing!\")\n    \n####\nmodel.eval()\nwith open(\"submission.csv\", \"wt\", newline=\"\") as csvfile:\n    writer = csv.DictWriter(csvfile, fieldnames=fieldnames)\n    writer.writeheader()\n\n    for inputs, oids_test in dl:\n        inputs = inputs.to(device)\n        with torch.inference_mode():\n            with torch.autocast(device_type=\"cuda\"):\n                outputs, _ = model(inputs) # TV regularization\n                \n        y_preds = outputs[:, 0].cpu().numpy()\n\n        for y_pred, oid_test in zip(y_preds, oids_test):\n            for y_pos in range(70):\n                row = dict(zip(x_cols, [y_pred[y_pos, x_pos] for x_pos in range(1, 70, 2)]))\n                row[\"oid_ypos\"] = f\"{oid_test}_y_{y_pos}\"\n\n                writer.writerow(row)\n\nt1 = format_time(time.time() - t0)\nprint(f\"Inference Time: {t1}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T13:59:06.837004Z","iopub.execute_input":"2025-05-21T13:59:06.83777Z","iopub.status.idle":"2025-05-21T13:59:37.179111Z","shell.execute_reply.started":"2025-05-21T13:59:06.83774Z","shell.execute_reply":"2025-05-21T13:59:37.177878Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Multi-GPU\n# !torchrun --nproc_per_node=2 _inference.py","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}