{"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":11334027,"sourceType":"datasetVersion","datasetId":7089850},{"sourceId":11430896,"sourceType":"datasetVersion","datasetId":7142133},{"sourceId":233737401,"sourceType":"kernelVersion"},{"sourceId":233737815,"sourceType":"kernelVersion"},{"sourceId":233738664,"sourceType":"kernelVersion"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n| Model Name | Dataset | LB | Local cv | Version |\n| :----: | :----: | :----: | :----: | :----: |\n| pretrained by flatfault_b_l2_480 (200 epochs) | Kaggle Dataset | 211.3 | 201.9 | v19(Training) v20(Inference) |\n| pretrained by flatfault_b_l2_480 | Kaggle Dataset | 215.6 | 209.8 | v16 |\n| flatfault_b_l2_480 | FlatFault-B | 249.6 | - | v10 |\n| - | Kaggle Dataset | 257.1 | 251.9 | v2 |\n| curvevel_b_l1_480 | CurveVel-B | 258.2 | - | v13 |\n| flatfault_b_l1_480 | FlatFault-B | 259.1 | - | v11 |\n| curvevel_b_l2_480 | CurveVel-B | 278.7 | - | v12 |\n| flatvel_a_l2_480 | FlatVel-A | 474.8 | - | v9 |","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport time\nimport datetime\nimport json\n\nimport numpy as np\n\nfrom pathlib import Path\n\nimport matplotlib.pyplot as plt\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import RandomSampler, DataLoader, Dataset\nfrom torch.utils.data.dataloader import default_collate\nfrom torch.utils.data.distributed import DistributedSampler\nfrom torch.utils.tensorboard import SummaryWriter\nimport torchvision\nfrom torchvision.transforms import Compose\nfrom bisect import bisect_right\n\nimport utils\nimport network\nimport transforms as T\n\nimport random\nimport gc\n\nfrom tqdm.auto import tqdm\n\n# Need to use parallel in apex, torch ddp can cause bugs when computing gradient penalty\n# import apex.parallel as parallel","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-16T01:42:36.234072Z","iopub.execute_input":"2025-04-16T01:42:36.234280Z","iopub.status.idle":"2025-04-16T01:42:57.243391Z","shell.execute_reply.started":"2025-04-16T01:42:36.234264Z","shell.execute_reply":"2025-04-16T01:42:57.242647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    # model related\n    model = 'InversionNet' #  'generator name'\n    model_d = 'Discriminator' # 'discriminator name'\n    up_mode = None # 'upsampling layer mode such as \"nearest\", \"bicubic\", etc.'\n    sample_spatial = 1.0 # 'spatial sampling ratio'\n    sample_temporal = 1 # 'temporal sampling ratio'\n\n    # Loss related\n    lambda_g1v = 100.0\n    lambda_g2v = 0.0\n    lambda_adv = 1.0\n    lambda_gp = 10.0\n\n    # Training ralted\n    k = 1 # 'k in log transformation'\n    weight_decay = 1e-4\n    batch_size = 64\n    n_critic = 5 # 'generator & discriminator update ratio'\n    lr_g = 0.0001 # 'initial learning rate of generator'\n    lr_d = 0.0001 # 'initial learning rate of discriminator'\n    lr_milestones = [] # 'decrease lr on milestones'\n    momentum = 0.9 # momentum\n    lr_gamma = 0.1 # 'decrease lr by a factor of lr-gamma'\n    lr_warmup_epochs = 0 # 'number of warmup epochs'\n    epoch_block = 40 # 'epochs in a saved block'\n    num_block = 5 # 'number of saved block'\n    workers = 4\n    print_freq = 20 # 'print frequency'\n    start_epoch = 0 # 'start epoch'\n\n    pretrained = True\n    pretrain_path = '/kaggle/input/waveform-inversion-models/pretrained_models/VelocityGAN/flatvel_b_l2_480.pth'\n\n    resume = None\n\n    output_path = '/kaggle/working/'\n\n    seed = 2025\n\n    run_train = False\n\nargs = CFG()\n\ndef seed_torch(seed_value):\n    random.seed(seed_value) # Python\n    np.random.seed(seed_value) # cpu vars\n    torch.manual_seed(seed_value) # cpu  vars    \n    if torch.cuda.is_available(): \n        torch.cuda.manual_seed(seed_value)\n        torch.cuda.manual_seed_all(seed_value) # gpu vars\n    if torch.backends.cudnn.is_available:\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n\nseed_torch(args.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T01:42:57.244161Z","iopub.execute_input":"2025-04-16T01:42:57.244631Z","iopub.status.idle":"2025-04-16T01:42:57.320299Z","shell.execute_reply.started":"2025-04-16T01:42:57.244611Z","shell.execute_reply":"2025-04-16T01:42:57.319599Z"}},"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-04-16T01:42:57.321902Z","iopub.execute_input":"2025-04-16T01:42:57.322131Z","iopub.status.idle":"2025-04-16T01:42:57.328054Z","shell.execute_reply.started":"2025-04-16T01:42:57.322114Z","shell.execute_reply":"2025-04-16T01:42:57.327388Z"}},"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):\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        if preload: \n            self.data_list, self.label_list = [], []\n            for batch in self.batches: \n                data, label = self.load_every(batch)\n                self.data_list.append(data)\n                if label is not None:\n                    self.label_list.append(label)\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)[:, :, ::self.sample_ratio, :]\n        data = data.astype('float32')\n        if len(batch) > 1:\n            label_path = batch[1]\n            label = np.load(label_path)\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            data = self.data_list[batch_idx][sample_idx]\n            label = self.label_list[batch_idx][sample_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        if self.transform_data:\n            data = self.transform_data(data)\n        if self.transform_label and label is not None:\n            label = self.transform_label(label)\n        return data, label if label is not None else np.array([])\n        \n    def __len__(self):\n        return len(self.batches) * self.file_size\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T01:42:57.328638Z","iopub.execute_input":"2025-04-16T01:42:57.328819Z","iopub.status.idle":"2025-04-16T01:42:57.350140Z","shell.execute_reply.started":"2025-04-16T01:42:57.328805Z","shell.execute_reply":"2025-04-16T01:42:57.349510Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_inputs = [\n    f\n    for f in\n    Path('/kaggle/input/waveform-inversion/train_samples').rglob('*.npy')\n    if ('seis' in f.stem) or ('data' in f.stem)\n]\n\ndef 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_outputs = inputs_files_to_output_files(all_inputs)\n\ntrain_inputs = [all_inputs[i] for i in range(0, len(all_inputs), 2)] # Sample every two\nvalid_inputs = [f for f in all_inputs if not f in train_inputs]\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-04-16T01:42:57.350766Z","iopub.execute_input":"2025-04-16T01:42:57.350982Z","iopub.status.idle":"2025-04-16T01:42:57.445603Z","shell.execute_reply.started":"2025-04-16T01:42:57.350967Z","shell.execute_reply":"2025-04-16T01:42:57.445157Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transform_data = Compose([\n    T.LogTransform(k=1),\n    T.MinMaxNormalize(T.log_transform(-61, k=1), T.log_transform(120, k=1))\n])\ntransform_label = Compose([\n    T.MinMaxNormalize(2000, 6000)\n])\ndataset = FWIDataset(train_inputs[:1], train_outputs[:1], transform_data=transform_data, transform_label=transform_label, file_size=500)\ndata, label = dataset[0]\nprint(data.shape)\nprint(label is None)\n\ndel dataset, data, label\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T01:42:57.446229Z","iopub.execute_input":"2025-04-16T01:42:57.446486Z","iopub.status.idle":"2025-04-16T01:43:01.739388Z","shell.execute_reply.started":"2025-04-16T01:42:57.446452Z","shell.execute_reply":"2025-04-16T01:43:01.738815Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ctx_dict = {\n    \"flatvel-a\": {\n        \"data_min\": -26.95,\n        \"data_max\": 52.77,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\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        \"label_min\": 1500,\n        \"label_max\": 4500,\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-b\": {\n        \"data_min\": -27.17,\n        \"data_max\": 56.05,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\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        \"label_min\": 1500,\n        \"label_max\": 4500,\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        \"label_min\": 1500,\n        \"label_max\": 4500,\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        \"label_min\": 1500,\n        \"label_max\": 4500,\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        \"label_min\": 1500,\n        \"label_max\": 4500,\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-b\": {\n        \"data_min\": -24.93,\n        \"data_max\": 50.98,\n        \"label_min\": 1500,\n        \"label_max\": 4500,\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        \"label_min\": 1500,\n        \"label_max\": 4500,\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        \"label_min\": 1500,\n        \"label_max\": 4500,\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},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ctx = ctx_dict['flatfault-b']\n\nlog_data_min = T.log_transform(ctx['data_min'], k=args.k)\nlog_data_max = T.log_transform(ctx['data_max'], k=args.k)\ntransform_data = Compose([\n    T.LogTransform(k=args.k),\n    T.MinMaxNormalize(log_data_min, log_data_max)\n])\ntransform_label = Compose([\n    T.MinMaxNormalize(ctx['label_min'], ctx['label_max'])\n])\n\nif not args.run_train:\n    train_inputs = train_inputs[:10]\n    train_outputs = train_outputs[:10]\n\n    valid_inputs = valid_inputs[:10]\n    valid_outputs = valid_outputs[:10]\n    \ndataset_train = FWIDataset(\n        train_inputs,\n        train_outputs,\n        preload=True,\n        sample_ratio=args.sample_temporal,\n        file_size=ctx['file_size'],\n        transform_data=transform_data,\n        transform_label=transform_label\n    )\n\ndataset_valid = FWIDataset(\n    valid_inputs,\n    valid_outputs,\n    preload=True,\n    sample_ratio=args.sample_temporal,\n    file_size=ctx['file_size'],\n    transform_data=transform_data,\n    transform_label=transform_label\n)\n\ntrain_sampler = RandomSampler(dataset_train)\nvalid_sampler = RandomSampler(dataset_valid)\n\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-04-16T01:43:01.740086Z","iopub.execute_input":"2025-04-16T01:43:01.740318Z","iopub.status.idle":"2025-04-16T01:44:09.836525Z","shell.execute_reply.started":"2025-04-16T01:43:01.740289Z","shell.execute_reply":"2025-04-16T01:44:09.835899Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nmodel = network.model_dict[args.model](upsample_mode=args.up_mode, \n        sample_spatial=args.sample_spatial, sample_temporal=args.sample_temporal).to(device)\nmodel_d = network.model_dict[args.model_d]().to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T01:44:09.837211Z","iopub.execute_input":"2025-04-16T01:44:09.837397Z","iopub.status.idle":"2025-04-16T01:44:10.386562Z","shell.execute_reply.started":"2025-04-16T01:44:09.837382Z","shell.execute_reply":"2025-04-16T01:44:10.386036Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"l1loss = nn.L1Loss()\nl2loss = nn.MSELoss()\n\ndef criterion_g(pred, gt, model_d=None):\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    if model_d is not None:\n        loss_adv = -torch.mean(model_d(pred))\n        loss += args.lambda_adv * loss_adv\n    return loss, loss_g1v, loss_g2v\ncriterion_d = utils.Wasserstein_GP(device, args.lambda_gp)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T01:44:10.388669Z","iopub.execute_input":"2025-04-16T01:44:10.388896Z","iopub.status.idle":"2025-04-16T01:44:10.393409Z","shell.execute_reply.started":"2025-04-16T01:44:10.388880Z","shell.execute_reply":"2025-04-16T01:44:10.392866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Scale lr according to effective batch size\nlr_g = args.lr_g\nlr_d = args.lr_d\noptimizer_g = torch.optim.AdamW(model.parameters(), lr=lr_g, betas=(0, 0.9), weight_decay=args.weight_decay)\noptimizer_d = torch.optim.AdamW(model_d.parameters(), lr=lr_d, betas=(0, 0.9), weight_decay=args.weight_decay)\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]\nlr_schedulers = [WarmupMultiStepLR(\n    optimizer, milestones=lr_milestones, gamma=args.lr_gamma,\n    warmup_iters=warmup_iters, warmup_factor=1e-5) for optimizer in [optimizer_g, optimizer_d]]\n\nmodel_without_ddp = model\nmodel_d_without_ddp = model_d\n\nif args.resume:\n    checkpoint = torch.load(args.resume, map_location='cpu')\n    model_without_ddp.load_state_dict(network.replace_legacy(checkpoint['model']))\n    model_d_without_ddp.load_state_dict(network.replace_legacy(checkpoint['model_d']))\n    optimizer_g.load_state_dict(checkpoint['optimizer_g'])\n    optimizer_d.load_state_dict(checkpoint['optimizer_d'])\n    args.start_epoch = checkpoint['epoch'] + 1\n    step = checkpoint['step']\n    for i in range(len(lr_schedulers)):\n        lr_schedulers[i].load_state_dict(checkpoint['lr_schedulers'][i])\n    for lr_scheduler in lr_schedulers:\n        lr_scheduler.milestones = lr_milestones\n\n\nif args.pretrained and args.run_train:\n    checkpoint = torch.load(args.pretrain_path)\n    model_without_ddp.load_state_dict(network.replace_legacy(checkpoint['model']))\n    model_d_without_ddp.load_state_dict(network.replace_legacy(checkpoint['model_d']))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T01:44:10.394090Z","iopub.execute_input":"2025-04-16T01:44:10.394250Z","iopub.status.idle":"2025-04-16T01:44:12.818974Z","shell.execute_reply.started":"2025-04-16T01:44:10.394237Z","shell.execute_reply":"2025-04-16T01:44:12.818384Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"step = 0\n\ndef train_one_epoch(model, model_d, criterion_g, criterion_d, optimizer_g, optimizer_d, \n                    lr_schedulers, dataloader, device, epoch, print_freq, n_critic=5):\n    global step\n    model.train()\n    model_d.train()\n\n    # Logger setup\n    metric_logger = utils.MetricLogger(delimiter='  ')\n    metric_logger.add_meter('lr_g', utils.SmoothedValue(window_size=1, fmt='{value}'))\n    metric_logger.add_meter('lr_d', 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    itr = 0 # step in this epoch\n    max_itr = len(dataloader)\n\n\n    for data, label in metric_logger.log_every(dataloader, print_freq, header):\n        start_time = time.time()\n        data, label = data.to(device), label.to(device)\n\n        # Update discribminator first\n        optimizer_d.zero_grad()\n        with torch.no_grad():\n            pred = model(data)\n        loss_d, loss_diff, loss_gp = criterion_d(label, pred, model_d)\n        loss_d.backward()\n        optimizer_d.step()\n        metric_logger.update(loss_diff=loss_diff, loss_gp=loss_gp)\n\n        # Update generator occasionally \n        if ((itr + 1) % n_critic == 0) or (itr == max_itr - 1):\n            optimizer_g.zero_grad()\n            pred = model(data)\n            loss_g, loss_g1v, loss_g2v = criterion_g(pred, label, model_d)\n            loss_g.backward()\n            optimizer_g.step()\n            metric_logger.update(loss_g1v=loss_g1v, loss_g2v=loss_g2v)\n\n        batch_size = data.shape[0]\n        metric_logger.update(lr_g=optimizer_g.param_groups[0]['lr'],\n                            lr_d=optimizer_d.param_groups[0]['lr'])\n        metric_logger.meters['samples/s'].update(batch_size / (time.time() - start_time))\n        step += 1\n        itr += 1\n        for lr_scheduler in lr_schedulers:\n            lr_scheduler.step()\n\n\ndef evaluate(model, criterion, dataloader, device, epoch):\n    model.eval()\n    metric_logger = utils.MetricLogger(delimiter='  ')\n    header = 'Test:'\n    \n    all_outputs = []\n    all_labels = []\n    with torch.no_grad():\n        for data, label in metric_logger.log_every(dataloader, 20, header):\n            data = data.to(device, non_blocking=True)\n            label = label.to(device, non_blocking=True)\n            pred = model(data)\n            loss, loss_g1v, loss_g2v = criterion(pred, label)\n            metric_logger.update(loss=loss.item(), \n                                 loss_g1v=loss_g1v.item(), loss_g2v=loss_g2v.item())\n\n            all_outputs.append(pred.cpu())\n            all_labels.append(label.cpu())\n\n\n    all_output = torch.concat(all_outputs, axis=0)\n    all_label = torch.concat(all_labels, axis=0)\n    all_output = T.minmax_denormalize(all_output, ctx['label_min'], ctx['label_max'])\n    all_label = T.minmax_denormalize(all_label, ctx['label_min'], ctx['label_max'])\n    l1loss_eval = l1loss(all_output, all_label)\n    # Gather the stats from all processes\n    metric_logger.synchronize_between_processes()\n    print(' * Loss {loss.global_avg:.8f}, L1_loss {l1loss_eval:.8f} \\n'.format(loss=metric_logger.loss, l1loss_eval=l1loss_eval))\n    \n    if epoch % 4 == 0:\n        y = all_label[0, 0].detach().cpu()\n        y_pred = all_output[0, 0].detach().cpu()\n        \n        fig, ax = plt.subplots(1, 2, figsize=(5, 2.5))\n        fig.suptitle(f'Epoch {epoch} | Valid: {l1loss_eval:.5f}')\n        ax[0].imshow(y)\n        ax[1].imshow(y_pred)\n        plt.show()\n\n    return metric_logger.loss.global_avg, l1loss_eval","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T01:44:12.819674Z","iopub.execute_input":"2025-04-16T01:44:12.819951Z","iopub.status.idle":"2025-04-16T01:44:12.832406Z","shell.execute_reply.started":"2025-04-16T01:44:12.819928Z","shell.execute_reply":"2025-04-16T01:44:12.831793Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if args.run_train:\n\n    print('Start training')\n    start_time = time.time()\n    args.epochs = args.epoch_block * args.num_block\n    \n    best_loss = 5000\n    for epoch in range(args.start_epoch, args.epochs):\n        train_one_epoch(model, model_d, criterion_g, criterion_d, optimizer_g, optimizer_d,\n                        lr_schedulers, dataloader_train, device, epoch, \n                        args.print_freq, args.n_critic)\n        loss_global_avg, l1loss_eval = evaluate(model, criterion_g, dataloader_valid, device, epoch)\n        checkpoint = {\n            'model': model_without_ddp.state_dict(),\n            'model_d': model_d_without_ddp.state_dict(),\n            'optimizer_g': optimizer_g.state_dict(),\n            'optimizer_d': optimizer_d.state_dict(),\n            'lr_schedulers': [scheduler.state_dict() for scheduler in lr_schedulers],\n            'epoch': epoch,\n            'step': step,\n            'args': args}\n    \n        if l1loss_eval < best_loss:\n            utils.save_on_master(\n                checkpoint,\n                os.path.join(args.output_path, 'best_model.pth'))\n            best_loss = l1loss_eval\n        \n        utils.save_on_master(\n            checkpoint,\n            os.path.join(args.output_path, 'checkpoint.pth'))\n        # Save checkpoint every epoch block\n        if args.output_path and (epoch + 1) % args.epoch_block == 0:\n            utils.save_on_master(\n                checkpoint,\n                os.path.join(args.output_path, 'model_{}.pth'.format(epoch + 1)))\n    \n    \n    total_time = time.time() - start_time\n    total_time_str = str(datetime.timedelta(seconds=int(total_time)))\n    print('Training time {}'.format(total_time_str))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T01:44:12.833217Z","iopub.execute_input":"2025-04-16T01:44:12.833475Z","execution_failed":"2025-04-16T01:59:50.692Z"}},"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-04-16T01:59:50.693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\ntest_files = list(Path('/kaggle/input/waveform-inversion/test').glob('*.npy'))\nlen(test_files)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-04-16T01:59:50.693Z"}},"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-04-16T01:59:50.693Z"}},"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-04-16T01:59:50.693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ctx_test = ctx_dict['flatfault-b']\n\n\nlog_data_min = T.log_transform(ctx_test['data_min'], k=args.k)\nlog_data_max = T.log_transform(ctx_test['data_max'], k=args.k)\ntransform_data = Compose([\n    T.LogTransform(k=args.k),\n    T.MinMaxNormalize(log_data_min, log_data_max)\n])\n\n\nds = TestDataset(test_files, transform_data)\ndl = DataLoader(ds, batch_size=8, num_workers=4, pin_memory=True)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-04-16T01:59:50.693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"checkpoint = torch.load('/kaggle/input/gwi-model/cv201_best_model.pth')\n\nmodel_without_ddp.load_state_dict(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            outputs = model(inputs)\n\n        y_preds = outputs[:, 0].cpu().numpy()\n        y_preds = T.minmax_denormalize(y_preds, ctx_test['label_min'], ctx_test['label_max'])\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-04-16T01:59:50.693Z"}},"outputs":[],"execution_count":null}]}