{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","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,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, random_split\n\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\n\nimport glob\nimport os\nimport numpy as np\nimport pandas as pd\n\nimport 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 colorama import init, Fore, Style\nfrom torch.utils.data import Dataset, DataLoader\nfrom matplotlib.gridspec import GridSpec\n\ninit(autoreset=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:19:14.640740Z","iopub.execute_input":"2025-11-20T08:19:14.641084Z","iopub.status.idle":"2025-11-20T08:19:14.646937Z","shell.execute_reply.started":"2025-11-20T08:19:14.641063Z","shell.execute_reply":"2025-11-20T08:19:14.646088Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# *Yale/UNC-CH - Geophysical Waveform Inversion*","metadata":{}},{"cell_type":"markdown","source":"##  Data Loading Function","metadata":{}},{"cell_type":"code","source":"def load_seismic_file(file_path, target_shape=(100, 70)):\n   \n    data = np.load(file_path).astype(np.float32)\n\n   \n    while data.ndim > 3:\n        data = data[0]\n\n   \n    if data.shape[1:] != target_shape:\n        data = data[:, :target_shape[0], :target_shape[1]]\n\n   \n    data = (data - data.mean()) / (data.std() + 1e-8)\n\n    return data  # shape: (S, T, R)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:19:14.648207Z","iopub.execute_input":"2025-11-20T08:19:14.648426Z","iopub.status.idle":"2025-11-20T08:19:14.663801Z","shell.execute_reply.started":"2025-11-20T08:19:14.648410Z","shell.execute_reply":"2025-11-20T08:19:14.663147Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##  PyTorch Dataset","metadata":{}},{"cell_type":"code","source":"class WaveformDataset(Dataset):\n    def __init__(self, input_files, target_arrays, transform_input=None, target_shape=(100, 70)):\n        self.input_files = input_files\n        self.target_arrays = target_arrays\n        self.transform_input = transform_input\n        self.target_shape = target_shape\n\n    def __len__(self):\n        return len(self.input_files)\n\n    def __getitem__(self, idx):\n        data = load_seismic_file(self.input_files[idx], self.target_shape)  # (S, T, R)\n        if self.transform_input:\n            data = self.transform_input(data)\n\n        data_tensor = torch.tensor(data, dtype=torch.float32).unsqueeze(0)  # (1, S, T, R)\n        data_tensor = data_tensor.permute(0, 3, 1, 2)  # (1, R, S, T)\n\n        target_tensor = torch.tensor(self.target_arrays[idx], dtype=torch.float32)  # (H, W)\n        return data_tensor, target_tensor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:19:14.664563Z","iopub.execute_input":"2025-11-20T08:19:14.664798Z","iopub.status.idle":"2025-11-20T08:19:14.684691Z","shell.execute_reply.started":"2025-11-20T08:19:14.664774Z","shell.execute_reply":"2025-11-20T08:19:14.683977Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##  Simplified 3D U-Net Model","metadata":{}},{"cell_type":"code","source":"class UNet3D(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super(UNet3D, self).__init__()\n        self.encoder = nn.Sequential(\n            nn.Conv3d(in_channels, 8, kernel_size=3, padding=1),\n            nn.BatchNorm3d(8),\n            nn.ReLU(True),\n            nn.Conv3d(8, 16, kernel_size=3, padding=1),\n            nn.BatchNorm3d(16),\n            nn.ReLU(True),\n            nn.AdaptiveAvgPool3d((1, 100, 70))\n        )\n        self.decoder = nn.Sequential(\n            nn.ConvTranspose3d(16, 8, kernel_size=3, padding=1),\n            nn.BatchNorm3d(8),\n            nn.ReLU(True),\n            nn.ConvTranspose3d(8, out_channels, kernel_size=3, padding=1),\n        )\n\n    def forward(self, x):\n        x = self.encoder(x)    # -> (B, 16, 1, 100, 70)\n        x = self.decoder(x)    # -> (B, 1, 1, 100, 70)\n        x = x.squeeze(2)       # -> (B, 1, 100, 70)\n        return x.squeeze(1)    # -> (B, 100, 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:19:14.686407Z","iopub.execute_input":"2025-11-20T08:19:14.686623Z","iopub.status.idle":"2025-11-20T08:19:14.698654Z","shell.execute_reply.started":"2025-11-20T08:19:14.686608Z","shell.execute_reply":"2025-11-20T08:19:14.697823Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##  Sample Data Setup","metadata":{}},{"cell_type":"code","source":"example_file   = \"/kaggle/input/waveform-inversion/train_samples/CurveFault_A/seis2_1_0.npy\"\nexample_target = np.zeros((100, 70), dtype=np.float32)\nfile_list      = [example_file] * 10\ntarget_list    = [example_target] * 10\n\ndataset = WaveformDataset(file_list, target_list)\ntrain_size = int(0.8 * len(dataset))\nval_size = len(dataset) - train_size\ntrain_dataset, val_dataset = random_split(dataset, [train_size, val_size])\n\ntrain_loader = DataLoader(train_dataset, batch_size=2, shuffle=True)\nval_loader   = DataLoader(val_dataset, batch_size=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:19:14.699453Z","iopub.execute_input":"2025-11-20T08:19:14.699690Z","iopub.status.idle":"2025-11-20T08:19:14.716443Z","shell.execute_reply.started":"2025-11-20T08:19:14.699668Z","shell.execute_reply":"2025-11-20T08:19:14.715894Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##  Model Training Setup¶\n","metadata":{}},{"cell_type":"code","source":"device    = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel     = UNet3D(in_channels=1, out_channels=1).to(device)\ncriterion = nn.L1Loss()\noptimizer = optim.Adam(model.parameters(), lr=5e-4)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer,\n                                                 mode='min',\n                                                 factor=0.5,\n                                                 patience=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:19:14.717067Z","iopub.execute_input":"2025-11-20T08:19:14.717242Z","iopub.status.idle":"2025-11-20T08:19:14.733316Z","shell.execute_reply.started":"2025-11-20T08:19:14.717228Z","shell.execute_reply":"2025-11-20T08:19:14.732569Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Loop with Validation & Early Stopping","metadata":{}},{"cell_type":"code","source":"num_epochs = 30\nbest_val_loss = float(\"inf\")\nearly_stop_counter = 0\ntrain_losses = []\nval_losses = []\n\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs} Training\")\n\n    for inputs, targets in progress_bar:\n        inputs  = inputs.to(device)\n        targets = targets.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss    = criterion(outputs, targets)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n        progress_bar.set_postfix({\"loss\": loss.item()})\n\n    avg_train_loss = running_loss / len(train_loader)\n    train_losses.append(avg_train_loss)\n\n    # Validation step\n    model.eval()\n    val_loss = 0.0\n    with torch.no_grad():\n        for inputs, targets in val_loader:\n            inputs = inputs.to(device)\n            targets = targets.to(device)\n            outputs = model(inputs)\n            loss = criterion(outputs, targets)\n            val_loss += loss.item()\n\n    avg_val_loss = val_loss / len(val_loader)\n    val_losses.append(avg_val_loss)\n\n    print(f\"Epoch {epoch+1}: Train Loss = {avg_train_loss:.4f}, Val Loss = {avg_val_loss:.4f}\")\n\n    # Save best model\n    if avg_val_loss < best_val_loss:\n        best_val_loss = avg_val_loss\n        torch.save(model.state_dict(), \"best_model.pth\")\n        early_stop_counter = 0\n    else:\n        early_stop_counter += 1\n\n    scheduler.step(avg_val_loss)\n\n    if early_stop_counter >= 5:\n        print(\"Early stopping triggered.\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:19:14.734259Z","iopub.execute_input":"2025-11-20T08:19:14.734466Z","iopub.status.idle":"2025-11-20T08:19:51.327223Z","shell.execute_reply.started":"2025-11-20T08:19:14.734450Z","shell.execute_reply":"2025-11-20T08:19:51.326553Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##  Plot Loss Curve","metadata":{}},{"cell_type":"code","source":"plt.plot(train_losses, label=\"Training Loss\")\nplt.plot(val_losses, label=\"Validation Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.legend()\nplt.title(\"Training and Validation Loss Curve\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:19:51.327996Z","iopub.execute_input":"2025-11-20T08:19:51.328279Z","iopub.status.idle":"2025-11-20T08:19:51.499912Z","shell.execute_reply.started":"2025-11-20T08:19:51.328262Z","shell.execute_reply":"2025-11-20T08:19:51.499245Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##  Inference & Submission","metadata":{}},{"cell_type":"code","source":"def predict(model, file_paths):\n    model.eval()\n    predictions = []\n    with torch.no_grad():\n        for fp in file_paths:\n            data = load_seismic_file(fp)\n            tensor = torch.tensor(data, dtype=torch.float32).unsqueeze(0).unsqueeze(0).to(device)\n            output = model(tensor)\n            predictions.append(output.squeeze().cpu().numpy())\n    return predictions\n\ndef create_submission(oids, predictions):\n    sample_path = '/kaggle/input/waveform-inversion/sample_submission.csv'\n    sample_df = pd.read_csv(sample_path)\n    id_col = sample_df.columns[0]\n\n    width = predictions[0].shape[1]\n    odd_indices = list(range(0, width, 2))\n\n    rows = []\n    for oid, pred in zip(oids, predictions):\n        for y in range(pred.shape[0]):\n            row_id = f\"{oid}_y_{y}\"\n            row = [row_id] + [float(pred[y, x]) for x in odd_indices]\n            rows.append(row)\n\n    columns = [id_col] + [f\"x_{i}\" for i in odd_indices]\n    df = pd.DataFrame(rows, columns=columns)\n    df.to_csv('/kaggle/working/submission.csv', index=False)\n    print(\"Submission saved to /kaggle/working/submission.csv.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:19:51.500665Z","iopub.execute_input":"2025-11-20T08:19:51.500898Z","iopub.status.idle":"2025-11-20T08:19:51.507823Z","shell.execute_reply.started":"2025-11-20T08:19:51.500882Z","shell.execute_reply":"2025-11-20T08:19:51.507095Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##  Generate Final Submission","metadata":{}},{"cell_type":"code","source":"test_files = glob.glob(\"/kaggle/input/waveform-inversion/test_samples/**/*.npy\", recursive=True)\nif test_files:\n    model.load_state_dict(torch.load(\"best_model.pth\"))\n    preds = predict(model, test_files)\n    oids = [os.path.splitext(os.path.basename(fp))[0] for fp in test_files]\n    create_submission(oids, preds)\nelse:\n    # Fallback to zero-filled submission\n    sample_path = '/kaggle/input/waveform-inversion/sample_submission.csv'\n    sample_df = pd.read_csv(sample_path)\n    sample_df.iloc[:, 1:] = 0.0\n    sample_df.to_csv('/kaggle/working/submission.csv', index=False)\n    print(\"Fallback submission written to /kaggle/working/submission.csv.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T08:19:51.509416Z","iopub.execute_input":"2025-11-20T08:19:51.509572Z","iopub.status.idle":"2025-11-20T08:21:22.588610Z","shell.execute_reply.started":"2025-11-20T08:19:51.509560Z","shell.execute_reply":"2025-11-20T08:21:22.587876Z"}},"outputs":[],"execution_count":null}]}