{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":11568812,"sourceType":"datasetVersion","datasetId":7253205},{"sourceId":11569667,"sourceType":"datasetVersion","datasetId":7253605},{"sourceId":11569755,"sourceType":"datasetVersion","datasetId":7253661},{"sourceId":11675752,"sourceType":"datasetVersion","datasetId":7327562},{"sourceId":233737815,"sourceType":"kernelVersion"},{"sourceId":237664225,"sourceType":"kernelVersion"},{"sourceId":237775960,"sourceType":"kernelVersion"},{"sourceId":238047738,"sourceType":"kernelVersion"}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Github Source: https://github.com/KGML-lab/Generalized-Forward-Inverse-Framework-for-DL4SI\n\n| Model Name | Dataset | LB | Local cv | Version |\n| :----: | :----: | :----: | :----: | :----: |\n| UNetInverseModel 33M (5 epochs) | All Dataset | 143.1456 | 148.6 | v3 |","metadata":{}},{"cell_type":"code","source":"!pip install -q iunets","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import transforms as T\n\nimport os\nimport sys\nimport time\nimport datetime\nimport json\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import RandomSampler, DataLoader, Dataset, random_split\nfrom torch.utils.data.dataloader import default_collate\nimport torchvision\nfrom torchvision.transforms import Compose\n\nfrom iunets import iUNet\n\nfrom tqdm.notebook import tqdm\nfrom tqdm import tqdm\n\nimport iunet_network\nimport utils\n\nimport random\nimport numpy as np\n\nfrom pathlib import Path\n\nfrom torchvision.models import vgg16\nimport torchvision.models.vgg as vgg\n\nfrom matplotlib.colors import ListedColormap\nimport matplotlib.pyplot as plt\nfrom matplotlib.colors import TwoSlopeNorm\nfrom mpl_toolkits.axes_grid1 import make_axes_locatable\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:31:58.501573Z","iopub.execute_input":"2025-05-05T12:31:58.501838Z","iopub.status.idle":"2025-05-05T12:32:03.142096Z","shell.execute_reply.started":"2025-05-05T12:31:58.501816Z","shell.execute_reply":"2025-05-05T12:32:03.141561Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    fil_size = 500 # number of samples in each npy file\n\n    # Path related\n    output_path = '/kaggle/working/Invnet_models'\n    plot_directory = 'visualisation'\n    cfg_path = '/kaggle/input/unet-configs'\n\n    # Model realted\n    model = 'UNetInverseModel'\n    latent_dim = 70\n    skip = 1 # [0, 1] Unet skip connections 0:False, and 1:True\n    up_mode = None # upsampling layer mode such as \"nearest\", \"bicubic\", etc.\n    sample_spatial = 1.0 \n    sample_temporal = 1\n    optimizer = 'Adam'\n    lr_scheduler = 'StepLR'\n    unet_depth = 2\n    unet_repeat_blocks = 2\n\n    # Training related\n    batch_size = 64\n    lr = 0.001\n    lr_milestones = []\n    momentum = 0.9\n    weight_decay = 13-4\n    lr_gamma = 0.1\n    lr_warmup_epochs = 0\n    epoch_block = 3\n    num_block = 5\n    workers = 4\n    k = 1\n    print_freq = 250\n    resume = '/kaggle/input/unetinverse-training-inference-with-float16/Invnet_models/model_5.pth'\n    start_epoch = 0\n    plot_interval = 1\n    num_images = 5\n    rm_direct_arrival=1 # 'Remove direct arrival from amplitude data.'\n    velocity_transform = 'min_max'\n    amplitude_transform = 'normalize'\n    mask_factor = 0.0\n\n    # Loss realted\n    lambda_g1v = 1.0\n    lambda_g2v = 1.0\n    lambda_vel = 1.0\n    lambda_vgg_vel = 0.1\n    vgg_layer_output = 2 # VGG16 pretrained model layer output for perceptual loss calculation.\n    lambda_reg = 0.1\n    lambda_recons = 0.0\n\n    train_frac = 2\n    valid_frac = 16\n\nargs = CFG()\n\nargs.epochs = args.epoch_block * args.num_block\nargs.skip = True if args.skip==1 else False\n\n\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\nset_seed(1234)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:32:03.143431Z","iopub.execute_input":"2025-05-05T12:32:03.143971Z","iopub.status.idle":"2025-05-05T12:32:03.198697Z","shell.execute_reply.started":"2025-05-05T12:32:03.143950Z","shell.execute_reply":"2025-05-05T12:32:03.197946Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Utility","metadata":{}},{"cell_type":"code","source":"joint_model_list = [iunet_network.IUnetModel, iunet_network.JointModel, iunet_network.Decouple_IUnetModel]\ninverse_model_list = [\n                        iunet_network.InversionNet, \n                        iunet_network.IUnetInverseModel, \n                        iunet_network.UNetInverseModel,\n                        iunet_network.IUnetInverseModel_Legacy, \n                        iunet_network.UNetInverseModel_Legacy,\n                      ]\nrainbow_cmap = ListedColormap(np.load('/kaggle/input/unet-configs/rainbow256.npy'))\ndef get_optimizer(args, model, lr):\n    if args.optimizer == \"AdamW\":\n        optimizer = torch.optim.AdamW(model.parameters(), lr=lr, betas=(0.9, 0.999), weight_decay=args.weight_decay)\n    elif args.optimizer == \"Adam\":\n        optimizer = torch.optim.Adam(model.parameters(), lr=lr)\n    elif args.optimizer == \"Adadelta\":\n        optimizer = torch.optim.Adadelta(model.parameters(), lr=lr, rho=0.9, eps=1e-06, weight_decay=args.weight_decay) #lr = 1.0 default\n    elif args.optimizer == \"Adamax\":\n        optimizer = torch.optim.Adamax(model.parameters(), lr=lr, betas=(0.9, 0.999), eps=1e-08, weight_decay=args.weight_decay) # deafult lr = 0.002\n    else:\n        print(\"Optimizer not found\")\n    return optimizer\n\ndef get_lr_scheduler(args, optimizer):\n    # LR-Schedulers\n    if args.lr_scheduler == \"StepLR\":\n        lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=(args.epochs-args.start_epoch)//3, gamma=0.5)\n    elif args.lr_scheduler == \"LinearLR\":\n        lr_scheduler = torch.optim.lr_scheduler.LinearLR(optimizer,\n                                                         start_factor=1,\n                                                         end_factor = 1e-4, \n                                                         total_iters=args.epochs-args.start_epoch)\n    elif args.lr_scheduler == \"ExponentialLR\":\n        lr_scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.99)\n    elif args.lr_scheduler == \"CosineAnneal\":\n        lr_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs-args.start_epoch)\n    elif args.lr_scheduler == \"None\":\n        lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=(args.epochs-args.start_epoch), gamma=1.0) #Fake LR Scheduler\n    else:\n        print(\"LR Scheduler not found\")\n    return lr_scheduler\n\ndef get_transforms(args, ctx):\n    # label transforms\n    if args.velocity_transform == \"min_max\":\n        transform_label = T.MinMaxNormalize(ctx['label_min'], ctx['label_max'])\n    elif args.velocity_transform == \"normalize\":\n        transform_label = T.Normalize(ctx['label_mean'], ctx['label_std'])\n    elif args.velocity_transform == \"quantile\":\n        transform_label = T.QuantileTransform(n_quantiles=100)\n    else:\n        transform_label = None\n\n    # data transforms\n    if args.amplitude_transform == \"min_max\":\n        transform_data = T.MinMaxNormalize(ctx['data_min'], ctx['data_max'])\n    elif args.amplitude_transform == \"normalize\":\n        transform_data = T.Normalize(ctx['data_mean'], ctx['data_std'])\n    elif args.amplitude_transform == \"quantile\":\n        transform_data = T.QuantileTransform(n_quantiles=100)      \n    else:\n        transform_data=None\n    return transform_data, transform_label\n\n\ndef count_parameters(model, verbose=True):\n    num_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    if num_params >= 1e6:\n        num_params /= 1e6\n        suffix = \"M\"\n    elif num_params >= 1e3:\n        num_params /= 1e3\n        suffix = \"K\"\n    else:\n        suffix = \"\"\n    if verbose:\n        print(f\"Number of trainable parameters: {num_params:.2f}{suffix}\")\n    return num_params\n\n\nclass VGG16FeatureExtractor(torch.nn.Module):\n    def __init__(self):\n        super(VGG16FeatureExtractor, self).__init__()\n\n        vgg_model = vgg.vgg16(pretrained=True)\n        self.vgg_layers = vgg_model.features\n        self.layer_name_mapping = {\n            '3': \"relu1_2\",\n            '8': \"relu2_2\",\n            '15': \"relu3_3\",\n            '22': \"relu4_3\"\n        }\n    \n    def forward(self, x, vgg_layer_output=2):\n        assert vgg_layer_output <= len(self.layer_name_mapping)\n        \n        count = 0\n        for name, module in self.vgg_layers._modules.items():\n            x = module(x)\n            if name in self.layer_name_mapping:\n                if count == vgg_layer_output:\n                    return x\n                count += 1\n        return None\n\ndef plot_images(num_images, dataset, model, epoch, vis_folder, device, transform_data, transform_label, plot=True, save_key=\"results_epoch\"):\n    items = np.random.choice(len(dataset), num_images)\n\n    # _, amp_true, vel_true = dataset[items]\n\n    samples = [dataset[i] for i in items]\n\n    _, amp_true, vel_true = zip(*samples)\n    amp_true, vel_true = torch.tensor(amp_true).to(device), torch.tensor(vel_true).to(device)\n\n    model = model.to(device)\n\n    if np.any([isinstance(model, item) for item in joint_model_list]):\n        # if isinstance(model, iunet_network.JointModel) and isinstance(model.forward_model, forward_network.FNO2d):\n        #     amp_true = torch.einsum(\"ijkl->iklj\", amp_true)\n        #     vel_true = torch.einsum(\"ijkl->iklj\", vel_true)\n        with torch.autocast(device_type=\"cuda\"):\n            amp_pred = model.forward(vel_true).detach()\n        if transform_data is not None:\n            amp_true_np = transform_data.inverse_transform(amp_true.detach().cpu().numpy())\n            amp_pred_np = transform_data.inverse_transform(amp_pred.detach().cpu().numpy())\n        with torch.autocast(device_type=\"cuda\"):\n            vel_pred = model.inverse(amp_true).detach()\n        if transform_label is not None:\n            vel_true_np = transform_label.inverse_transform(vel_true.detach().cpu().numpy())\n            vel_pred_np = transform_label.inverse_transform(vel_pred.detach().cpu().numpy())\n        \n        fig, axes = plt.subplots(num_images, 6, figsize=(20, int(3*num_images)), dpi=150)\n        for i in range(num_images):\n            vel = np.concatenate([vel_pred_np[i, 0], vel_true_np[i, 0]], axis=1)\n            amp = np.concatenate([amp_pred_np[i, 0], amp_true_np[i, 0]], axis=1)\n            \n            min_vel, max_vel = vel.min(), vel.max()\n            \n            diff_vel = vel_pred_np[i, 0] - vel_true_np[i, 0]\n            diff_amp = amp_pred_np[i, 0] - amp_true_np[i, 0]\n            \n            ax = axes[i, 0]\n            img = ax.imshow(vel_true_np[i, 0], aspect='auto', vmin=min_vel, vmax=max_vel, cmap=rainbow_cmap)\n            divider = make_axes_locatable(ax)\n            cax = divider.append_axes(\"right\", size=\"10%\", pad=0.05)\n            plt.colorbar(img, cax=cax)\n            ax.set_title(f\"Velocity GT {i}\", fontsize=12)\n            ax.set_xticks([])\n            ax.set_yticks([])\n            \n            ax = axes[i, 1]\n            img = ax.imshow(vel_pred_np[i, 0], aspect='auto', vmin=min_vel, vmax=max_vel, cmap=rainbow_cmap)\n            divider = make_axes_locatable(ax)\n            cax = divider.append_axes(\"right\", size=\"10%\", pad=0.05)\n            plt.colorbar(img, cax=cax)\n            ax.set_title(f\"Velocity Predicted {i}\", fontsize=12)\n            ax.set_xticks([])\n            ax.set_yticks([])\n            \n            ax = axes[i, 2]\n            img = ax.imshow(diff_vel, aspect='auto', cmap=\"coolwarm\", norm=TwoSlopeNorm(vcenter=0))\n            divider = make_axes_locatable(ax)\n            cax = divider.append_axes(\"right\", size=\"10%\", pad=0.05)\n            plt.colorbar(img, cax=cax)\n            ax.set_title(f\"Velocity Difference {i}\", fontsize=12)\n            ax.set_xticks([])\n            ax.set_yticks([])\n            \n            ax = axes[i, 3]\n            img = ax.imshow(amp_true_np[i, 0], aspect='auto', vmin=-1, vmax=1, cmap=\"seismic\")\n            divider = make_axes_locatable(ax)\n            cax = divider.append_axes(\"right\", size=\"10%\", pad=0.05)\n            plt.colorbar(img, cax=cax)\n            ax.set_title(f\"Waveform GT {i}\", fontsize=12)\n            ax.set_xticks([])\n            ax.set_yticks([])\n            \n            ax = axes[i, 4]\n            img = ax.imshow(amp_pred_np[i, 0], aspect='auto', vmin=-1, vmax=1, cmap=\"seismic\")\n            divider = make_axes_locatable(ax)\n            cax = divider.append_axes(\"right\", size=\"10%\", pad=0.05)\n            plt.colorbar(img, cax=cax)\n            ax.set_title(f\"Waveform Predicted {i}\", fontsize=12)\n            ax.set_xticks([])\n            ax.set_yticks([])\n            \n            ax = axes[i, 5]\n            img = ax.imshow(diff_amp, aspect='auto', cmap=\"seismic\", vmin=-1, vmax=1)\n            divider = make_axes_locatable(ax)\n            cax = divider.append_axes(\"right\", size=\"10%\", pad=0.05)\n            plt.colorbar(img, cax=cax)\n            ax.set_title(f\"Waveform Difference {i}\", fontsize=12)\n            ax.set_xticks([])\n            ax.set_yticks([])\n\n    # plotting for inverse problem only\n    elif np.any([isinstance(model, item) for item in inverse_model_list]):\n        with torch.autocast(device_type=\"cuda\"):\n            vel_pred = model(amp_true).detach()\n        \n        if transform_label is not None:\n            vel_true_np = transform_label.inverse_transform(vel_true.detach().cpu().numpy())\n            vel_pred_np = transform_label.inverse_transform(vel_pred.detach().cpu().numpy())\n        \n        fig, axes = plt.subplots(num_images, 3, figsize=(10, int(3*num_images)), dpi=150)\n        for i in range(num_images):\n            vel = np.concatenate([vel_pred_np[i, 0], vel_true_np[i, 0]], axis=1)\n            \n            min_vel, max_vel = vel.min(), vel.max()\n            \n            diff_vel = vel_pred_np[i, 0] - vel_true_np[i, 0]\n            \n            ax = axes[i, 0]\n            img = ax.imshow(vel_true_np[i, 0], aspect='auto', vmin=min_vel, vmax=max_vel, cmap=rainbow_cmap)\n            divider = make_axes_locatable(ax)\n            cax = divider.append_axes(\"right\", size=\"10%\", pad=0.05)\n            plt.colorbar(img, cax=cax)\n            ax.set_title(f\"Velocity GT {i}\", fontsize=12)\n            ax.set_xticks([])\n            ax.set_yticks([])\n            \n            ax = axes[i, 1]\n            img = ax.imshow(vel_pred_np[i, 0], aspect='auto', vmin=min_vel, vmax=max_vel, cmap=rainbow_cmap)\n            divider = make_axes_locatable(ax)\n            cax = divider.append_axes(\"right\", size=\"10%\", pad=0.05)\n            plt.colorbar(img, cax=cax)\n            ax.set_title(f\"Velocity Predicted {i}\", fontsize=12)\n            ax.set_xticks([])\n            ax.set_yticks([])\n            \n            ax = axes[i, 2]\n            img = ax.imshow(diff_vel, aspect='auto', cmap=\"coolwarm\", norm=TwoSlopeNorm(vcenter=0))\n            divider = make_axes_locatable(ax)\n            cax = divider.append_axes(\"right\", size=\"10%\", pad=0.05)\n            plt.colorbar(img, cax=cax)\n            ax.set_title(f\"Velocity Difference {i}\", fontsize=12)\n            ax.set_xticks([])\n            ax.set_yticks([])\n     \n\n    # plotting for forward problem only\n    elif np.any([isinstance(model, item) for item in forward_model_list]):\n#         if isinstance(model, forward_network.FNO2d):\n#             amp_true = torch.einsum(\"ijkl->iklj\", amp_true)\n#             vel_true = torch.einsum(\"ijkl->iklj\", vel_true)\n        with torch.autocast(device_type=\"cuda\"):\n            amp_pred = model(vel_true).detach()\n        if transform_data is not None:\n            amp_true_np = transform_data.inverse_transform(amp_true.detach().cpu().numpy())\n            amp_pred_np = transform_data.inverse_transform(amp_pred.detach().cpu().numpy())\n\n        fig, axes = plt.subplots(num_images, 3, figsize=(10, int(3*num_images)), dpi=150)\n        for i in range(num_images):\n            amp = np.concatenate([amp_pred_np[i, 0], amp_true_np[i, 0]], axis=1)\n            diff_amp = amp_pred_np[i, 0] - amp_true_np[i, 0]\n            \n            ax = axes[i, 0]\n            img = ax.imshow(amp_true_np[i, 0], aspect='auto', vmin=-1, vmax=1, cmap=\"seismic\")\n            divider = make_axes_locatable(ax)\n            cax = divider.append_axes(\"right\", size=\"10%\", pad=0.05)\n            plt.colorbar(img, cax=cax)\n            ax.set_title(f\"Waveform GT {i}\", fontsize=12)\n            ax.set_xticks([])\n            ax.set_yticks([])\n            \n            ax = axes[i, 1]\n            img = ax.imshow(amp_pred_np[i, 0], aspect='auto', vmin=-1, vmax=1, cmap=\"seismic\")\n            divider = make_axes_locatable(ax)\n            cax = divider.append_axes(\"right\", size=\"10%\", pad=0.05)\n            plt.colorbar(img, cax=cax)\n            ax.set_title(f\"Waveform Predicted {i}\", fontsize=12)\n            ax.set_xticks([])\n            ax.set_yticks([])\n            \n            ax = axes[i, 2]\n            img = ax.imshow(diff_amp, aspect='auto', cmap=\"seismic\", vmin=-1, vmax=1)\n            divider = make_axes_locatable(ax)\n            cax = divider.append_axes(\"right\", size=\"10%\", pad=0.05)\n            plt.colorbar(img, cax=cax)\n            ax.set_title(f\"Waveform Difference {i}\", fontsize=12)\n            ax.set_xticks([])\n            ax.set_yticks([])\n    plt.tight_layout()\n    plt.savefig(os.path.join(vis_folder, f\"{save_key}_{epoch}.pdf\"))\n    if plot:\n        plt.show()\n    plt.close()","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-05-05T12:32:03.199613Z","iopub.execute_input":"2025-05-05T12:32:03.199873Z","iopub.status.idle":"2025-05-05T12:32:03.236563Z","shell.execute_reply.started":"2025-05-05T12:32:03.199855Z","shell.execute_reply":"2025-05-05T12:32:03.235814Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Ctx","metadata":{}},{"cell_type":"code","source":"ctx_dict = {\n    \"flatvel-a\": {\n        \"data_min\": -26.95,\n        \"data_max\": 52.77,\n        \"data_mean\": -3.3804e-05,\n        \"data_std\": 1.4797,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\n        \"label_mean\": 2782.0442,\n        \"label_std\": 786.1557,\n        \"file_size\": 500,\n        \"nbc\": 120,\n        \"dx\": 10,\n        \"nt\": 1000,\n        \"dt\": 1e-3,\n        \"f\": 15,\n        \"n_grid\": 70,\n        \"ns\": 5,\n        \"ng\": 70,\n        \"sz\": 10,\n        \"gz\": 10\n    },\n    \"curvevel-a\": {\n        \"data_min\": -27.11,\n        \"data_max\": 55.10,\n        \"data_mean\": -5.0961e-05,\n        \"data_std\": 1.4870,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\n        \"file_size\": 500,\n        \"label_mean\": 2788.1562,\n        \"label_std\": 794.2432,\n        \"nbc\": 120,\n        \"dx\": 10,\n        \"nt\": 1000,\n        \"dt\": 1e-3,\n        \"f\": 15,\n        \"n_grid\": 70,\n        \"ns\": 5,\n        \"ng\": 70,\n        \"sz\": 10,\n        \"gz\": 10\n    },\n    \"flatvel-b\": {\n        \"data_min\": -27.17,\n        \"data_max\": 56.05,\n        \"data_mean\": -0.0002,\n        \"data_std\": 1.7832,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\n        \"label_mean\": 3001.3389,\n        \"label_std\": 866.6636,\n        \"file_size\": 500,\n        \"nbc\": 120,\n        \"dx\": 10,\n        \"nt\": 1000,\n        \"dt\": 1e-3,\n        \"f\": 15,\n        \"n_grid\": 70,\n        \"ns\": 5,\n        \"ng\": 70,\n        \"sz\": 10,\n        \"gz\": 10\n    },\n    \"curvevel-b\": {\n        \"data_min\": -29.04,\n        \"data_max\": 57.03,\n        \"data_mean\": -0.0002,\n        \"data_std\": 1.7648,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\n        \"label_mean\": 3000.5669,\n        \"label_std\": 865.4404,\n        \"file_size\": 500,\n        \"nbc\": 120,\n        \"dx\": 10,\n        \"nt\": 1000,\n        \"dt\": 1e-3,\n        \"f\": 15,\n        \"n_grid\": 70,\n        \"ns\": 5,\n        \"ng\": 70,\n        \"sz\": 10,\n        \"gz\": 10\n    },\n\t\"flatfault-a\": {\n        \"data_min\": -26.10,\n        \"data_max\": 50.86,\n        \"data_mean\": -0.00043503073,\n        \"data_std\": 1.5410482,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\n        \"label_mean\": 3088.6873,\n        \"label_std\": 855.37024,\n        \"file_size\": 500,\n        \"nbc\": 120,\n        \"dx\": 10,\n        \"nt\": 1000,\n        \"dt\": 1e-3,\n        \"f\": 15,\n        \"n_grid\": 70,\n        \"ns\": 5,\n        \"ng\": 70,\n        \"sz\": 10,\n        \"gz\": 10\n    },\n    \"curvefault-a\": {\n        \"data_min\": -26.48,\n        \"data_max\": 52.32,\n        \"data_mean\": -0.00045603843,\n        \"data_std\": 1.5448948,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\n        \"label_mean\": 3082.6616,\n        \"label_std\": 852.38995,\n        \"file_size\": 500,\n        \"nbc\": 120,\n        \"dx\": 10,\n        \"nt\": 1000,\n        \"dt\": 1e-3,\n        \"f\": 15,\n        \"n_grid\": 70,\n        \"ns\": 5,\n        \"ng\": 70,\n        \"sz\": 10,\n        \"gz\": 10\n    },\n    \"flatfault-b\": {\n        \"data_min\": -24.86,\n        \"data_max\": 50.28,\n        \"data_mean\": -0.0001,\n        \"data_std\": 1.4952,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\n        \"file_size\": 500,\n        \"label_mean\": 3055.4231,\n        \"label_std\": 875.8992,\n        \"nbc\": 120,\n        \"dx\": 10,\n        \"nt\": 1000,\n        \"dt\": 1e-3,\n        \"f\": 15,\n        \"n_grid\": 70,\n        \"ns\": 5,\n        \"ng\": 70,\n        \"sz\": 10,\n        \"gz\": 10\n    },\n    \"curvefault-b\": {\n        \"data_min\": -24.93,\n        \"data_max\": 50.98,\n        \"data_mean\": -8.882544e-05,\n        \"data_std\": 1.50228,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\n        \"label_mean\": 3035.5576,\n        \"label_std\": 890.48785,\n        \"file_size\": 500,\n        \"nbc\": 120,\n        \"dx\": 10,\n        \"nt\": 1000,\n        \"dt\": 1e-3,\n        \"f\": 15,\n        \"n_grid\": 70,\n        \"ns\": 5,\n        \"ng\": 70,\n        \"sz\": 10,\n        \"gz\": 10\n    },\n    \"style-a\": {\n        \"data_min\": -24.96,\n        \"data_max\": 48.93,\n        \"data_mean\": 0.00024991733,\n        \"data_std\": 1.4618,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\n        \"label_mean\": 2728.5144,\n        \"label_std\": 665.83215,\n        \"file_size\": 500,\n        \"nbc\": 120,\n        \"dx\": 10,\n        \"nt\": 1000,\n        \"dt\": 1e-3,\n        \"f\": 15,\n        \"n_grid\": 70,\n        \"ns\": 5,\n        \"ng\": 70,\n        \"sz\": 10,\n        \"gz\": 10\n    },\n    \"style-b\": {\n        \"data_min\": -23.76,\n        \"data_max\": 46.01,\n        \"data_mean\": 0.00013498258,\n        \"data_std\": 1.4579,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\n        \"label_mean\": 2837.3164,\n        \"label_std\": 637.6763,\n        \"file_size\": 500,\n        \"nbc\": 120,\n        \"dx\": 10,\n        \"nt\": 1000,\n        \"dt\": 1e-3,\n        \"f\": 15,\n        \"n_grid\": 70,\n        \"ns\": 5,\n        \"ng\": 70,\n        \"sz\": 10,\n        \"gz\": 10\n    },\n    \"flatvel-tutorial\": {\n        \"data_min\": -26.95,\n        \"data_max\": 52.77,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\n        \"file_size\": 120,\n        \"nbc\": 120,\n        \"dx\": 10,\n        \"nt\": 1000,\n        \"dt\": 1e-3,\n        \"f\": 15,\n        \"n_grid\": 70,\n        \"ns\": 5,\n        \"ng\": 70,\n        \"sz\": 10,\n        \"gz\": 10\n    }\n}","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-05-05T12:32:03.238424Z","iopub.execute_input":"2025-05-05T12:32:03.238990Z","iopub.status.idle":"2025-05-05T12:32:03.251102Z","shell.execute_reply.started":"2025-05-05T12:32:03.238973Z","shell.execute_reply":"2025-05-05T12:32:03.250454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ctx = ctx_dict['flatvel-a']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:32:03.251768Z","iopub.execute_input":"2025-05-05T12:32:03.252027Z","iopub.status.idle":"2025-05-05T12:32:03.268241Z","shell.execute_reply.started":"2025-05-05T12:32:03.252006Z","shell.execute_reply":"2025-05-05T12:32:03.267482Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# WarmupMultiStepLR","metadata":{}},{"cell_type":"code","source":"class WarmupMultiStepLR(torch.optim.lr_scheduler._LRScheduler):\n    def __init__(\n        self,\n        optimizer,\n        milestones,\n        gamma=0.1,\n        warmup_factor=1.0 / 3,\n        warmup_iters=5,\n        warmup_method=\"linear\",\n        last_epoch=-1,\n    ):\n        if not milestones == sorted(milestones):\n            raise ValueError(\n                \"Milestones should be a list of\" \" increasing integers. Got {}\",\n                milestones,\n            )\n\n        if warmup_method not in (\"constant\", \"linear\"):\n            raise ValueError(\n                \"Only 'constant' or 'linear' warmup_method accepted\"\n                \"got {}\".format(warmup_method)\n            )\n        self.milestones = milestones\n        self.gamma = gamma\n        self.warmup_factor = warmup_factor\n        self.warmup_iters = warmup_iters\n        self.warmup_method = warmup_method\n        super(WarmupMultiStepLR, self).__init__(optimizer, last_epoch)\n\n    def get_lr(self):\n        warmup_factor = 1\n        if self.last_epoch < self.warmup_iters:\n            if self.warmup_method == \"constant\":\n                warmup_factor = self.warmup_factor\n            elif self.warmup_method == \"linear\":\n                alpha = float(self.last_epoch) / self.warmup_iters\n                warmup_factor = self.warmup_factor * (1 - alpha) + alpha\n        return [\n            base_lr *\n            warmup_factor *\n            self.gamma ** bisect_right(self.milestones, self.last_epoch)\n            for base_lr in self.base_lrs\n        ]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:32:03.269019Z","iopub.execute_input":"2025-05-05T12:32:03.269214Z","iopub.status.idle":"2025-05-05T12:32:03.282790Z","shell.execute_reply.started":"2025-05-05T12:32:03.269169Z","shell.execute_reply":"2025-05-05T12:32:03.282216Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class FWIDataset(Dataset):\n    ''' FWI dataset\n    For convenience, in this class, a batch refers to a npy file \n    instead of the batch used during training.\n\n    Args:\n        preload: whether to load the whole dataset into memory\n        sample_ratio: downsample ratio for seismic data\n        file_size: # of samples in each npy file\n        transform_data|label: transformation applied to data or label\n    '''\n    def __init__(self, inputs, outputs, preload=True, sample_ratio=1, file_size=500,\n                    transform_data=None, transform_label=None, mask_factor=0.0):\n        self.preload = preload\n        self.sample_ratio = sample_ratio\n        self.file_size = file_size\n        self.transform_data = transform_data\n        self.transform_label = transform_label\n        if outputs is not None:\n            self.batches = [str(inputs[i])+'&'+str(outputs[i]) for i in range(len(inputs))]\n        else:\n            self.batches = [str(inputs[i]) for i in range(len(inputs))]\n            \n        if preload:\n            self.data_list, self.label_list= (), ()\n            for batch in tqdm(self.batches):\n                data, label = self.load_every(batch) \n\n                self.data_list = self.data_list + (data,)\n                self.label_list = self.label_list + (label,)\n\n            self.data_list = np.concatenate(self.data_list, 0)\n            self.label_list = np.concatenate(self.label_list, 0)\n\n            mask_indices = np.random.choice(len(self.data_list), \n                                              int(mask_factor*len(self.data_list)),\n                                              replace=False)\n            \n            self.mask_list = np.ones(len(self.data_list), dtype=np.int8)\n            self.mask_list[mask_indices] = 0\n\n            print(\"Data concatenation complete.\")\n            if self.transform_data is not None:\n                self.data_list = self.transform_data(self.data_list)\n            if self.transform_label is not None:\n                self.label_list = self.transform_label(self.label_list)\n\n    # Load from one line\n    def load_every(self, batch):\n        batch = batch.split('&')\n        data_path = batch[0] if len(batch) > 1 else batch[0][:-1]\n        data = np.load(data_path, mmap_mode='r')[:, :, ::self.sample_ratio, :]\n        # data = data.astype('float32')\n        if len(batch) > 1:\n            label_path = batch[1]\n            label = np.load(label_path, mmap_mode='r')\n            # label = label.astype('float32')\n        else:\n            label = None\n        \n        return data, label\n        \n    def __getitem__(self, idx):\n        batch_idx, sample_idx = idx // self.file_size, idx % self.file_size\n        if self.preload:\n            mask = self.mask_list[idx]\n            data = self.data_list[idx]\n            label = self.label_list[idx] if len(self.label_list) != 0 else None\n        else:\n            data, label = self.load_every(self.batches[batch_idx])\n            data = data[sample_idx]\n            label = label[sample_idx] if label is not None else None\n            mask=np.array([1])\n\n        \n            if self.transform_data is not None:\n                data = self.transform_data(data)\n            if self.transform_label is not None:\n                label = self.transform_label(label)\n        return mask, data, label if label is not None else np.array([])\n        \n    def __len__(self):\n        return len(self.batches) * self.file_size","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:32:03.283497Z","iopub.execute_input":"2025-05-05T12:32:03.284061Z","iopub.status.idle":"2025-05-05T12:32:03.300800Z","shell.execute_reply.started":"2025-05-05T12:32:03.284038Z","shell.execute_reply":"2025-05-05T12:32:03.300290Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class FWIDataset(Dataset):\n#     def __init__(self, inputs_files, output_files, n_examples_per_file=500, transform_data=None, transform_label=None):\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.transform_data = transform_data\n#         self.transform_label =transform_label\n\n#     def __len__(self):\n#         return len(self.inputs_files) * self.n_examples_per_file\n\n#     def __getitem__(self, idx):\n#         # Calculate file offset and sample offset within file\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#         try:\n#             data = X[sample_idx].copy()\n#             label = y[sample_idx].copy()\n#             mask = np.array([1])\n#             if self.transform_data is not None:\n#                 data = self.transform_data(data)\n#             if self.transform_label is not None:\n#                 label = self.transform_data(label)\n#             return mask, data, label\n#         finally:\n#             del X, y","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-05-05T12:32:03.301555Z","iopub.execute_input":"2025-05-05T12:32:03.301810Z","iopub.status.idle":"2025-05-05T12:32:03.317867Z","shell.execute_reply.started":"2025-05-05T12:32:03.301783Z","shell.execute_reply":"2025-05-05T12:32:03.317298Z"}},"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, all_outputs = [], []\nfor path in [\"/kaggle/input/open-wfi-1/openfwi_float16_1\", \"/kaggle/input/open-wfi-2/openfwi_float16_2\"]:\n\n    all_inputs1 = [\n        f\n        for f in\n        Path(path).rglob('*.npy')\n        if ('seis' in f.stem) or ('data' in f.stem)\n    ]\n    \n    \n    all_outputs1 = inputs_files_to_output_files(all_inputs)\n\n    all_inputs.extend(all_inputs1)\n    all_outputs.extend(all_outputs1)\n\nvalid_inputs = [all_inputs[i] for i in range(0, len(all_inputs), args.valid_frac)]\ntrain_inputs = [f for f in all_inputs if not f in valid_inputs]\nif args.train_frac > 1:\n    train_inputs = [train_inputs[i] for i in range(0, len(train_inputs), args.train_frac)]\n\ntrain_outputs = inputs_files_to_output_files(train_inputs)\nvalid_outputs = inputs_files_to_output_files(valid_inputs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:32:03.318486Z","iopub.execute_input":"2025-05-05T12:32:03.318671Z","iopub.status.idle":"2025-05-05T12:32:04.249591Z","shell.execute_reply.started":"2025-05-05T12:32:03.318656Z","shell.execute_reply":"2025-05-05T12:32:04.248821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transform_data, transform_label = get_transforms(args, ctx)\n\nprint('Loading data')\nprint('Loading training data')\ndataset_train = FWIDataset(\n            train_inputs,\n            train_outputs,\n            preload=False,\n            sample_ratio=args.sample_temporal,\n            file_size=ctx['file_size'],\n            transform_data=transform_data,\n            transform_label=transform_label,\n            mask_factor=args.mask_factor\n        )\n\nprint('Loading validation data')\ndataset_valid = FWIDataset(\n            valid_inputs,\n            valid_outputs,\n            preload=False,\n            sample_ratio=args.sample_temporal,\n            file_size=ctx['file_size'],\n            transform_data=transform_data,\n            transform_label=transform_label,\n            mask_factor=args.mask_factor\n        )\n\n\n# dataset_size = len(dataset_train)\n# train_size = int(0.9 * dataset_size)  \n# val_size = dataset_size - train_size \n\n# dataset_train, dataset_valid = random_split(dataset_train, [train_size, val_size])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:32:04.251952Z","iopub.execute_input":"2025-05-05T12:32:04.252241Z","iopub.status.idle":"2025-05-05T12:32:04.258156Z","shell.execute_reply.started":"2025-05-05T12:32:04.252216Z","shell.execute_reply":"2025-05-05T12:32:04.257427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('Creating data loaders')\ntrain_sampler = RandomSampler(dataset_train)\nvalid_sampler = RandomSampler(dataset_valid)\n\ndataloader_train = DataLoader(\n        dataset_train, batch_size=args.batch_size,\n        sampler=train_sampler, num_workers=args.workers,\n        pin_memory=True, drop_last=True, collate_fn=default_collate)\n\ndataloader_valid = DataLoader(\n    dataset_valid, batch_size=args.batch_size,\n    sampler=valid_sampler, num_workers=args.workers,\n    pin_memory=True, collate_fn=default_collate)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:32:04.258939Z","iopub.execute_input":"2025-05-05T12:32:04.259137Z","iopub.status.idle":"2025-05-05T12:32:04.273705Z","shell.execute_reply.started":"2025-05-05T12:32:04.259123Z","shell.execute_reply":"2025-05-05T12:32:04.273084Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"step = 0\n\ndef train_one_epoch(model, vgg_model, masked_criterion, optimizer, lr_scheduler, \n                    dataloader, device, epoch, print_freq):\n    global step\n    model.train()\n\n    # Logger setup\n    metric_logger = utils.MetricLogger(delimiter='  ')\n    metric_logger.add_meter('lr', utils.SmoothedValue(window_size=1, fmt='{value}'))\n    metric_logger.add_meter('samples/s', utils.SmoothedValue(window_size=10, fmt='{value:.3f}'))\n    header = 'Epoch: [{}]'.format(epoch)\n\n    vel_vgg_loss = torch.tensor([0.], device=device)\n    \n    # for VGG amp loss\n    upsample = torch.nn.Upsample(size=(70, 70), mode=\"bicubic\")\n    \n    for mask, amp, vel in metric_logger.log_every(dataloader, print_freq, header):\n        start_time = time.time()\n        \n        optimizer.zero_grad()\n\n        mask, amp, vel = mask.to(device), amp.to(device), vel.to(device)\n        identity_mask = torch.ones_like(mask)\n\n        with torch.autocast(device_type=\"cuda\"):\n            vel_pred = model(amp)\n\n            vel_loss, vel_loss_g1v, vel_loss_g2v = masked_criterion(vel_pred, vel, mask)\n\n            # Calculating the perceptual loss using VGG-16 model for velocity and amplitude\n            if args.lambda_vgg_vel>0:\n                vgg_vel = vel.repeat(1,3,1,1)\n                vgg_vel_pred = vel_pred.repeat(1,3,1,1)\n    \n                with torch.no_grad():\n                    vgg_features_vel = vgg_model(vgg_vel, vgg_layer_output=args.vgg_layer_output)   \n                vgg_features_vel_pred = vgg_model(vgg_vel_pred, vgg_layer_output=args.vgg_layer_output)\n    \n                vel_vgg_loss, vel_vgg_loss_g1v, vel_vgg_loss_g2v = masked_criterion(vgg_features_vel, vgg_features_vel_pred, mask)\n    \n                vel_vgg_loss_g1v_val = vel_vgg_loss_g1v.item()\n                vel_vgg_loss_g2v_val = vel_vgg_loss_g2v.item()\n    \n                metric_logger.update(vel_vgg_loss_g1v = vel_vgg_loss_g1v_val, \n                                 vel_vgg_loss_g2v = vel_vgg_loss_g2v_val)\n                \n    \n            # Calcultaing the reconstruction loss on encoder-decoder for both amp and vel     \n            amp_loss_recons = 0  \n            vel_loss_recons = 0 \n            if args.lambda_recons>0:\n                # print(\"applying reconstruction\")\n                vel_recons = model.vel_model.forward(vel)\n                amp_recons = model.amp_model.forward(amp)\n                amp_loss_recons = nn.MSELoss()(amp_recons, amp)  # Compute amplitude loss\n                vel_loss_recons = nn.MSELoss()(vel_recons, vel)  # Compute velocity loss\n                \n                metric_logger.update(amp_loss_recons = amp_loss_recons, \n                                 vel_loss_recons = vel_loss_recons)\n                \n    \n    \n            \n            loss = args.lambda_vel * vel_loss + args.lambda_vgg_vel * vel_vgg_loss + args.lambda_recons * amp_loss_recons + args.lambda_recons * vel_loss_recons\n\n        loss.backward()\n        optimizer.step()\n\n        loss_val = loss.item()\n        vel_loss_g1v_val = vel_loss_g1v.item()\n        vel_loss_g2v_val = vel_loss_g2v.item()\n\n        batch_size = amp.shape[0]\n        metric_logger.update(\n                             loss = loss_val, \n                             lr = optimizer.param_groups[0]['lr'],\n                             vel_loss_g1v = vel_loss_g1v_val,\n                             vel_loss_g2v = vel_loss_g2v_val,\n                            )\n        \n        metric_logger.meters['samples/s'].update(batch_size / (time.time() - start_time))\n\n        step += 1\n\n\ndef evaluate(model, criterion, dataloader, device, ctx, transform_data=None, transform_label=None):\n    l1_loss = nn.L1Loss()\n    l2_loss = nn.MSELoss()\n\n    metric_logger = utils.MetricLogger(delimiter='  ')\n    header = 'Test:'\n\n    model.eval()\n\n    all_outputs = []\n    all_labels = []\n    with torch.no_grad():\n        total_samples = 0\n        \n        eval_metrics = [\"vel_sum_abs_error\", \"vel_sum_squared_error\"]\n        eval_dict = {}\n        for metric in eval_metrics:\n            eval_dict[metric] = 0\n\n        val_loss = 0\n        for _, amp, vel in metric_logger.log_every(dataloader, 100, header):\n            \n            amp = amp.to(device, non_blocking=True)\n            vel = vel.to(device, non_blocking=True)\n\n            batch_size = amp.shape[0]\n            total_samples += batch_size\n            with torch.autocast(device_type=\"cuda\"):\n                vel_pred = model(amp)\n\n            vel_loss, vel_loss_g1v, vel_loss_g2v = criterion(vel_pred, vel)\n\n            loss = vel_loss\n            val_loss += loss.item()\n\n            eval_dict[\"vel_sum_abs_error\"] += (l1_loss(vel, vel_pred) * batch_size).item()\n            eval_dict[\"vel_sum_squared_error\"] += (l2_loss(vel, vel_pred) * batch_size).item()\n\n            all_outputs.append(vel_pred.detach().cpu())\n            all_labels.append(vel.detach().cpu())\n        \n        for metric in eval_metrics:\n            eval_dict[metric] /= total_samples \n\n    val_loss /= len(dataloader)\n\n    all_output = torch.concat(all_outputs, axis=0)\n    all_label = torch.concat(all_labels, axis=0)\n\n    all_output = transform_label.inverse_transform(all_output.numpy())\n    all_label = transform_label.inverse_transform(all_label.numpy())\n\n    \n    l1loss_eval = l1loss(torch.tensor(all_output), torch.tensor(all_label))\n\n    eval_dict[\"Val_Loss\"] = val_loss\n    eval_dict[\"L1_Loss\"] = l1loss_eval.item()\n    \n    return val_loss, l1loss_eval, eval_dict","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:32:04.274332Z","iopub.execute_input":"2025-05-05T12:32:04.274537Z","iopub.status.idle":"2025-05-05T12:32:04.292590Z","shell.execute_reply.started":"2025-05-05T12:32:04.274522Z","shell.execute_reply":"2025-05-05T12:32:04.291877Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(args.device)\ntorch.backends.cudnn.benchmark = True\n\nprint('Creating model')\nif args.model not in iunet_network.model_dict:\n    print('Unsupported model.')\n    sys.exit()\n\n\ndef set_inverse_params(args, inverse_model_params):\n        inverse_model_params.setdefault('IUnetInverseModel', {})\n        inverse_model_params['IUnetInverseModel']['cfg_path'] = args.cfg_path\n        inverse_model_params['IUnetInverseModel']['latent_dim'] = args.latent_dim\n        \n        inverse_model_params.setdefault('UNetInverseModel', {})\n        inverse_model_params['UNetInverseModel']['cfg_path'] = args.cfg_path\n        inverse_model_params['UNetInverseModel']['latent_dim'] = args.latent_dim\n        inverse_model_params['UNetInverseModel']['unet_depth'] = args.unet_depth\n        inverse_model_params['UNetInverseModel']['unet_repeat_blocks'] = args.unet_repeat_blocks\n        inverse_model_params['UNetInverseModel']['skip'] = args.skip\n        return inverse_model_params\n    \n# creating inverse model\ninverse_model_params = iunet_network.inverse_params\ninverse_model_params = set_inverse_params(args, inverse_model_params)\n\nmodel = iunet_network.model_dict[args.model](**inverse_model_params[args.model]).to(device)\n\nprint(count_parameters(model))\n\nvgg_model = VGG16FeatureExtractor().to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:32:04.293328Z","iopub.execute_input":"2025-05-05T12:32:04.293522Z","iopub.status.idle":"2025-05-05T12:32:09.508651Z","shell.execute_reply.started":"2025-05-05T12:32:04.293507Z","shell.execute_reply":"2025-05-05T12:32:09.508066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define loss function\nl1loss = nn.L1Loss()\nl2loss = nn.MSELoss()\n\ndef masked_criterion(pred, gt, mask):\n    B, C, H, W = pred.shape\n    mask = mask.view(B, 1, 1, 1)\n    num_elements = mask.sum() + 1\n\n    squared_diff = ((pred - gt)**2) * mask\n    abs_diff = (pred-gt).abs() * mask\n    norm_l2_loss = torch.sum(squared_diff.mean(dim=[1, 2, 3])/num_elements)\n    norm_l1_loss = torch.sum(abs_diff.mean(dim=[1, 2, 3])/num_elements)\n    loss = args.lambda_g1v * norm_l1_loss + args.lambda_g2v * norm_l2_loss\n    return loss, norm_l1_loss, norm_l2_loss\n\ndef criterion(pred, gt):\n    loss_g1v = l1loss(pred, gt)\n    loss_g2v = l2loss(pred, gt)\n    loss = args.lambda_g1v * loss_g1v + args.lambda_g2v * loss_g2v\n    return loss, loss_g1v, loss_g2v\n\ndef relative_l2_error(pred, gt):\n    batch_size = gt.shape[0]\n    pred = pred.view(batch_size, -1)\n    gt = gt.view(batch_size, -1)\n\n    numerator = torch.linalg.norm(pred - gt, ord=2, dim=1)\n    denominator = torch.linalg.norm(gt, ord=2, dim=1)\n    relative_loss = (numerator/denominator).mean()\n    return relative_loss\n\n# Scale lr according to effective batch size\nlr = args.lr\noptimizer = get_optimizer(args, model, lr)\nlr_scheduler = get_lr_scheduler(args, optimizer)\n\n# Convert scheduler to be per iteration instead of per epoch\nwarmup_iters = args.lr_warmup_epochs * len(dataloader_train)\nlr_milestones = [len(dataloader_train) * m for m in args.lr_milestones]\n\nmodel_without_ddp = model\n\nif args.resume:\n    checkpoint = torch.load(args.resume, map_location='cpu')\n    model_without_ddp.load_state_dict(iunet_network.replace_legacy(checkpoint['model']))\n    optimizer.load_state_dict(checkpoint['optimizer'])\n    lr_scheduler.load_state_dict(checkpoint['lr_scheduler'])\n    args.start_epoch = checkpoint['epoch'] + 1\n    step = checkpoint['step']\n    lr_scheduler.milestones = lr_milestones","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:32:09.509369Z","iopub.execute_input":"2025-05-05T12:32:09.509636Z","iopub.status.idle":"2025-05-05T12:32:09.518833Z","shell.execute_reply.started":"2025-05-05T12:32:09.509617Z","shell.execute_reply":"2025-05-05T12:32:09.518039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('Start training')\nstart_time = time.time()\nbest_loss = 10000\nchp = 1 \n\ntrain_vis_folder = os.path.join(args.output_path, args.plot_directory, 'train')\nvalid_vis_folder = os.path.join(args.output_path, args.plot_directory, 'validation')\n\nos.makedirs(train_vis_folder, exist_ok=True)\nos.makedirs(valid_vis_folder, exist_ok=True)\nfor epoch in tqdm(range(args.start_epoch, args.epochs)):\n    train_one_epoch(model, vgg_model, masked_criterion, optimizer, lr_scheduler, dataloader_train,\n                    device, epoch, args.print_freq)\n    \n    lr_scheduler.step()    \n   \n    val_loss, loss, eval_dict = evaluate(model, criterion, dataloader_valid, device, ctx, transform_data, transform_label)\n    print(\"Test Metrics:\", eval_dict)\n\n    checkpoint = {\n        'model': model_without_ddp.state_dict(),\n        'optimizer': optimizer.state_dict(),\n        'lr_scheduler': lr_scheduler.state_dict(),\n        'epoch': epoch,\n        'step': step\n        }\n\n    if (epoch+1)%args.plot_interval == 0:\n        plot_images(args.num_images, dataset_train, model, epoch, train_vis_folder, device, transform_data, transform_label)\n        plot_images(args.num_images, dataset_valid, model, epoch, valid_vis_folder, device, transform_data, transform_label)\n\n        torch.save(\n            checkpoint,\n            os.path.join(args.output_path, 'latest_checkpoint.pth'))\n\n    \n    # Save checkpoint per epoch\n    if loss < best_loss:\n        torch.save(\n            checkpoint,\n            os.path.join(args.output_path, 'checkpoint.pth'))\n        print('saving checkpoint at epoch: ', epoch)\n        chp = epoch\n        best_loss = loss\n        \n    # Save checkpoint every epoch block\n    print('current best loss: ', best_loss)\n    print('current best epoch: ', chp)\n    if args.output_path and (epoch + 1) % args.epoch_block == 0:\n        torch.save(\n            checkpoint,\n            os.path.join(args.output_path, 'model_{}.pth'.format(epoch + 1)))\n\ntotal_time = time.time() - start_time\ntotal_time_str = str(datetime.timedelta(seconds=int(total_time)))\nprint('Training time {}'.format(total_time_str))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:32:09.519500Z","iopub.execute_input":"2025-05-05T12:32:09.519660Z","execution_failed":"2025-05-05T13:22:08.898Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Predict","metadata":{}},{"cell_type":"code","source":"import csv  # Use \"low-level\" CSV to save memory on predictions","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-05T13:22:08.899Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\ntest_files = list(Path('/kaggle/input/open-wfi-test/test').glob('*.npy'))\nlen(test_files)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-05T13:22:08.899Z"}},"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","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-05T13:22:08.899Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, test_files, transform_data=None):\n        self.test_files = test_files\n        self.transform_data = transform_data\n\n\n    def __len__(self):\n        return len(self.test_files)\n\n\n    def __getitem__(self, i):\n        test_file = self.test_files[i]\n        data = np.load(test_file)\n        if self.transform_data:\n            data = self.transform_data(data)\n\n        return data, test_file.stem","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-05T13:22:08.899Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = TestDataset(test_files, transform_data)\ndl = DataLoader(ds, batch_size=16, num_workers=4, pin_memory=True)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-05T13:22:08.899Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"checkpoint = torch.load('/kaggle/working/Invnet_models/checkpoint.pth')\n\nmodel_without_ddp.load_state_dict(iunet_network.replace_legacy(checkpoint['model']))\n\n# Train\nmodel.eval()\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(dl, desc='test'):\n        inputs = inputs.to(device)\n        with torch.inference_mode():\n            with torch.autocast(device_type=\"cuda\"):\n                outputs = model(inputs)\n\n        y_preds = outputs[:, 0].cpu().numpy()\n        y_preds = transform_label.inverse_transform(y_preds)\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)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-05T13:22:08.899Z"}},"outputs":[],"execution_count":null}]}