{"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":"gpu","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":8007490,"sourceType":"datasetVersion","datasetId":4476729}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:31.759795Z","iopub.execute_input":"2025-04-12T12:42:31.760137Z","iopub.status.idle":"2025-04-12T12:42:33.864146Z","shell.execute_reply.started":"2025-04-12T12:42:31.760099Z","shell.execute_reply":"2025-04-12T12:42:33.863441Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp /kaggle/input/torchvideo/hgnet.py .\nfrom hgnet import hgnetv2_b5, hgnetv2_b0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:33.868189Z","iopub.execute_input":"2025-04-12T12:42:33.868455Z","iopub.status.idle":"2025-04-12T12:42:35.983726Z","shell.execute_reply.started":"2025-04-12T12:42:33.868432Z","shell.execute_reply":"2025-04-12T12:42:35.982762Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Find files to load and create Dataset","metadata":{}},{"cell_type":"code","source":"all_inputs = [\n    f\n    for f in\n    Path('/kaggle/input/waveform-inversion/train_samples').rglob('*.npy')\n    if ('seis' in f.stem) or ('data' in f.stem)\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:35.984697Z","iopub.execute_input":"2025-04-12T12:42:35.984950Z","iopub.status.idle":"2025-04-12T12:42:36.016795Z","shell.execute_reply.started":"2025-04-12T12:42:35.984929Z","shell.execute_reply":"2025-04-12T12:42:36.015899Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def 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\nall_outputs = inputs_files_to_output_files(all_inputs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.017771Z","iopub.execute_input":"2025-04-12T12:42:36.018139Z","iopub.status.idle":"2025-04-12T12:42:36.022702Z","shell.execute_reply.started":"2025-04-12T12:42:36.018106Z","shell.execute_reply":"2025-04-12T12:42:36.021868Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"assert all(f.exists() for f in all_outputs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.023524Z","iopub.execute_input":"2025-04-12T12:42:36.023828Z","iopub.status.idle":"2025-04-12T12:42:36.047782Z","shell.execute_reply.started":"2025-04-12T12:42:36.023777Z","shell.execute_reply":"2025-04-12T12:42:36.046903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_inputs = [all_inputs[i] for i in range(0, len(all_inputs), 2)] # Sample every two\nvalid_inputs = [f for f in all_inputs if not f in train_inputs]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.048606Z","iopub.execute_input":"2025-04-12T12:42:36.048886Z","iopub.status.idle":"2025-04-12T12:42:36.053174Z","shell.execute_reply.started":"2025-04-12T12:42:36.048852Z","shell.execute_reply":"2025-04-12T12:42:36.052404Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_outputs = inputs_files_to_output_files(train_inputs)\nvalid_outputs = inputs_files_to_output_files(valid_inputs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.055725Z","iopub.execute_input":"2025-04-12T12:42:36.055966Z","iopub.status.idle":"2025-04-12T12:42:36.066733Z","shell.execute_reply.started":"2025-04-12T12:42:36.055946Z","shell.execute_reply":"2025-04-12T12:42:36.066069Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class 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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.067872Z","iopub.execute_input":"2025-04-12T12:42:36.068102Z","iopub.status.idle":"2025-04-12T12:42:36.082888Z","shell.execute_reply.started":"2025-04-12T12:42:36.068082Z","shell.execute_reply":"2025-04-12T12:42:36.082120Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dstrain = SeismicDataset(train_inputs, train_outputs)\ndltrain = DataLoader(dstrain, batch_size=128, shuffle=True, pin_memory=True, drop_last=True, num_workers=4, persistent_workers=True)\n\ndsvalid = SeismicDataset(valid_inputs, valid_outputs)\ndlvalid = DataLoader(dsvalid, batch_size=128, shuffle=False, pin_memory=True, drop_last=False, num_workers=4, persistent_workers=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.083779Z","iopub.execute_input":"2025-04-12T12:42:36.084033Z","iopub.status.idle":"2025-04-12T12:42:36.099908Z","shell.execute_reply.started":"2025-04-12T12:42:36.084013Z","shell.execute_reply":"2025-04-12T12:42:36.099065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dstrain[0][0].shape, dstrain[0][1].shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.100740Z","iopub.execute_input":"2025-04-12T12:42:36.100974Z","iopub.status.idle":"2025-04-12T12:42:36.127202Z","shell.execute_reply.started":"2025-04-12T12:42:36.100955Z","shell.execute_reply":"2025-04-12T12:42:36.126363Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## DumbNet","metadata":{}},{"cell_type":"code","source":"class hgnet_model(nn.Module):\n    def __init__(self, pool_size=(4, 4), input_chans=5, output_size=70 * 70):\n        super().__init__()\n\n        self.backbone = hgnetv2_b0(in_chans=input_chans)\n        self.avg = nn.AdaptiveAvgPool2d((4, 4))\n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(1024 * 4 * 4, 2048),\n            nn.GELU(),\n            nn.Dropout(0.2),\n\n            nn.Linear(2048, 1024),\n            nn.GELU(),\n            nn.Dropout(0.2),\n\n            nn.Linear(1024, output_size)\n        )\n\n    def forward(self, x):\n        bs = x.size(0)\n\n        x = self.backbone.forward_features(x)\n        x = self.avg(x)\n\n        x = self.classifier(x)\n        return x.view(bs, 1, 70, 70) * 1000 + 1500","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.128060Z","iopub.execute_input":"2025-04-12T12:42:36.128283Z","iopub.status.idle":"2025-04-12T12:42:36.133497Z","shell.execute_reply.started":"2025-04-12T12:42:36.128263Z","shell.execute_reply":"2025-04-12T12:42:36.132834Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# class ResidualBlock(nn.Module):\n#     def __init__(self, in_channels, out_channels, downsample=False):\n#         super().__init__()\n#         stride = 2 if downsample else 1\n\n#         self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1, stride=stride)\n#         self.bn1 = nn.BatchNorm2d(out_channels)\n#         self.relu = nn.ReLU(inplace=True)\n\n#         self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)\n#         self.bn2 = nn.BatchNorm2d(out_channels)\n\n#         self.downsample = None\n#         if downsample or in_channels != out_channels:\n#             self.downsample = nn.Sequential(\n#                 nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride),\n#                 nn.BatchNorm2d(out_channels)\n#             )\n\n#     def forward(self, x):\n#         identity = x\n\n#         out = self.relu(self.bn1(self.conv1(x)))\n#         out = self.bn2(self.conv2(out))\n\n#         if self.downsample is not None:\n#             identity = self.downsample(identity)\n\n#         out += identity\n#         return self.relu(out)\n\n# class DumbNet(nn.Module):\n#     '''Deep CNN with residual blocks and dense classifier'''\n#     def __init__(self, input_channels=5, output_size=70 * 70):\n#         super().__init__()\n\n#         self.stem = nn.Sequential(\n#             nn.Conv2d(input_channels, 64, kernel_size=7, stride=2, padding=3),\n#             nn.BatchNorm2d(64),\n#             nn.ReLU(inplace=True),\n#             nn.MaxPool2d(kernel_size=3, stride=2, padding=1)  # 1000x70 -> ~250x18\n#         )\n\n#         self.layer1 = ResidualBlock(64, 128, downsample=True)  # ~125x9\n#         self.layer2 = ResidualBlock(128, 256, downsample=True)  # ~63x5\n#         self.layer3 = ResidualBlock(256, 512, downsample=True)  # ~32x3\n#         self.layer4 = ResidualBlock(512, 1024, downsample=True)\n#         self.layer5 = ResidualBlock(1024, 1024, downsample = False)# same spatial\n\n#         self.global_pool = nn.AdaptiveAvgPool2d((4, 4))  # fixed output\n\n#         self.classifier = nn.Sequential(\n#             nn.Flatten(),\n#             nn.Linear(1024 * 4 * 4, 2048),\n#             nn.GELU(),\n#             nn.Dropout(0.5),\n\n#             nn.Linear(2048, 1024),\n#             nn.GELU(),\n#             nn.Dropout(0.25),\n\n#             nn.Linear(1024, output_size)\n#         )\n\n#     def forward(self, x):\n#         bs = x.size(0)\n\n#         x = self.stem(x)\n#         x = self.layer1(x)\n#         x = self.layer2(x)\n#         x = self.layer3(x)\n#         x = self.layer4(x)\n#         x = self.layer5(x)\n#         x = self.global_pool(x)\n\n#         x = self.classifier(x)\n#         return x.view(bs, 1, 70, 70) * 1000 + 1500","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.134280Z","iopub.execute_input":"2025-04-12T12:42:36.134558Z","iopub.status.idle":"2025-04-12T12:42:36.149470Z","shell.execute_reply.started":"2025-04-12T12:42:36.134529Z","shell.execute_reply":"2025-04-12T12:42:36.148789Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.150276Z","iopub.execute_input":"2025-04-12T12:42:36.150555Z","iopub.status.idle":"2025-04-12T12:42:36.193470Z","shell.execute_reply.started":"2025-04-12T12:42:36.150525Z","shell.execute_reply":"2025-04-12T12:42:36.192449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = hgnet_model().to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.194373Z","iopub.execute_input":"2025-04-12T12:42:36.194666Z","iopub.status.idle":"2025-04-12T12:42:36.828606Z","shell.execute_reply.started":"2025-04-12T12:42:36.194631Z","shell.execute_reply":"2025-04-12T12:42:36.827637Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Loop","metadata":{}},{"cell_type":"code","source":"criterion = nn.L1Loss()\noptim = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.829489Z","iopub.execute_input":"2025-04-12T12:42:36.829747Z","iopub.status.idle":"2025-04-12T12:42:36.834703Z","shell.execute_reply.started":"2025-04-12T12:42:36.829724Z","shell.execute_reply":"2025-04-12T12:42:36.833786Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n_epochs = 50\npatience = 7  \nbest_loss = float('inf')  \ncounter = 0  \n\nhistory = []\n\nfor epoch in range(1, n_epochs+1):\n    print(f'[{epoch:02d}] Begin train')\n\n    # Train\n    model.train()\n    train_losses = []\n    for inputs, targets in tqdm(dltrain, desc='train', leave=False):\n        inputs = inputs.to(device)\n        targets = targets.to(device)\n        \n        optim.zero_grad()\n        \n        outputs = model(inputs)\n        \n        loss = criterion(outputs, targets)\n        \n        loss.backward()\n        \n        optim.step()\n\n        train_losses.append(loss.item())\n\n    print('Train loss: {:.5f}'.format( np.mean(train_losses) ))\n\n    # Valid\n    model.eval()\n    valid_losses = []\n    for inputs, targets in tqdm(dlvalid, desc='valid', leave=False):\n        inputs = inputs.to(device)\n        targets = targets.to(device)\n\n        with torch.inference_mode():\n            outputs = model(inputs)\n        \n        loss = criterion(outputs, targets)\n\n        valid_losses.append(loss.item())\n    \n    print('Valid loss: {:.5f}'.format( np.mean(valid_losses)) )\n    history.append({\n        'train': np.mean(train_losses),\n        'valid': np.mean(valid_losses)\n    })\n\n    # Early stop\n    valid_loss_mean = np.mean(valid_losses)\n    if valid_loss_mean < best_loss:\n        best_loss = valid_loss_mean\n        counter = 0  \n        torch.save(model.state_dict(), \"best_model.pt\")\n    else:\n        counter += 1\n        if counter >= patience:\n            print('Early stopping triggered')\n            break\n\n    # Plot last result\n    if epoch % 4 == 0:\n        y = targets[0, 0].detach().cpu()\n        y_pred = outputs[0, 0].detach().cpu()\n        \n        fig, ax = plt.subplots(1, 2, figsize=(5, 2.5))\n        fig.suptitle(f'Epoch {epoch} | Valid: {np.mean(valid_losses):.5f}')\n        ax[0].imshow(y)\n        ax[1].imshow(y_pred)\n        plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T12:42:36.835574Z","iopub.execute_input":"2025-04-12T12:42:36.835776Z","iopub.status.idle":"2025-04-12T13:04:42.292739Z","shell.execute_reply.started":"2025-04-12T12:42:36.835758Z","shell.execute_reply":"2025-04-12T13:04:42.291387Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pd.DataFrame(history).plot();","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T13:04:45.823033Z","iopub.execute_input":"2025-04-12T13:04:45.823354Z","iopub.status.idle":"2025-04-12T13:04:46.045452Z","shell.execute_reply.started":"2025-04-12T13:04:45.823329Z","shell.execute_reply":"2025-04-12T13:04:46.044478Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Predict test","metadata":{}},{"cell_type":"code","source":"import csv  # Use \"low-level\" CSV to save memory on predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T13:04:48.968561Z","iopub.execute_input":"2025-04-12T13:04:48.968942Z","iopub.status.idle":"2025-04-12T13:04:48.972621Z","shell.execute_reply.started":"2025-04-12T13:04:48.968910Z","shell.execute_reply":"2025-04-12T13:04:48.971713Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\ntest_files = list(Path('/kaggle/input/waveform-inversion/test').glob('*.npy'))\nlen(test_files)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T13:04:49.935378Z","iopub.execute_input":"2025-04-12T13:04:49.935692Z","iopub.status.idle":"2025-04-12T13:04:50.768259Z","shell.execute_reply.started":"2025-04-12T13:04:49.935666Z","shell.execute_reply":"2025-04-12T13:04:50.767326Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x_cols = [f'x_{i}' for i in range(1, 70, 2)]\nfieldnames = ['oid_ypos'] + x_cols","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T13:04:53.393335Z","iopub.execute_input":"2025-04-12T13:04:53.393644Z","iopub.status.idle":"2025-04-12T13:04:53.397603Z","shell.execute_reply.started":"2025-04-12T13:04:53.393619Z","shell.execute_reply":"2025-04-12T13:04:53.396769Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class 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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T13:04:54.445749Z","iopub.execute_input":"2025-04-12T13:04:54.446071Z","iopub.status.idle":"2025-04-12T13:04:54.450731Z","shell.execute_reply.started":"2025-04-12T13:04:54.446048Z","shell.execute_reply":"2025-04-12T13:04:54.449769Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = TestDataset(test_files)\ndl = DataLoader(ds, batch_size=8, num_workers=4, pin_memory=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T13:04:56.105692Z","iopub.execute_input":"2025-04-12T13:04:56.106065Z","iopub.status.idle":"2025-04-12T13:04:56.110407Z","shell.execute_reply.started":"2025-04-12T13:04:56.106033Z","shell.execute_reply":"2025-04-12T13:04:56.109575Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(torch.load('/kaggle/working/best_model.pt'))\n\n# Train\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 tqdm(dl, desc='test'):\n        inputs = inputs.to(device)\n        with torch.inference_mode():\n            outputs = model(inputs)\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(\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            \n                writer.writerow(row)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T13:05:41.348504Z","iopub.execute_input":"2025-04-12T13:05:41.348879Z","iopub.status.idle":"2025-04-12T13:06:01.578516Z","shell.execute_reply.started":"2025-04-12T13:05:41.348848Z","shell.execute_reply":"2025-04-12T13:06:01.577351Z"}},"outputs":[],"execution_count":null}]}