{"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":11460962,"sourceType":"datasetVersion","datasetId":7179812}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-18T11:45:52.552198Z","iopub.execute_input":"2025-04-18T11:45:52.552428Z","iopub.status.idle":"2025-04-18T11:45:59.761345Z","shell.execute_reply.started":"2025-04-18T11:45:52.552408Z","shell.execute_reply":"2025-04-18T11:45:59.760501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T09:29:51.546358Z","iopub.execute_input":"2025-04-18T09:29:51.546978Z","iopub.status.idle":"2025-04-18T09:29:51.565839Z","shell.execute_reply.started":"2025-04-18T09:29:51.546952Z","shell.execute_reply":"2025-04-18T09:29:51.564938Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"headers = ['oid_ypos'] + [f\"x_{xpos}\" for xpos in range(1, 70, 2)]\nheaders = ','.join(headers)\nwith open(\"submission.csv\", 'w+') as file:\n    file.write(headers)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T09:29:51.566685Z","iopub.execute_input":"2025-04-18T09:29:51.566997Z","iopub.status.idle":"2025-04-18T09:29:51.588191Z","shell.execute_reply.started":"2025-04-18T09:29:51.566969Z","shell.execute_reply":"2025-04-18T09:29:51.587318Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T09:29:51.590015Z","iopub.execute_input":"2025-04-18T09:29:51.590305Z","iopub.status.idle":"2025-04-18T09:29:51.608362Z","shell.execute_reply.started":"2025-04-18T09:29:51.590284Z","shell.execute_reply":"2025-04-18T09:29:51.607414Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch_size = 100\n\nTEST_DIR = \"/kaggle/input/waveform-inversion/test\"\ntest_filenames = [filename for filename in os.listdir(TEST_DIR)]\nprint(\"number of test files:\", len(test_filenames))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T09:29:51.609179Z","iopub.execute_input":"2025-04-18T09:29:51.609491Z","iopub.status.idle":"2025-04-18T09:29:51.644909Z","shell.execute_reply.started":"2025-04-18T09:29:51.609468Z","shell.execute_reply":"2025-04-18T09:29:51.644050Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Model","metadata":{}},{"cell_type":"code","source":"class FWIEncoder(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.pre_mlp = nn.Sequential(\n            nn.Linear(1000, 512),\n            nn.Dropout(0.1),\n            nn.LeakyReLU()\n        )\n        self.positional_embedding = nn.Embedding(70*5, 512)\n        self.transformer_layer1 = nn.TransformerEncoderLayer(\n            512,\n            8,\n            dim_feedforward=512,\n            batch_first=True\n        )\n        self.tanh = nn.Tanh()\n    def forward(self, x):\n        x = x.view(x.shape[0], -1, x.shape[-1])\n        xpos = torch.arange(5*70, dtype=torch.long, device=device).view(1, -1)\n        xpos = self.positional_embedding(xpos)\n        \n        x = self.pre_mlp(x) + xpos\n        x = self.transformer_layer1(x)\n        \n        x = x.mean(dim=1)\n        x = self.tanh(x)\n        return x\n\nclass FWIModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = FWIEncoder()\n        self.decoder = nn.Sequential(\n            nn.Linear(512, 1024),\n            nn.BatchNorm1d(1024),\n            nn.Softplus(),\n            nn.Dropout(0.1),\n\n            nn.Linear(1024, 1024),\n            nn.BatchNorm1d(1024),\n            nn.Softplus(),\n            nn.Dropout(0.1),\n\n            nn.Linear(1024, 70*70)\n        )\n    def forward(self, x):\n        xencoder = self.encoder(x)\n        xdecoder = self.decoder(xencoder)\n        xdecoder = xdecoder.view(-1, 70, 70)\n        return xdecoder","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T09:29:51.645949Z","iopub.execute_input":"2025-04-18T09:29:51.646247Z","iopub.status.idle":"2025-04-18T09:29:51.657581Z","shell.execute_reply.started":"2025-04-18T09:29:51.646219Z","shell.execute_reply":"2025-04-18T09:29:51.656543Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### load models","metadata":{}},{"cell_type":"code","source":"models = [\n    torch.load(\"/kaggle/input/fwi-baseline-models/model0.pth\", map_location=device),\n    #torch.load(\"/kaggle/input/fwi-baseline-models/model1.pth\", map_location=device),\n    #torch.load(\"/kaggle/input/fwi-baseline-models/model2.pth\", map_location=device),\n    #torch.load(\"/kaggle/input/fwi-baseline-models/model3.pth\", map_location=device),\n    #torch.load(\"/kaggle/input/fwi-baseline-models/model4.pth\", map_location=device)\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T09:29:51.658553Z","iopub.execute_input":"2025-04-18T09:29:51.658843Z","iopub.status.idle":"2025-04-18T09:29:52.866796Z","shell.execute_reply.started":"2025-04-18T09:29:51.658821Z","shell.execute_reply":"2025-04-18T09:29:52.865764Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Inference","metadata":{}},{"cell_type":"code","source":"def load_testdata(subfilenames):\n    test_data = [np.load(os.path.join(TEST_DIR, filename))[np.newaxis, :] for filename in subfilenames]\n    test_data = np.concatenate(test_data)\n    test_data = torch.tensor(test_data, dtype=torch.float32)\n    test_data = test_data/10\n    test_data = test_data.to(device)\n    test_data = test_data.permute(0,1, 3, 2)\n    test_data = test_data.contiguous()\n    return test_data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T09:29:52.868188Z","iopub.execute_input":"2025-04-18T09:29:52.868512Z","iopub.status.idle":"2025-04-18T09:29:52.873741Z","shell.execute_reply.started":"2025-04-18T09:29:52.868482Z","shell.execute_reply":"2025-04-18T09:29:52.872754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def infer_model(data):\n    batch_size=len(data)\n    pred_velocity = np.zeros((batch_size, 70, 70), dtype=np.float32)\n\n    for model in models:\n        model.eval()\n        with torch.no_grad():\n            pred_velocity += model(data).detach().cpu().numpy()\n    pred_velocity = pred_velocity/len(models)\n    pred_velocity = (pred_velocity*1000)+3000\n    return pred_velocity","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T09:29:52.874804Z","iopub.execute_input":"2025-04-18T09:29:52.875395Z","iopub.status.idle":"2025-04-18T09:29:52.895958Z","shell.execute_reply.started":"2025-04-18T09:29:52.875362Z","shell.execute_reply":"2025-04-18T09:29:52.894953Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def append_results_to_file(filenames, pred_velocity_datalist):\n    for k,filename in enumerate(filenames):\n        oid = filename.replace(\".npy\", \"\")\n        yhat = pred_velocity_datalist[k]\n        \n        for y_pos in range(70):\n            oidpos = oid+\"_y_\"+str(y_pos)\n            cur_data=[]\n            cur_data.append(oidpos)\n            for x_pos in range(1, 70, 2):\n                pred_value = yhat[y_pos][x_pos]\n                pred_value = str(pred_value)\n                cur_data.append( pred_value )\n            cur_data = ','.join(cur_data)\n            with open(\"submission.csv\", 'a') as file:\n                file.write(\"\\n\")\n                file.writelines(cur_data)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T09:29:52.898473Z","iopub.execute_input":"2025-04-18T09:29:52.898769Z","iopub.status.idle":"2025-04-18T09:29:52.912931Z","shell.execute_reply.started":"2025-04-18T09:29:52.898745Z","shell.execute_reply":"2025-04-18T09:29:52.911944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\nfor k in range(0, len(test_filenames), batch_size):\n    if (k+batch_size)%500 == 0:\n        print(k , k+batch_size)\n    subfilenames = test_filenames[k: k+batch_size]\n    test_data = load_testdata(subfilenames)\n    pred_velocity_datalist = infer_model(test_data)\n    append_results_to_file(subfilenames, pred_velocity_datalist)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T09:29:52.913921Z","iopub.execute_input":"2025-04-18T09:29:52.914241Z","iopub.status.idle":"2025-04-18T09:30:05.620543Z","shell.execute_reply.started":"2025-04-18T09:29:52.914209Z","shell.execute_reply":"2025-04-18T09:30:05.619593Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}