{"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":"gpu","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom pathlib import Path\nfrom torch.utils.data import Dataset, DataLoader\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport csv\nfrom tqdm import tqdm\nfrom skimage.metrics import structural_similarity as ssim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom scipy.stats import pearsonr\nfrom scipy.ndimage import rotate","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-26T22:50:34.471478Z","iopub.execute_input":"2025-04-26T22:50:34.471762Z","iopub.status.idle":"2025-04-26T22:50:34.476421Z","shell.execute_reply.started":"2025-04-26T22:50:34.471739Z","shell.execute_reply":"2025-04-26T22:50:34.475765Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"velocity = np.load('/kaggle/input/waveform-inversion/train_samples/FlatVel_A/model/model2.npy')\nseismic = np.load('/kaggle/input/waveform-inversion/train_samples/FlatVel_A/data/data2.npy')\n\nprint(\"Velocity map shape:\", velocity.shape)  \nprint(\"Seismic data shape:\", seismic.shape)   ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T22:50:34.477527Z","iopub.execute_input":"2025-04-26T22:50:34.477732Z","iopub.status.idle":"2025-04-26T22:50:40.046121Z","shell.execute_reply.started":"2025-04-26T22:50:34.477715Z","shell.execute_reply":"2025-04-26T22:50:40.045399Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_shot_gather(data, title):\n    fig, ax = plt.subplots(figsize=(12, 6))\n    norm = np.max(np.abs(data))\n    for i in range(data.shape[1]):\n        trace = data[:, i] / norm  \n        ax.plot(trace + i, color='black') \n    ax.set_title(title)\n    ax.set_xlabel(\"Receiver Index (shifted)\")\n    ax.set_ylabel(\"Time\")\n    ax.invert_yaxis()\n    plt.show()\n\n\ndef plot_validation_results(epoch, outputs, targets, valid_losses):\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, cmap='viridis')\n        ax[1].imshow(y_pred, cmap='viridis')\n        plt.show()  \n\n\ndef plot_seismic_waveform(seismic, sample, source_idx, receiver_idx):\n    waveform = seismic[sample, source_idx, :, receiver_idx]  \n\n    plt.figure(figsize=(12, 4))\n    plt.plot(waveform, lw=0.8)\n    plt.title(f\"Seismic Waveform - Sample {sample}, Source {source_idx}, Receiver {receiver_idx}\")\n    plt.xlabel(\"Time step\")\n    plt.ylabel(\"Amplitude\")\n    plt.grid(True)\n    plt.show()\n\n\ndef plot_velocity_map(velocity, sample):\n    fig, ax = plt.subplots(figsize=(10, 6))\n    img = ax.imshow(velocity[sample, 0], cmap='jet', origin='upper')\n    plt.colorbar(img, ax=ax, label=\"Velocity (km/s)\")\n    ax.set_title(f\"Velocity Map - Sample {sample}\")\n    ax.set_xlabel(\"Horizontal Position (x)\")\n    ax.set_ylabel(\"Depth (z)\")\n    plt.show()\n\n\ndef plot_seismic_data(seis, sample_id):\n    plt.figure(figsize=(10, 6))\n    plt.title(f\"Seismic Data - Batch 0, Source 0\")\n    plt.imshow(seis[0, 0], aspect='auto', cmap='seismic')\n    plt.colorbar(label=\"Amplitude\")\n    plt.xlabel(\"Receivers\")\n    plt.ylabel(\"Timesteps\")\n    plt.show()\n\n\nsample = 11\nplot_shot_gather(seismic[sample, 0], f\"Shot Gather - Sample {sample}, Source 0\")\nplot_seismic_waveform(seismic, sample, 0, 35)\nplot_velocity_map(velocity, sample)\nplot_seismic_data(seismic, sample)  \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T22:50:40.046843Z","iopub.execute_input":"2025-04-26T22:50:40.047124Z","iopub.status.idle":"2025-04-26T22:50:41.009760Z","shell.execute_reply.started":"2025-04-26T22:50:40.047104Z","shell.execute_reply":"2025-04-26T22:50:41.009026Z"}},"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_inputs = [\n    f\n    for f in Path('/kaggle/input/waveform-inversion/train_samples').rglob('*.npy')\n    if ('seis' in f.stem) or ('data' in f.stem)\n]\n\nall_outputs = inputs_files_to_output_files(all_inputs)\nassert all(f.exists() for f in all_outputs)\n\ntrain_inputs = [all_inputs[i] for i in range(0, len(all_inputs), 2)]  \nvalid_inputs = [f for f in all_inputs if f not in train_inputs]\ntrain_outputs = inputs_files_to_output_files(train_inputs)\nvalid_outputs = inputs_files_to_output_files(valid_inputs)\n\n\nclass SeismicDataset(Dataset):\n    def __init__(self, inputs_files, output_files, n_examples_per_file=500, augmentation_prob=0.3, crop_size=None, rotate_prob=0.5, flip_prob=0.5, normalize=True):\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        self.augmentation_prob = augmentation_prob\n        self.crop_size = crop_size\n        self.rotate_prob = rotate_prob\n        self.flip_prob = flip_prob\n        self.normalize = normalize  \n\n    def __len__(self):\n        return len(self.inputs_files) * self.n_examples_per_file\n\n    def __getitem__(self, idx):\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        X_sample = X[sample_idx].copy()\n\n        if np.random.rand() < self.augmentation_prob:\n            noise = np.random.normal(0, 0.01, X_sample.shape)\n            X_sample += noise\n        \n        if self.crop_size is not None:\n            X_sample = self.random_crop(X_sample, self.crop_size)\n\n        if np.random.rand() < self.rotate_prob:\n            X_sample = self.random_rotate(X_sample)\n\n        if np.random.rand() < self.flip_prob:\n            X_sample = self.random_flip(X_sample)\n\n        if self.normalize:\n            X_sample = (X_sample - np.mean(X_sample)) / np.std(X_sample)  \n\n        return X_sample.copy(), y[sample_idx].copy()\n\n    def random_crop(self, X_sample, crop_size):\n        h, w = X_sample.shape\n        crop_h, crop_w = crop_size\n        top = np.random.randint(0, h - crop_h)\n        left = np.random.randint(0, w - crop_w)\n        cropped = X_sample[top:top+crop_h, left:left+crop_w]\n        return cropped\n\n    def random_rotate(self, X_sample):\n        angle = np.random.uniform(-45, 45)\n        rotated = rotate(X_sample, angle, mode='nearest', reshape=False)  \n        return rotated\n\n    def random_flip(self, X_sample):\n        flip_choice = np.random.choice([0, 1])  \n        if flip_choice == 0:\n            X_sample = np.fliplr(X_sample)  \n        else:\n            X_sample = np.flipud(X_sample)  \n        return X_sample\n\n\ndstrain = SeismicDataset(train_inputs, train_outputs)\ndltrain = DataLoader(dstrain, batch_size=64, 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=64, shuffle=False, pin_memory=True, drop_last=False, num_workers=4, persistent_workers=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T22:50:41.010821Z","iopub.execute_input":"2025-04-26T22:50:41.011207Z","iopub.status.idle":"2025-04-26T22:50:41.214952Z","shell.execute_reply.started":"2025-04-26T22:50:41.011185Z","shell.execute_reply":"2025-04-26T22:50:41.214171Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ChannelAttention(nn.Module):\n    def __init__(self, in_planes, ratio=8):\n        super(ChannelAttention, self).__init__()\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.max_pool = nn.AdaptiveMaxPool2d(1)\n\n        self.fc = nn.Sequential(\n            nn.Conv2d(in_planes, in_planes // ratio, 1, bias=False),\n            nn.ReLU(),\n            nn.Conv2d(in_planes // ratio, in_planes, 1, bias=False)\n        )\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg_out = self.fc(self.avg_pool(x))\n        max_out = self.fc(self.max_pool(x))\n        out = avg_out + max_out\n        return self.sigmoid(out)\n\n\nclass SpatialAttention(nn.Module):\n    def __init__(self, kernel_size=7):\n        super(SpatialAttention, self).__init__()\n        assert kernel_size in (3, 7), 'kernel_size must be 3 or 7'\n        padding = 3 if kernel_size == 7 else 1\n\n        self.conv = nn.Conv2d(2, 1, kernel_size, padding=padding, bias=False)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg_out = torch.mean(x, dim=1, keepdim=True)\n        max_out, _ = torch.max(x, dim=1, keepdim=True)\n        x = torch.cat([avg_out, max_out], dim=1)\n        x = self.conv(x)\n        return self.sigmoid(x)\n\n\nclass SeismicModel(nn.Module):\n    def __init__(self):\n        super(SeismicModel, self).__init__()\n        \n        self.conv_block1 = nn.Sequential(\n            nn.Conv2d(5, 16, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(16),\n            nn.ReLU(),\n            nn.Dropout2d(0.2)\n        )\n        \n        self.conv_block2 = nn.Sequential(\n            nn.Conv2d(16, 32, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            nn.MaxPool2d(2),\n            nn.Dropout2d(0.2)\n        )\n        \n        self.conv_block3 = nn.Sequential(\n            nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Dropout2d(0.2)\n        )\n        \n        self.conv_block4 = nn.Sequential(\n            nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.MaxPool2d(2),\n            nn.Dropout2d(0.2)\n        )\n\n        self.ca = ChannelAttention(128)\n        self.sa = SpatialAttention()\n\n        self.avgpool = nn.AdaptiveAvgPool2d((8, 8))\n\n        self.fc1 = nn.Linear(128 * 8 * 8, 256)\n        self.fc2 = nn.Linear(256, 70 * 70)\n\n    def forward(self, x):\n        x = self.conv_block1(x)\n        x = self.conv_block2(x)\n        x = self.conv_block3(x)\n        x = self.conv_block4(x)\n\n        x = self.ca(x) * x\n        x = self.sa(x) * x\n\n        x = self.avgpool(x)  \n        x = x.view(x.size(0), -1)  \n\n        x = F.relu(self.fc1(x))\n        x = self.fc2(x)\n\n        x = x.view(x.size(0), 1, 70, 70)\n        return x\n\n\nclass EarlyStopping:\n    def __init__(self, patience=5, delta=0):\n        self.patience = patience  \n        self.delta = delta        \n        self.counter = 0          \n        self.best_loss = None     \n        self.early_stop = False   \n\n    def __call__(self, val_loss):\n        if self.best_loss is None:\n            self.best_loss = val_loss\n        elif val_loss < self.best_loss - self.delta:\n            self.best_loss = val_loss\n            self.counter = 0\n        else:\n            self.counter += 1\n            if self.counter >= self.patience:\n                self.early_stop = True\n        return self.early_stop\n\n\ndef train(model, train_dataloader, valid_dataloader, epochs=10, patience=5):\n    model.to(device)\n    optimizer = torch.optim.Adam(model.parameters())\n    criterion = torch.nn.L1Loss()\n\n    scheduler = ReduceLROnPlateau(optimizer, 'min', patience=5, factor=0.5, verbose=True)\n    early_stopping = EarlyStopping(patience=patience)\n    \n    train_losses = []\n    valid_losses = []\n    ssim_scores = []\n    corr_scores = []\n    learning_rates = [] \n\n    for epoch in range(epochs):\n        model.train()\n        running_loss = 0.0\n\n        for inputs, targets in tqdm(train_dataloader, desc=f'Training Epoch {epoch+1}'):\n            inputs, targets = inputs.to(device), targets.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(inputs)\n\n            loss = criterion(outputs, targets)\n            loss.backward()\n            optimizer.step()\n\n            running_loss += loss.item()\n\n        avg_train_loss = running_loss / len(train_dataloader)\n        train_losses.append(avg_train_loss)\n\n        model.eval()\n        valid_loss = 0.0\n        with torch.no_grad():\n            for inputs, targets in tqdm(valid_dataloader, desc=f'Validating Epoch {epoch+1}'):\n                inputs, targets = inputs.to(device), targets.to(device)\n                outputs = model(inputs)\n                loss = criterion(outputs, targets)\n                valid_loss += loss.item()\n\n        avg_valid_loss = valid_loss / len(valid_dataloader)\n        valid_losses.append(avg_valid_loss)\n\n        target_np = targets[0, 0].detach().cpu().numpy()\n        output_np = outputs[0, 0].detach().cpu().numpy()\n        ssim_score = ssim(target_np, output_np, data_range=output_np.max() - output_np.min())\n        corr_score, _ = pearsonr(target_np.flatten(), output_np.flatten())\n\n        current_lr = optimizer.param_groups[0]['lr']\n        learning_rates.append(current_lr)\n        \n        ssim_scores.append(ssim_score)\n        corr_scores.append(corr_score)\n\n        print(f\"Epoch {epoch+1}: Train Loss: {avg_train_loss:.5f}, Valid Loss: {avg_valid_loss:.5f}\")\n        print(f\"          SSIM: {ssim_score:.4f}, Pearson Corr: {corr_score:.4f}\")\n\n        scheduler.step(avg_valid_loss)\n        \n        if early_stopping(avg_valid_loss):\n            print(\"Early stopping triggered.\")\n            break\n            \n        plot_validation_results(epoch, outputs, targets, valid_losses)\n\n    plt.figure(figsize=(16,5))\n\n    plt.subplot(1,4,1)\n    plt.plot(range(1, epochs+1), train_losses, label='Train Loss')\n    plt.plot(range(1, epochs+1), valid_losses, label='Valid Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.title('Train/Validation Loss')\n    plt.legend()\n\n    plt.subplot(1,4,2)\n    plt.plot(range(1, epochs+1), ssim_scores, marker='o')\n    plt.xlabel('Epoch')\n    plt.ylabel('SSIM')\n    plt.title('SSIM over Epochs')\n\n    plt.subplot(1,4,3)\n    plt.plot(range(1, epochs+1), corr_scores, marker='o')\n    plt.xlabel('Epoch')\n    plt.ylabel('Pearson Correlation')\n    plt.title('Pearson Corr over Epochs')\n\n    plt.subplot(1,4,4)\n    plt.plot(range(1, epochs+1), learning_rates, marker='o')\n    plt.xlabel('Epoch')\n    plt.ylabel('Learning Rate')\n    plt.title('Learning Rate over Epochs')\n\n    plt.tight_layout()\n    plt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T22:50:41.216739Z","iopub.execute_input":"2025-04-26T22:50:41.217048Z","iopub.status.idle":"2025-04-26T22:50:41.239382Z","shell.execute_reply.started":"2025-04-26T22:50:41.217019Z","shell.execute_reply":"2025-04-26T22:50:41.238654Z"}},"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    def __len__(self):\n        return len(self.test_files)\n\n    def __getitem__(self, i):\n        test_file = self.test_files[i]\n        return np.load(test_file), test_file.stem\n\nx_cols = [f'x_{i}' for i in range(1, 70, 2)]\nfieldnames = ['oid_ypos'] + x_cols\n\ntest_files = [f for f in Path('/kaggle/input/waveform-inversion/test').rglob('*.npy')]\n\nds_test = TestDataset(test_files)\ndl_test = DataLoader(ds_test, batch_size=8, num_workers=4, pin_memory=True)\n\ndef test_and_save_results(model, test_dataloader):\n    model.eval()\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n    with 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(test_dataloader, desc='Testing'):\n            inputs = inputs.to(device)\n            with torch.no_grad():\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)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nmodel = SeismicModel()\ntrain(model, dltrain, dlvalid, epochs=50)  \ntest_and_save_results(model, dl_test)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T22:50:41.240131Z","iopub.execute_input":"2025-04-26T22:50:41.240318Z"}},"outputs":[],"execution_count":null}]}