{"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"},{"sourceId":11801643,"sourceType":"datasetVersion","datasetId":7377931}],"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 pandas as pd \nimport os\nimport random\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport torch.nn as nn\nfrom tqdm import tqdm\nimport torch.nn.functional as F\n\nimport torch \nfrom torch.utils.data import DataLoader, Dataset, RandomSampler\nfrom torchvision import transforms\nfrom torchvision.transforms import Compose, ToTensor\nfrom sklearn.model_selection import train_test_split\nfrom pathlib import Path\nimport csv\n\ntorch.cuda.empty_cache()\n\nplt.style.use('seaborn-v0_8-whitegrid')\nOUTPUT_DIR = \"models/\"\nos.makedirs(OUTPUT_DIR, exist_ok =True)\n\nMODEL_PATH = \"models/model.pth\"\nTRAIN_DIRS = [\"/kaggle/input/openfwi-preprocessed-72x72/openfwi_72x72\"]\nBATCH_SIZE = 32 # TODO CHANGE LATER \nSEED = 42\nN_EXAMPLES_PER_FILE = 500\n\nINPUT_SHAPE = (1, 5, 72, 72)\nOUTPUT_SHAPE = (1, 70, 70)\nTEST_PATH =\"/kaggle/input/waveform-inversion/test\"\nTRAIN_RATIO = 0.8\n\nlr = 3e-4 \nweight_decay = 0.00001\nEPOCHS = 3\nEARLY_STOPPING_EPOCH = 3 # the number of epochs to wait for improvement in the model\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # Auto-detect GPU\nprint(\"Device:\", device)\ncuda_count = torch. cuda. device_count()\nprint(f\"We have {cuda_count} cuda devices\")\n\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:42.602804Z","iopub.execute_input":"2025-05-15T21:47:42.603088Z","iopub.status.idle":"2025-05-15T21:47:52.053471Z","shell.execute_reply.started":"2025-05-15T21:47:42.603059Z","shell.execute_reply":"2025-05-15T21:47:52.052759Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Resources :\n\n- https://www.kaggle.com/code/kayrahanozcan/yale-unc-ch-geophysical-waveform-inversion-v1-1\n- https://github.com/lanl/OpenFWI/","metadata":{}},{"cell_type":"markdown","source":"### Check one sample","metadata":{}},{"cell_type":"markdown","source":"Inside the training samples, we have families of syles, lets pick random one ","metadata":{}},{"cell_type":"code","source":"random_style = \"FlatFault_A\"\nrandom_style\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:52.054817Z","iopub.execute_input":"2025-05-15T21:47:52.055207Z","iopub.status.idle":"2025-05-15T21:47:52.061257Z","shell.execute_reply.started":"2025-05-15T21:47:52.055184Z","shell.execute_reply":"2025-05-15T21:47:52.060413Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"lets check the files of the style","metadata":{}},{"cell_type":"markdown","source":"The files that started with `seis` is the input, while `vel` is the output\n\nFor example `seis4_1_0.npy` correspond to `vel4_1_0.npy`.\n\nLets check it","metadata":{}},{"cell_type":"markdown","source":"Load the input first using `numpy`","metadata":{}},{"cell_type":"code","source":"loaded_example_input = np.load(os.path.join(os.path.join(TRAIN_DIRS[0], random_style), \"seis4_1_0.npy\"))\nloaded_example_output = np.load(os.path.join(os.path.join(TRAIN_DIRS[0], random_style), \"vel4_1_0.npy\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:52.062175Z","iopub.execute_input":"2025-05-15T21:47:52.062506Z","iopub.status.idle":"2025-05-15T21:47:52.441159Z","shell.execute_reply.started":"2025-05-15T21:47:52.062467Z","shell.execute_reply":"2025-05-15T21:47:52.440418Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Check the shape","metadata":{}},{"cell_type":"code","source":"loaded_example_input.shape, loaded_example_output.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:52.443212Z","iopub.execute_input":"2025-05-15T21:47:52.443539Z","iopub.status.idle":"2025-05-15T21:47:52.448783Z","shell.execute_reply.started":"2025-05-15T21:47:52.443519Z","shell.execute_reply":"2025-05-15T21:47:52.448033Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Each input and output contains 500 files . \n\n**Seismic data(inputs) (data/*.npy or seis_*.npy)** are 4D arrays of shape (batch_size, num_sources, time_steps, num_receivers) representing wave recordings over time.\n\n**Velocity maps (model/*.npy or vel_*.npy)** are 3D arrays of shape (batch_size, height, width) representing the subsurface velocity distribution.\n\nPick the first observation","metadata":{}},{"cell_type":"markdown","source":"The output is **2D(x, z)** where **x** is the offset and **z** is the depth of the wave in m(Metre) . Lets visualize the output .","metadata":{}},{"cell_type":"code","source":"sample_number = 14\nsample_y = loaded_example_output[sample_number]\nprint(\"Velocity shape:\", sample_y.shape)\n\nfig, ax = plt.subplots(1, figsize=(9, 6))\nim = ax.imshow(sample_y[0], cmap=\"jet\")\nax.set_xticks(range(0, 70, 10))\nax.set_xticklabels(range(0, 700, 100))\nax.set_yticks(range(0, 70, 10))\nax.set_yticklabels(range(0, 700, 100))\n\nax.set_xlabel(\"Offset\")\nax.set_ylabel(\"Depth\")\n\nplt.colorbar(im).ax.set_title(\"km/s\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:52.449450Z","iopub.execute_input":"2025-05-15T21:47:52.449651Z","iopub.status.idle":"2025-05-15T21:47:52.841315Z","shell.execute_reply.started":"2025-05-15T21:47:52.449636Z","shell.execute_reply":"2025-05-15T21:47:52.840587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_number = 14\n\nsample_x = loaded_example_input[sample_number]\nprint(\"Data shape:\", sample_x.shape)\nfig, axes = plt.subplots(1, 5, figsize=(20, 5))\n\nfor i, ax in enumerate(axes):\n   ax.imshow(sample_x[i], cmap=\"seismic\", aspect=\"auto\", vmin=-0.5,vmax=0.5)\n\n   ax.set_ylabel('Time (ms)', fontsize=16)\n   ax.set_xlabel('Signals', fontsize=16)\nplt.show()\n\nsample_y = loaded_example_output[sample_number]\nprint(\"Velocity shape:\", sample_y.shape)\n\nfig, ax = plt.subplots(1, figsize=(9, 6))\nim = ax.imshow(sample_y[0], cmap=\"jet\")\nax.set_xticks(range(0, 70, 10))\nax.set_xticklabels(range(0, 700, 100))\nax.set_yticks(range(0, 70, 10))\nax.set_yticklabels(range(0, 700, 100))\n\nax.set_xlabel(\"Offset\")\nax.set_ylabel(\"Depth\")\n\nplt.colorbar(im).ax.set_title(\"km/s\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:52.842055Z","iopub.execute_input":"2025-05-15T21:47:52.842318Z","iopub.status.idle":"2025-05-15T21:47:53.668209Z","shell.execute_reply.started":"2025-05-15T21:47:52.842269Z","shell.execute_reply":"2025-05-15T21:47:53.667607Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The input is 3D array(g, t, s) where \n\n**g** : is the geaphone\n\n**t** : is the time \n\n**s**: is the signals recorded by the geaphone","metadata":{}},{"cell_type":"markdown","source":"## Dataset Loading","metadata":{}},{"cell_type":"code","source":"class FWIDataset(Dataset):\n    def __init__(self, input_files, output_files , n_examples_per_files, data_transform=None, label_transform=None):\n        self.input_files = input_files\n        self.output_files = output_files\n        self.n_examples_per_files = n_examples_per_files\n        self.data_transform = data_transform\n        self.label_transform = label_transform\n        \n    def __getitem__(self, idx):\n        \n        # get the file index and the sample index\n        batch_idx, sample_idx = idx // self.n_examples_per_files, idx % self.n_examples_per_files\n        # check if it is exists\n        if batch_idx >= len(self.input_files):\n            raise IndexError(\"File doasn't exists\")\n        # load the files\n        data = np.load(self.input_files[batch_idx], mmap_mode='r')\n        label = np.load(self.output_files[batch_idx], mmap_mode='r')\n        # load the exact sample\n        seism_data = data[sample_idx].copy().astype('float32')\n        velocity_data = label[sample_idx].copy().astype('float32')\n\n        # remove the variables from the memory\n        del data, label\n        \n        if self.data_transform:\n           seism_data = self.data_transform(seism_data)\n        if self.label_transform:\n            velocity_data = self.label_transform(velocity_data)\n        return torch.from_numpy(seism_data) ,torch.from_numpy(velocity_data)\n\n    def __len__(self):\n        return self.n_examples_per_files * len(self.input_files)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:53.669004Z","iopub.execute_input":"2025-05-15T21:47:53.669217Z","iopub.status.idle":"2025-05-15T21:47:53.676889Z","shell.execute_reply.started":"2025-05-15T21:47:53.669199Z","shell.execute_reply":"2025-05-15T21:47:53.675928Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Normalize the data with Min Max","metadata":{}},{"cell_type":"code","source":"def log_transform(data, k=1, c=0):\n    return (np.log1p(np.abs(k * data) + c)) * np.sign(data)\n\nclass LogTransform(object):\n    def __init__(self, k=1, c=0):\n        self.k = k\n        self.c = c\n\n    def __call__(self, data):\n        return log_transform(data, k=self.k, c=self.c)\n \n\ndata_transform = Compose([\n    LogTransform(),\n])\n\nlabel_transform = Compose([\n    LogTransform(),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:53.677789Z","iopub.execute_input":"2025-05-15T21:47:53.678086Z","iopub.status.idle":"2025-05-15T21:47:53.695091Z","shell.execute_reply.started":"2025-05-15T21:47:53.678060Z","shell.execute_reply":"2025-05-15T21:47:53.694331Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"To load the data , we have first to work with two families given the limited resources we have.\nStart small end big","metadata":{}},{"cell_type":"code","source":"output_files = []\ninput_files = []\nfor train_dir in TRAIN_DIRS:\n    for dirname, _, filenames in os.walk(train_dir):\n        for filename in filenames:\n            path = os.path.join(dirname, filename)\n            if \"model\" in filename or \"vel\" in filename:\n                output_files.append(path)\n            elif \"csv\" not in filename:\n                input_files.append(path)\n\nprint(\"Length of input files: \", len(input_files), \" Length of output files: \", len(output_files))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:53.696334Z","iopub.execute_input":"2025-05-15T21:47:53.696599Z","iopub.status.idle":"2025-05-15T21:47:57.159651Z","shell.execute_reply.started":"2025-05-15T21:47:53.696575Z","shell.execute_reply":"2025-05-15T21:47:57.158917Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Split the data into train and val","metadata":{}},{"cell_type":"code","source":"train_input_files, val_input_files, train_output_files, val_output_files = train_test_split(input_files,  output_files, test_size=0.2, shuffle=True, random_state=42)\nprint(f\"Length of training: {len(train_input_files)}, Length of validation : {len(val_input_files)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:57.161812Z","iopub.execute_input":"2025-05-15T21:47:57.162036Z","iopub.status.idle":"2025-05-15T21:47:57.167994Z","shell.execute_reply.started":"2025-05-15T21:47:57.162019Z","shell.execute_reply":"2025-05-15T21:47:57.167202Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Loading the data","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"train_fwi_dataset = FWIDataset(train_input_files, train_output_files, N_EXAMPLES_PER_FILE, \n                         # data_transform=data_transform, \n                               # label_transform=label_transform\n                              )\n\nval_fwi_dataset = FWIDataset(val_input_files, val_output_files, N_EXAMPLES_PER_FILE, \n                         # data_transform=data_transform, \n                             #label_transform=label_transform\n                            )\nprint(f\"Length of training dataset : {len(train_fwi_dataset)}, Length of validation dataset : {len(val_fwi_dataset)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:57.169050Z","iopub.execute_input":"2025-05-15T21:47:57.169372Z","iopub.status.idle":"2025-05-15T21:47:57.198680Z","shell.execute_reply.started":"2025-05-15T21:47:57.169345Z","shell.execute_reply":"2025-05-15T21:47:57.197894Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_random_sampler = RandomSampler(train_fwi_dataset)\nval_random_sampler = RandomSampler(val_fwi_dataset)\n\ntrain_fwi_dataloader = DataLoader(train_fwi_dataset, batch_size=BATCH_SIZE, pin_memory=True, num_workers=4, sampler=train_random_sampler)\nval_fwi_dataloader = DataLoader(val_fwi_dataset, batch_size=BATCH_SIZE, pin_memory=True, num_workers=4, sampler=val_random_sampler)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:57.199465Z","iopub.execute_input":"2025-05-15T21:47:57.199671Z","iopub.status.idle":"2025-05-15T21:47:57.213019Z","shell.execute_reply.started":"2025-05-15T21:47:57.199655Z","shell.execute_reply":"2025-05-15T21:47:57.212447Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Building","metadata":{}},{"cell_type":"code","source":"class _BidirectionalLSTM(nn.Module):\n  def __init__(self, input_size: int, hidden_size: int, output_size: int):\n    super(_BidirectionalLSTM, self).__init__()\n    self.lstm = nn.LSTM(input_size=input_size, hidden_size=hidden_size, bidirectional=True, batch_first =True)\n    self.linear = nn.Linear(in_features = hidden_size* 2, out_features=output_size)\n  def forward(self, x: torch.Tensor)-> torch.Tensor:\n\n    recurrent, _ = self.lstm(x)\n    seq_lenght, batch_size, inputs_size = recurrent.size()\n    seq_lenght2 = recurrent.reshape(seq_lenght * batch_size, inputs_size)\n\n    out = self.linear(seq_lenght2)\n    out = out.reshape(seq_lenght, batch_size, -1)\n    return out\n\n# (batch_size, 5, 72, 72) => (batch_size, 70, 70)\nclass CRNN(nn.Module):\n  def __init__(self, in_channels: int, output_size:int):\n    super(CRNN, self).__init__()\n    self.conv_layer = nn.Sequential(\n        nn.Conv2d(in_channels = in_channels, out_channels=64, kernel_size=3, padding=1, stride=1, bias=True),\n        nn.ReLU(0.3),\n        nn.MaxPool2d(kernel_size=2, stride=2),\n        nn.Conv2d(in_channels=64, out_channels=128, kernel_size=3, padding=1, stride=1, bias=True),\n        nn.MaxPool2d(kernel_size=2, stride=2),\n        nn.Conv2d(in_channels=128, out_channels=256, kernel_size=3, padding=1, stride=1, bias=False),\n        nn.BatchNorm2d(256),\n        nn.ReLU(0.3),\n        nn.Conv2d(in_channels=256, out_channels=256, kernel_size=3, padding=1, stride=1, bias=True),\n        nn.MaxPool2d(kernel_size=(2, 2), stride=(2, 1), padding=(0, 1)),\n        nn.Conv2d(in_channels=256, out_channels=512, kernel_size=3, padding=1, stride=1, bias=False),\n        nn.BatchNorm2d(512),\n        # nn.ReLU(0.2),\n        nn.Conv2d(in_channels=512, out_channels=512, kernel_size=3, padding=1, stride=1, bias=True),\n        # nn.ReLU(0.2),\n        nn.MaxPool2d(kernel_size=(2, 2), stride=(2, 1), padding=(0, 1)),\n        nn.Conv2d(in_channels=512, out_channels=512, kernel_size=2, padding=0, stride=1, bias=False),\n        nn.BatchNorm2d(512),\n        # nn.ReLU(0.2),\n\n    )\n    self.output_size = output_size\n    self.recurrent_layer = nn.Sequential(\n        _BidirectionalLSTM(512, 256, 256),\n        _BidirectionalLSTM(256, 256,  256),\n    )\n    self.fcl =  nn.Sequential(\n        nn.Linear(in_features = 14592, out_features=516) ,\n        nn.Linear(in_features = 516, out_features=self.output_size*self.output_size) ,\n        \n    )   \n  def forward(self, x: torch.Tensor) -> torch.Tensor:\n      batch_size = x.shape[0]\n      x = x.float()\n      mean = torch.mean(x, dim=(2, 3), keepdim=True)\n      std = torch.std(x, dim=(2, 3), keepdim=True)\n      x_norm = (x - mean) / (std + 1e-8) # Epsilon for numerical stability\n      \n      features = self.conv_layer(x_norm) # squeze the first [32, 512, 3, 19]\n\n      features = features.reshape(batch_size, 512, 57) # torch.Size([batch_size, 512, 57])\n      features = features.permute(0, 2, 1) # chagnge the dimension\n      recurrent = self.recurrent_layer(features) \n      flattened = recurrent.reshape(batch_size, -1) # flatten\n      linear_features = self.fcl(flattened)\n      output = linear_features.reshape(batch_size, 1, self.output_size, self.output_size) # reshape to (batch_size, 1, output_size, output_size)\n      output = output * 1000.0 + 1500.0\n      return output ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:57.213816Z","iopub.execute_input":"2025-05-15T21:47:57.214336Z","iopub.status.idle":"2025-05-15T21:47:57.230109Z","shell.execute_reply.started":"2025-05-15T21:47:57.214310Z","shell.execute_reply":"2025-05-15T21:47:57.229341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, dataloader, loss_fn, optimizer, device):\n    model.to(device)\n    model.train()\n    total_loss = 0\n    progress = tqdm(dataloader, desc=\"Training Epoch\")\n    for data, labels in progress:\n        optimizer.zero_grad()\n        data = data.to(device).float()\n        labels = labels.to(device).float()\n\n        with torch.autocast(device_type=str(device)):\n            outputs = model(data)\n            loss = loss_fn(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n        progress.set_postfix(loss=f'{loss.item():3f}')\n\n    return total_loss / len(dataloader)\n\n\ndef val(model, loss_fn, dataloader, device):\n    model.eval()\n    total_loss = 0\n    progress = tqdm(dataloader, desc=\"Validation\")\n    i = 0 # index to tracking storing the validation output of the first output\n    first_output = None\n    first_label = None\n    with torch.no_grad():\n        for data, labels in progress:\n            data = data.to(device).float()\n            labels = labels.to(device).float()\n\n            outputs = model(data)\n            loss = loss_fn(outputs, labels)\n\n            total_loss += loss.item()\n            progress.set_postfix(loss=f'{loss.item():3f}')\n            if i== 0:\n                first_output = outputs[0].detach().cpu()\n                first_label = labels[0].detach().cpu()\n            \n    return total_loss / len(dataloader), first_output, first_label\n\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:57.230887Z","iopub.execute_input":"2025-05-15T21:47:57.231086Z","iopub.status.idle":"2025-05-15T21:47:57.243798Z","shell.execute_reply.started":"2025-05-15T21:47:57.231072Z","shell.execute_reply":"2025-05-15T21:47:57.243196Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":" torch.cuda.empty_cache()\ngpu_devices = ','.join([str(id) for id in range(0, cuda_count)])\nos.environ[\"CUDA_VISIBLE_DEVICES\"] = gpu_devices\n# define the model\nnetwork = CRNN(5, 70)\n# load the model into all cuda available (in our case 2)\nnetowrk = nn.DataParallel(network)\nnetwork.to(device)\n\n# Apply initialization\n# define the Adam optmizer \noptimizer = torch.optim.AdamW( network.parameters(), lr=lr, weight_decay=weight_decay)\n# scheduler \nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer)\n# define the loss\nloss_fn = nn.L1Loss()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:57.244501Z","iopub.execute_input":"2025-05-15T21:47:57.244710Z","iopub.status.idle":"2025-05-15T21:47:57.754087Z","shell.execute_reply.started":"2025-05-15T21:47:57.244686Z","shell.execute_reply":"2025-05-15T21:47:57.752783Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training the model","metadata":{}},{"cell_type":"code","source":"def plot_output(ground_truth, ground_output):\n    values = [ground_truth, ground_output]\n    titles = [\"Prediction\", \"Ground Output\"]\n    fig, axes = plt.subplots(1, 2, figsize=(9, 6))\n    for i in range(0, 2):\n        im = axes[i].imshow(values[i], cmap=\"jet\")\n        axes[i].set_xticks(range(0, 70, 10))\n        axes[i].set_xticklabels(range(0, 700, 100))\n        axes[i].set_yticks(range(0, 70, 10))\n        axes[i].set_yticklabels(range(0, 700, 100))\n        \n        axes[i].set_xlabel(\"Offset\")\n        axes[i].set_ylabel(\"Depth\")\n        axes[i].set_title(titles[i])\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:47:57.754934Z","iopub.execute_input":"2025-05-15T21:47:57.755208Z","iopub.status.idle":"2025-05-15T21:47:57.761143Z","shell.execute_reply.started":"2025-05-15T21:47:57.755182Z","shell.execute_reply":"2025-05-15T21:47:57.760351Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":" torch.cuda.empty_cache()\nlosses = {\n    \"val_loss\":[],\n    \"train_loss\":[]\n}\nbest_loss = 10000\nepoch_waited = 0\nfor epoch in range(0, 10):\n    torch.cuda.empty_cache()\n    print(f\"Training for {epoch} epoch\")\n    avg_loss = train_one_epoch(network, train_fwi_dataloader, loss_fn, optimizer, device)\n    avg_loss_val, val_output, val_label = val(network,loss_fn, val_fwi_dataloader , device)\n    scheduler.step(avg_loss_val)\n    # Printing the training and validation loss\n    print(f\"==Training loss:{avg_loss} Validation loss:{avg_loss_val}===\")\n    # Ploting every 5 epoch the results\n    if epoch % 5 == 0:\n        plot_output(val_output[0],val_label[0])\n    # Append the loss for later ploting\n    losses[\"train_loss\"].append([avg_loss, i])\n    losses[\"val_loss\"].append([avg_loss_val, i])\n    if best_loss > avg_loss_val:\n        best_loss = avg_loss_val\n        print(f\"Saving the best model to {MODEL_PATH}\")\n        torch.save(network.state_dict().copy(), MODEL_PATH)\n    else:\n        epoch_waited += 1\n    # break if no improvement happen\n    if epoch_waited >= EARLY_STOPPING_EPOCH:\n        print(\"Breaking... No improvement\")\n        break","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"network.load_state_dict(torch.load(MODEL_PATH, weights_only=True))\nnetwork.eval()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"avg_loss_val, val_output, val_label = val(network,loss_fn, val_fwi_dataloader , device)\nprint(\"Loss:\", avg_loss_val)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport numpy as np\n\ndef _preprocess(x):\n    x = F.interpolate(x, size=(70, 70), mode='area')\n    x = F.pad(x, (1,1,1,1), mode='replicate')\n    return x\n\ndef _helper(x, ):\n    before_shape = x.shape\n    before_mem = x.nbytes / 1e6\n    x = torch.from_numpy(x).float()\n\n    # Interpolate and pad\n    x = _preprocess(x)\n    x = x.reshape(5, 72, 72)\n    return x\n    \nclass TestDataset(Dataset):\n    def __init__(self, files):\n        self.files = files\n\n\n    def __len__(self):\n        return len(self.files)\n\n\n    def __getitem__(self, i):\n        test_file = self.files[i]\n        x = _helper(np.load(test_file).reshape(-1, 5, 1000, 70))\n\n        return x, test_file.stem","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:48:04.172425Z","iopub.execute_input":"2025-05-15T21:48:04.172692Z","iopub.status.idle":"2025-05-15T21:48:04.179529Z","shell.execute_reply.started":"2025-05-15T21:48:04.172673Z","shell.execute_reply":"2025-05-15T21:48:04.178564Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_files = list(Path(TEST_PATH).glob(\"*.npy\"))\ntest_dataset = TestDataset(test_files)\ntest_dataloader = DataLoader(test_dataset, batch_size=BATCH_SIZE, pin_memory=False, num_workers=4)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:48:04.340848Z","iopub.execute_input":"2025-05-15T21:48:04.341126Z","iopub.status.idle":"2025-05-15T21:48:05.189589Z","shell.execute_reply.started":"2025-05-15T21:48:04.341104Z","shell.execute_reply":"2025-05-15T21:48:05.188981Z"}},"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\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(test_dataloader, desc=\"Testing data\"):\n        inputs = inputs.to(device)\n        with torch.inference_mode():\n            with torch.autocast(device_type=\"cuda\"):\n                outputs = network(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(zip(x_cols, [y_pred[y_pos, x_pos] for x_pos in range(1, 70, 2)]))\n                row[\"oid_ypos\"] = f\"{oid_test}_y_{y_pos}\"\n\n                writer.writerow(row)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:48:05.673777Z","iopub.execute_input":"2025-05-15T21:48:05.674052Z","iopub.status.idle":"2025-05-15T21:48:15.741675Z","shell.execute_reply.started":"2025-05-15T21:48:05.674031Z","shell.execute_reply":"2025-05-15T21:48:15.740574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}