{"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"},{"sourceId":11367935,"sourceType":"datasetVersion","datasetId":7116013},{"sourceId":11368083,"sourceType":"datasetVersion","datasetId":7116134},{"sourceId":11368169,"sourceType":"datasetVersion","datasetId":7116196},{"sourceId":11368268,"sourceType":"datasetVersion","datasetId":7116272},{"sourceId":11368499,"sourceType":"datasetVersion","datasetId":7116445},{"sourceId":11368545,"sourceType":"datasetVersion","datasetId":7116479},{"sourceId":11368547,"sourceType":"datasetVersion","datasetId":7116481},{"sourceId":11376433,"sourceType":"datasetVersion","datasetId":7122462},{"sourceId":11376448,"sourceType":"datasetVersion","datasetId":7122476},{"sourceId":11376464,"sourceType":"datasetVersion","datasetId":7122489},{"sourceId":11376742,"sourceType":"datasetVersion","datasetId":7122712},{"sourceId":11376868,"sourceType":"datasetVersion","datasetId":7122812},{"sourceId":11376871,"sourceType":"datasetVersion","datasetId":7122814},{"sourceId":11376872,"sourceType":"datasetVersion","datasetId":7122815},{"sourceId":11376935,"sourceType":"datasetVersion","datasetId":7122866},{"sourceId":11377083,"sourceType":"datasetVersion","datasetId":7122981},{"sourceId":11377231,"sourceType":"datasetVersion","datasetId":7123090},{"sourceId":11377291,"sourceType":"datasetVersion","datasetId":7123138},{"sourceId":11377325,"sourceType":"datasetVersion","datasetId":7123163},{"sourceId":11377334,"sourceType":"datasetVersion","datasetId":7123172},{"sourceId":11377594,"sourceType":"datasetVersion","datasetId":7123380},{"sourceId":11377614,"sourceType":"datasetVersion","datasetId":7123394},{"sourceId":11377741,"sourceType":"datasetVersion","datasetId":7123490},{"sourceId":11377752,"sourceType":"datasetVersion","datasetId":7123499},{"sourceId":11377756,"sourceType":"datasetVersion","datasetId":7123503},{"sourceId":11377935,"sourceType":"datasetVersion","datasetId":7123649},{"sourceId":11377970,"sourceType":"datasetVersion","datasetId":7123675},{"sourceId":11378141,"sourceType":"datasetVersion","datasetId":7123813},{"sourceId":11378162,"sourceType":"datasetVersion","datasetId":7123831},{"sourceId":11378178,"sourceType":"datasetVersion","datasetId":7123842}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"#### Seismic input is downsampled prior to model training runs to save memory. \n#### Predictions are adequate but limited by low resolution input. ","metadata":{}},{"cell_type":"markdown","source":"## Initialize","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nfrom pathlib import Path\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch import nn\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import mean_squared_error\nfrom sklearn.model_selection import train_test_split\nimport seaborn as sns\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nimport matplotlib.colors as mcolors\nimport gc\n\n# Verify GPU availability\nprint(f\"PyTorch Version: {torch.__version__}\")\nprint(f\"CUDA Available: {torch.cuda.is_available()}\")\n\nif torch.cuda.is_available(): \n    device = torch.device(\"cuda:0\") \n    print(f\"GPU Name: {torch.cuda.get_device_name(0)}\") \n    print(f\"CUDA Version: {torch.version.cuda}\")\n    print(f\"Number of GPUs: {torch.cuda.device_count()}\")\nelse: \n    device = torch.device(\"cpu\") \n    print(f\"Using device: {device}\")\n    \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T01:24:07.647519Z","iopub.execute_input":"2025-05-11T01:24:07.647915Z","iopub.status.idle":"2025-05-11T01:24:16.785383Z","shell.execute_reply.started":"2025-05-11T01:24:07.647813Z","shell.execute_reply":"2025-05-11T01:24:16.784504Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Find all sample directories\n### Use Kaggle Original and Open FWI database","metadata":{}},{"cell_type":"code","source":"# New updated \n# Collect Original Kaggle input files\ninput_dir = Path('/kaggle/input/waveform-inversion/train_samples')\n\ndata_files = []\nmodel_files = []\n\n# Add files from Vel and Style families\n\ndata_files = sorted(input_dir.rglob('data*.npy'))\nmodel_files = sorted(input_dir.rglob('model*.npy'))\n\n# Add Fault family\nfault_data_files = sorted(input_dir.rglob('seis*.npy'))\nfault_model_files = sorted(input_dir.rglob('vel*.npy'))\ndata_files += fault_data_files\nmodel_files += fault_model_files\n\n#LOAD the OPEN FWI training files\nroot_dir = Path(\"/kaggle/input\")\n\nall_data_files = []\nall_model_files = []\nseis_files = []\nvel_files = []\nseis2_files = []\nvel2_files = []\n\n# Loop over waveform-inversion-1 to waveform-inversion-30\nfor i in range(1,31):\n    subroot = root_dir / f\"waveform-inversion-{i}\"\n    seis_files = sorted(subroot.rglob(\"seis*.npy\"))\n    vel_files = sorted(subroot.rglob(\"vel*.npy\"))\n    all_data_files.extend(seis_files)\n    all_model_files.extend(vel_files)\n\n    seis2_files = sorted(subroot.rglob(\"data*.npy\"))\n    vel2_files = sorted(subroot.rglob(\"model*.npy\"))\n    all_data_files.extend(seis2_files)\n    all_model_files.extend(vel2_files)\n    seis_files = []\n    vel_files = [] \n    seis2_files = []\n    vel2_files = []\n  \n## File check\n#for i in range(1, 20):\n#    print(f\"index: {i}\" )\n#    print(\"data\", all_data_files[i])\n#    print(\"model\", all_model_files[i])   \n\n# Randomly select some Open FWI consistent pairs\n\nprint(\"Number of Open FWI seismic training files found:\", len(all_data_files))\n\nimport random\n# Set a fixed seed for reproducibility\nrandom.seed(42)\ndata_FWIfiles = []\nmodel_FWIfiles = []\n\n#Max number tested is 80, larger may crash memory\nselect_FWI  = 0\nprint(\"Number of Open FWI to add to original set:\", select_FWI)\nindices = random.sample(range(len(all_data_files)), select_FWI)\ndata_FWIfiles = [all_data_files[i] for i in indices]\nmodel_FWIfiles = [all_model_files[i] for i in indices]\n\n# Combine Original Kaggle and select Open FWI files\ndata_files += data_FWIfiles\nmodel_files += model_FWIfiles\n\nprint(\"Number of combined origianl/Open FWI training files found:\", len(data_files))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T01:28:51.835375Z","iopub.execute_input":"2025-05-11T01:28:51.836298Z","iopub.status.idle":"2025-05-11T01:28:57.584342Z","shell.execute_reply.started":"2025-05-11T01:28:51.836264Z","shell.execute_reply":"2025-05-11T01:28:57.583303Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load and preprocess sample files ","metadata":{}},{"cell_type":"code","source":"\nX_list = []\ny_list = []\n\nfor count, (data_path, model_path) in enumerate(zip(data_files, model_files), start=1):\n    try:\n        print(f\"Loading file {count} of {len(data_files)}\")\n        print(\"File\", data_path)\n        \n        # Memory-mapped loading\n        X = np.load(data_path, mmap_mode='r')  # (500, 5, 980, 70)\n        y = np.load(model_path, mmap_mode='r') # (500, 1, 70, 70)\n\n        # Preprocessing\n        X = X[:, :, 20:1000, :]        # -> (500, 5, 980, 70)\n        X = X[:, :, ::14, :]           # -> (500, 5, 70, 70)\n        X = X.transpose(0, 2, 1, 3)    # -> (500, 70, 5, 70)\n        X = X.reshape(500, 70, 350)    # -> (500, 70, 350)\n        X = X[:, :, ::5]               # -> (500, 70, 70)\n\n        if y.ndim == 4:\n            y = np.squeeze(y, axis=1)  # -> (500, 70, 70)\n        y = (y - 1500.0) / (4500.0 - 1500.0)\n\n        X_list.append(X)\n        y_list.append(y)\n\n    except Exception as e:\n        print(f\"Error loading {data_path} or {model_path}: {str(e)}\")\n        continue\n\n# Final conversion (only once, fast)\nX_data = np.concatenate(X_list, axis=0)\ny_data = np.concatenate(y_list, axis=0)\n\nprint(\"X_data shape:\", X_data.shape)\nprint(\"y_data shape:\", y_data.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T01:29:01.555223Z","iopub.execute_input":"2025-05-11T01:29:01.555529Z","iopub.status.idle":"2025-05-11T01:32:12.610052Z","shell.execute_reply.started":"2025-05-11T01:29:01.555507Z","shell.execute_reply":"2025-05-11T01:32:12.608935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Train/Val Split (80/20)\nX_train, X_val, y_train, y_val = train_test_split(\n    X_data, y_data, test_size=0.2, random_state=42\n)\n\nprint(\"X_train shape:\", X_train.shape)\nprint(\"X_val shape:\", X_val.shape)\nprint(\"y_train shape:\", y_train.shape)\nprint(\"y_val shape:\", y_val.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T01:38:14.885790Z","iopub.execute_input":"2025-05-11T01:38:14.886312Z","iopub.status.idle":"2025-05-11T01:38:15.079787Z","shell.execute_reply.started":"2025-05-11T01:38:14.886252Z","shell.execute_reply":"2025-05-11T01:38:15.078748Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visualize training samples","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Loop through all samples\nfor sample_idx in range(0, len(X_train), 5000):\n    seismic = X_train[sample_idx]      # shape (70, 70)\n    velocity = y_train[sample_idx]     # shape (70, 70)\n\n    # Plot them\n    plt.figure(figsize=(12, 5))\n\n    # Seismic Input\n    plt.subplot(1, 2, 1)\n    plt.imshow(seismic, aspect='auto', cmap='seismic', origin='lower')\n    plt.title(f\"Seismic Input (Sample {sample_idx})\")\n    plt.xlabel(\"Fused Shots Offset\")\n    plt.ylabel(\"Time\")\n    plt.colorbar(label=\"Amplitude\")\n\n    # Velocity Model (Target)\n    plt.subplot(1, 2, 2)\n    plt.imshow(velocity, aspect='equal', cmap='viridis', origin='lower', vmin=0, vmax=1)\n    plt.title(\"Velocity Model (Target)\")\n    plt.xlabel(\"X Position\")\n    plt.ylabel(\"Depth\")\n    plt.colorbar(label=\"Velocity\")\n\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T01:38:17.678529Z","iopub.execute_input":"2025-05-11T01:38:17.679192Z","iopub.status.idle":"2025-05-11T01:38:19.519857Z","shell.execute_reply.started":"2025-05-11T01:38:17.679162Z","shell.execute_reply":"2025-05-11T01:38:19.518798Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##  UNet - 5 Layer - attention bottleneck","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nimport torch\n\n# Custom dataset\nclass SeismicVelocityDataset(Dataset):\n    def __init__(self, X, y):\n        self.X = torch.tensor(X, dtype=torch.float32)  # (N, 70, 70)\n        self.y = torch.tensor(y, dtype=torch.float32)  # (N, 70, 70)\n\n    def __len__(self):\n        return len(self.X)\n\n    def __getitem__(self, idx):\n        return self.X[idx].unsqueeze(0), self.y[idx].unsqueeze(0)  # Add channel dimension (1, 70, 70)\n\n# Create datasets\ntrain_dataset = SeismicVelocityDataset(X_train, y_train)\nval_dataset = SeismicVelocityDataset(X_val, y_val)\n\n# DataLoaders\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T01:38:22.519005Z","iopub.execute_input":"2025-05-11T01:38:22.519331Z","iopub.status.idle":"2025-05-11T01:38:22.976571Z","shell.execute_reply.started":"2025-05-11T01:38:22.519305Z","shell.execute_reply":"2025-05-11T01:38:22.975798Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n# --- DoubleConv block ---\nclass DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(DoubleConv, self).__init__()\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, 3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, 3, padding=1),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)\n\n# --- Attention bottleneck ---\nclass BottleneckAttention(nn.Module):\n    def __init__(self, dim, heads=4):\n        super().__init__()\n        self.conv_in = nn.Conv2d(dim, dim, kernel_size=1)\n        self.norm = nn.LayerNorm(dim)\n        self.attn = nn.MultiheadAttention(embed_dim=dim, num_heads=heads, batch_first=True)\n        self.conv_out = nn.Conv2d(dim, dim, kernel_size=1)\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        x = self.conv_in(x)\n        x_flat = x.view(B, C, -1).transpose(1, 2)  # (B, HW, C)\n        x_flat = self.norm(x_flat)                # Apply LN on (B, HW, C)\n        attn_out, _ = self.attn(x_flat, x_flat, x_flat)\n        attn_out = attn_out.transpose(1, 2).view(B, C, H, W)\n        return self.conv_out(attn_out)\n\n# --- UNet with 5 levels and Attention ---\nclass UNet5Attention(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super(UNet5Attention, self).__init__()\n\n        self.enc1 = DoubleConv(in_channels, 16)\n        self.pool1 = nn.MaxPool2d(2)\n        \n        self.enc2 = DoubleConv(16, 32)\n        self.pool2 = nn.MaxPool2d(2)\n\n        self.enc3 = DoubleConv(32, 64)\n        self.pool3 = nn.MaxPool2d(2)\n\n        self.enc4 = DoubleConv(64, 128)\n        self.pool4 = nn.MaxPool2d(2)\n\n        self.enc5 = DoubleConv(128, 256)\n        self.pool5 = nn.MaxPool2d(2)\n\n        self.bottleneck_pre = nn.Conv2d(256, 512, kernel_size=3, padding=1)\n        self.bottleneck = BottleneckAttention(dim=512, heads=4)\n\n        self.up5 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.dec5 = DoubleConv(512 + 256, 256)\n\n        self.up4 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.dec4 = DoubleConv(256 + 128, 128)\n\n        self.up3 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.dec3 = DoubleConv(128 + 64, 64)\n\n        self.up2 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.dec2 = DoubleConv(64 + 32, 32)\n\n        self.up1 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.dec1 = DoubleConv(32 + 16, 16)\n\n        self.final = nn.Sequential(\n            nn.Conv2d(16, out_channels, kernel_size=1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        e1 = self.enc1(x)\n        p1 = self.pool1(e1)\n\n        e2 = self.enc2(p1)\n        p2 = self.pool2(e2)\n\n        e3 = self.enc3(p2)\n        p3 = self.pool3(e3)\n\n        e4 = self.enc4(p3)\n        p4 = self.pool4(e4)\n\n        e5 = self.enc5(p4)\n        p5 = self.pool5(e5)\n\n        b = self.bottleneck_pre(p5)\n        b = self.bottleneck(b)\n\n        u5 = self._pad_to_match(self.up5(b), e5)\n        d5 = self.dec5(torch.cat([u5, e5], dim=1))\n\n        u4 = self._pad_to_match(self.up4(d5), e4)\n        d4 = self.dec4(torch.cat([u4, e4], dim=1))\n\n        u3 = self._pad_to_match(self.up3(d4), e3)\n        d3 = self.dec3(torch.cat([u3, e3], dim=1))\n\n        u2 = self._pad_to_match(self.up2(d3), e2)\n        d2 = self.dec2(torch.cat([u2, e2], dim=1))\n\n        u1 = self._pad_to_match(self.up1(d2), e1)\n        d1 = self.dec1(torch.cat([u1, e1], dim=1))\n\n        return self.final(d1)\n\n    def _pad_to_match(self, upsampled, target):\n        diffY = target.size(2) - upsampled.size(2)\n        diffX = target.size(3) - upsampled.size(3)\n        return F.pad(upsampled, [\n            diffX // 2, diffX - diffX // 2,\n            diffY // 2, diffY - diffY // 2\n        ])\n\n\n# Instantiate\nmodel = UNet5Attention(in_channels=1, out_channels=1).to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T01:38:27.698970Z","iopub.execute_input":"2025-05-11T01:38:27.700394Z","iopub.status.idle":"2025-05-11T01:38:27.862198Z","shell.execute_reply.started":"2025-05-11T01:38:27.700348Z","shell.execute_reply":"2025-05-11T01:38:27.860821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.optim as optim\n\ncriterion = nn.L1Loss()  #\noptimizer = optim.Adam(model.parameters(), lr=5e-4)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T01:38:29.746434Z","iopub.execute_input":"2025-05-11T01:38:29.747444Z","iopub.status.idle":"2025-05-11T01:38:33.319924Z","shell.execute_reply.started":"2025-05-11T01:38:29.747408Z","shell.execute_reply":"2025-05-11T01:38:33.318987Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm.notebook import tqdm  # Use notebook-specific tqdm\nimport matplotlib.pyplot as plt\n\n# Simple training loop\nnum_epochs = 30\n\ntrain_losses = []\nval_losses = []\n\n# Training loop for Jupyter\nfor epoch in range(num_epochs):\n    model.train()\n    train_loss = 0\n    progress_bar = tqdm(train_loader, \n                        desc=f\"Epoch {epoch+1}/{num_epochs}\", \n                        total=len(train_loader), \n                        leave=False)\n\n    for batch_idx, (X_batch, y_batch) in enumerate(progress_bar):\n        X_batch, y_batch = X_batch.to(device), y_batch.to(device)\n        optimizer.zero_grad()\n        preds = model(X_batch)\n        loss = criterion(preds, y_batch)\n        loss.backward()\n        optimizer.step()\n\n        train_loss += loss.item()\n        progress_bar.set_postfix({\n            'loss': f\"{loss.item():.4f}\",\n            'avg_loss': f\"{train_loss/(batch_idx+1):.4f}\",\n            'lr': f\"{optimizer.param_groups[0]['lr']:.6f}\"\n        })\n        \n    train_loss /= len(train_loader)\n    train_losses.append(train_loss)  # Save train loss\n\n    # Validation\n    model.eval()\n    val_loss = 0\n    with torch.no_grad():\n        for X_batch, y_batch in val_loader:\n            X_batch, y_batch = X_batch.to(device), y_batch.to(device)\n            preds = model(X_batch)\n            loss = criterion(preds, y_batch)\n            val_loss += loss.item()\n\n    val_loss /= len(val_loader)\n    val_losses.append(val_loss)  # Save val loss\n\n    print(f\"Epoch {epoch+1}/{num_epochs} - Train Loss: {train_loss:.4f} - Val Loss: {val_loss:.4f}\")\n\n# 📈 Plot loss curves after training\nplt.figure(figsize=(8, 5))\nplt.plot(range(6, num_epochs + 1), train_losses[5:], label=\"Train Loss\")  # Start from epoch 3\nplt.plot(range(6, num_epochs + 1), val_losses[5:], label=\"Validation Loss\")  # Start from epoch 3\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training and Validation Loss Over Time (Starting from Epoch 3)\")\nplt.legend()\nplt.grid(True)\nplt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T01:38:35.620836Z","iopub.execute_input":"2025-05-11T01:38:35.621347Z","iopub.status.idle":"2025-05-11T03:40:57.451110Z","shell.execute_reply.started":"2025-05-11T01:38:35.621320Z","shell.execute_reply":"2025-05-11T03:40:57.449533Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Prediction vs Validation Plots\n","metadata":{}},{"cell_type":"code","source":"### NEW plot added \nimport matplotlib.pyplot as plt\nimport random\nimport torch\n\n# Pick random samples from validation set\nnum_samples = 5  # number of examples to visualize\nsample_indices = random.sample(range(len(X_val)), num_samples)\n\n# Plot predictions vs ground truth vs difference\nplt.figure(figsize=(18, 3 * num_samples))\n\nfor i, idx in enumerate(sample_indices):\n    X_sample = X_val[idx]  # shape (70, 70)\n    y_sample = y_val[idx]  # shape (70, 70)\n\n    X_sample_tensor = torch.tensor(X_sample, dtype=torch.float32).unsqueeze(0).unsqueeze(0).to(device)  # (1, 1, 70, 70)\n    y_sample = y_sample * (4500.0 - 1500.0) + 1500.0\n\n    with torch.no_grad():\n        pred = model(X_sample_tensor).squeeze().cpu().numpy()\n        pred = pred * (4500.0 - 1500.0) + 1500.0\n\n    diff = pred - y_sample  # Difference map\n\n    # Plot Ground Truth\n    plt.subplot(num_samples, 3, 3*i + 1)\n    plt.imshow(y_sample, cmap='viridis', aspect='equal', origin='lower', vmin=1500, vmax=4500)\n    plt.title(f\"Ground Truth #{idx}\")\n    plt.colorbar()\n\n    # Plot Prediction\n    plt.subplot(num_samples, 3, 3*i + 2)\n    plt.imshow(pred, cmap='viridis', aspect='equal', origin='lower', vmin=1500, vmax=4500)\n    plt.title(f\"Prediction #{idx}\")\n    plt.colorbar()\n\n    # Plot Difference\n    plt.subplot(num_samples, 3, 3*i + 3)\n    plt.imshow(diff, cmap='seismic', aspect='equal', origin='lower', vmin=-500, vmax=500)\n    plt.title(f\"Prediction - Ground Truth #{idx}\")\n    plt.colorbar(label=\"Error (m/s)\")\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T11:11:22.295250Z","iopub.execute_input":"2025-05-11T11:11:22.295552Z","iopub.status.idle":"2025-05-11T11:11:27.755538Z","shell.execute_reply.started":"2025-05-11T11:11:22.295528Z","shell.execute_reply":"2025-05-11T11:11:27.754079Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Prepare Submission File","metadata":{}},{"cell_type":"code","source":"\n# Set to evaluation mode\nmodel.eval()\n\n# Path to test files\ntest_dir = Path('/kaggle/input/waveform-inversion/test')\n\n# Initialize an empty list to collect prediction rows\nsubmission_rows = []\ncount = 0\ntest_examples = []\n\n# Go through test files one-by-one (no memory crash)\nfor file_path in test_dir.rglob('*.npy'):\n    # 1. Load test file\n    X = np.load(file_path)  # shape (5, 1000, 70)\n    count = count + 1\n    if count % 5000 == 0:\n        print(\"Loading file number:\", count)\n        test_examples.append(file_path)\n\n    # 2. Preprocess exactly like you did in training\n    X = X[:, 20:1000, :]\n    X = X[:, ::14, :]\n    X = X.transpose(1, 0, 2).reshape(70, -1)  # (70, 350)\n    X = X[:, ::5]  # (70, 70)\n\n    # 4. Convert to tensor\n    X_tensor = torch.tensor(X, dtype=torch.float32).unsqueeze(0).unsqueeze(0).to(device)  # (1, 1, 70, 70)\n\n    # 5. Predict\n    with torch.no_grad():\n        pred = model(X_tensor).squeeze(0).squeeze(0).cpu().numpy()  # (70, 70)\n        pred = pred * (4500.0 - 1500.0) + 1500.0\n    # 6. Format for submission\n    file_stem = file_path.stem  # e.g., \"2c56e45d6e\"\n    for i in range(pred.shape[0]):  # for each Y-slice\n        oid_ypos = f\"{file_stem}_y_{i}\"\n        x_values = pred[i, 1::2]  # take x_1, x_3, ..., x_69\n        row = [oid_ypos] + x_values.tolist()\n        submission_rows.append(row)\n    del X, X_tensor\n\n# After all files processed, make DataFrame\ncolumns = ['oid_ypos'] + [f'x_{i}' for i in range(1, 70, 2)]\nsubmission_df = pd.DataFrame(submission_rows, columns=columns)\n\n# Save CSV\nsubmission_df.to_csv('submission.csv', index=False)\nprint(\"✅ Submission file saved!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visualize test examples","metadata":{}},{"cell_type":"code","source":"\n\n# Put model in evaluation mode\nmodel.eval()\n\nfor i in range (len(test_examples)):\n    X = np.load(test_examples[i])  # shape (5, 1000, 70)\n    print(test_examples[i])\n    X = X[:, 20:1000, :]\n    X = X[:, ::14, :]\n    X = X.transpose(1, 0, 2).reshape(70, -1)  # (70, 350)\n    X = X[:, ::5]  # (70, 70)\n\n    X_tensor = torch.tensor(X, dtype=torch.float32).unsqueeze(0).unsqueeze(0).to(device)  # (1, 1, 70, 70)\n\n    # Predict\n    with torch.no_grad():\n        pred = model(X_tensor).squeeze(0).squeeze(0).cpu().numpy()  # (70, 70)\n        pred = pred * (4500.0 - 1500.0) + 1500.0\n\n    # Plot Prediction\n    plt.imshow(pred, cmap='viridis', aspect='equal', origin='lower', vmin=1500, vmax=4500)\n    plt.colorbar()\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T01:49:08.605319Z","iopub.execute_input":"2025-05-08T01:49:08.605655Z","iopub.status.idle":"2025-05-08T01:49:12.664390Z","shell.execute_reply.started":"2025-05-08T01:49:08.605632Z","shell.execute_reply":"2025-05-08T01:49:12.663391Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}