{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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"},{"sourceId":11867571,"sourceType":"datasetVersion","datasetId":7155401},{"sourceId":332421,"sourceType":"modelInstanceVersion","modelInstanceId":278650,"modelId":299553}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# UNet with Classification\n\nThis uses the model found [here](https://www.kaggle.com/code/egortrushin/gwi-improved-unet-pipepline-with-larger-dataset), and the classifier found [here](https://www.kaggle.com/code/johnnyhyland/waveform-classifier). The classifier selects the model, each model was trained on one type.","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.data import random_split, DataLoader\nimport torch.optim as optim\nfrom sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay\nfrom torchvision.models import efficientnet_b0\nimport csv\n\nmodel_category = ['CurveFault_A', \n                  'CurveFault_B',\n                  'CurveVel_A',\n                  'CurveVel_B',\n                  'FlatFault_A',\n                  'FlatFault_B',\n                  'FlatVel_A',\n                  'FlatVel_A',\n                  'Style_A',\n                  'Style_B',\n                 ]\nmodel_number = 1\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        # 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 ResidualDoubleConv(nn.Module):\n    \"\"\"(Convolution => [BN] => ReLU) * 2 + Residual Connection\"\"\"\n\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        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        self.conv2 = nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n\n        if in_channels == out_channels:\n            self.shortcut = nn.Identity()\n        else:\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\n        out = self.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n        identity_mapped = self.shortcut(identity)\n        out += identity_mapped\n        return self.relu(out)\n\n\nclass Up(nn.Module):\n    \"\"\"Upscaling then ResidualDoubleConv\"\"\"\n\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            self.conv = ResidualDoubleConv(in_channels + out_channels, out_channels)\n        else:\n            self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)\n            conv_in_channels = in_channels // 2\n            skip_channels = out_channels\n            total_in_channels = conv_in_channels + skip_channels\n            self.conv = ResidualDoubleConv(total_in_channels, out_channels)\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        diffY = x2.size(2) - x1.size(2)\n        diffX = x2.size(3) - x1.size(3)\n        x1 = F.pad(\n            x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]\n        )\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    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 architecture implementation with Residual Blocks\"\"\"\n    def __init__(\n        self,\n        n_channels=5,    # Default to 5 input channels\n        n_classes=1,     # Default to 1 output class (for regression)\n        init_features=32,\n        depth=5, \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\n        self.initial_pool = nn.AvgPool2d(kernel_size=(14, 1), stride=(14, 1)) \n\n        self.encoder_convs = nn.ModuleList()\n        self.encoder_pools = nn.ModuleList()\n\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            conv = ResidualDoubleConv(current_features, current_features * 2)\n            pool = nn.MaxPool2d(2)\n            self.encoder_convs.append(conv)\n            self.encoder_pools.append(pool)\n            current_features *= 2\n        \n        self.bottleneck = ResidualDoubleConv(current_features, current_features)\n\n        self.decoder_blocks = nn.ModuleList()\n        for _ in range(depth):\n            up_block = Up(current_features, current_features // 2, bilinear)\n            self.decoder_blocks.append(up_block)\n            current_features //= 2\n            \n        self.outc = OutConv(current_features, n_classes)\n\n    def _pad_or_crop(self, x, target_h=70, target_w=70):\n        _, _, h, w = x.shape\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))\n        elif h > target_h:\n            crop_top = (h - target_h) // 2\n            x = x[:, :, crop_top : crop_top + target_h, :]\n        \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))\n        elif w > target_w:\n            crop_left = (w - target_w) // 2\n            x = x[:, :, :, crop_left : crop_left + target_w]\n        return x\n\n    def forward(self, x):\n        # x initial shape: (bs, 5, 1000, 70)\n        x_pooled = self.initial_pool(x) # (bs, 5, 71, 70)\n        x_resized = self._pad_or_crop(x_pooled, target_h=70, target_w=70) # (bs, 5, 70, 70)\n\n        skip_connections = []\n        xi = x_resized\n\n        xi = self.encoder_convs[0](xi) # inc\n        skip_connections.append(xi)\n\n        for i in range(self.depth):\n            xi = self.encoder_convs[i+1](xi)\n            skip_connections.append(xi)\n            xi = self.encoder_pools[i](xi)\n        \n        xi = self.bottleneck(xi)\n\n        xu = xi\n        for i, block in enumerate(self.decoder_blocks):\n            skip_index = self.depth - 1 - i \n            skip = skip_connections[skip_index]\n            xu = block(xu, skip)\n            \n        logits = self.outc(xu) # Should be (bs, n_classes, 70, 70) e.g. (bs, 1, 70, 70)\n        \n        # Apply scaling and offset specific to the problem's target range\n        output = logits * 1000.0 + 1500.0 \n        return output\n# --- End of UNet Model Definition ---\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n# model = DumbNet().to(device)\n\nclass SeismicEfficientNetClassifier(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n        # Load EfficientNet-B0 backbone\n        self.backbone = efficientnet_b0(weights=None)\n        # Adapt input conv layer if needed (e.g., 1 channel instead of 3)\n        if self.backbone.features[0][0].in_channels != 1:\n            self.backbone.features[0][0] = nn.Conv2d(\n                1, 32, kernel_size=3, stride=2, padding=1, bias=False\n            )\n        # Replace classifier head\n        in_features = self.backbone.classifier[1].in_features\n        self.backbone.classifier[1] = nn.Linear(in_features, num_classes)\n\n    def forward(self, x):\n        # x: (batch, channels=1, H, W)\n        return self.backbone(x)\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-15T20:21:04.244507Z","iopub.execute_input":"2025-04-15T20:21:04.244809Z","iopub.status.idle":"2025-04-15T20:21:11.505104Z","shell.execute_reply.started":"2025-04-15T20:21:04.244775Z","shell.execute_reply":"2025-04-15T20:21:11.504430Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Predict test","metadata":{}},{"cell_type":"code","source":"import torch\nimport csv\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\n\ntest_files = list(Path('/kaggle/input/waveform-inversion/test').glob('*.npy'))\nx_cols = [f'x_{i}' for i in range(1, 70, 2)]\nfieldnames = ['oid_ypos'] + x_cols\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        return np.load(test_file), test_file.stem\n\nclassifier_path = '/kaggle/input/waveform-inversion-models/classifier.pth'\n\nclassifier = SeismicEfficientNetClassifier(num_classes=10)\nclassifier.load_state_dict(torch.load(classifier_path, map_location=device))\nclassifier.to(device)\nclassifier.eval()\n\nall_models = []\nfor i, cat in enumerate(model_category):\n    model_path = f'/kaggle/input/waveform-inversion-models/{cat}_unet_model.pth'\n    m = UNet().to(device)\n    m.load_state_dict(torch.load(model_path, map_location=device))\n    m.eval()\n    all_models.append(m)\n\nds = TestDataset(test_files)\ndl = DataLoader(ds, batch_size=1, num_workers=4, pin_memory=True)\n\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 tqdm(dl, desc='test'):\n        # Classify the test file\n        x = inputs.mean(dim=1, keepdim=True).to(device)  # (1, 1, time, receivers)\n        with torch.no_grad():\n            class_logits = classifier(x)\n            pred_class = class_logits.argmax(dim=1).item()\n\n        # Use the correct model for this class\n        model = all_models[pred_class]\n        with torch.no_grad():\n            # DumbNet expects (batch, 5, time, receivers)\n            y_pred = model(inputs.to(device)).cpu().numpy()[0, 0]  # shape (70, 70)\n\n        oid_test = oids_test[0]\n        for y_pos in range(70):\n            row = dict(\n                zip(\n                    x_cols,\n                    [y_pred[y_pos, x_pos] for x_pos in range(1, 70, 2)]\n                )\n            )\n            row['oid_ypos'] = f\"{oid_test}_y_{y_pos}\"\n            writer.writerow(row)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T20:23:16.833457Z","iopub.execute_input":"2025-04-15T20:23:16.833754Z","execution_failed":"2025-04-15T20:24:14.931Z"}},"outputs":[],"execution_count":null}]}