{"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":11334027,"sourceType":"datasetVersion","datasetId":7089850},{"sourceId":234861093,"sourceType":"kernelVersion"},{"sourceId":234861398,"sourceType":"kernelVersion"},{"sourceId":234861685,"sourceType":"kernelVersion"},{"sourceId":234862033,"sourceType":"kernelVersion"},{"sourceId":349771,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":291080,"modelId":311767}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Importing our libraries","metadata":{}},{"cell_type":"code","source":"import os\nimport glob\nimport gc\nimport csv\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom typing import Tuple, Optional, Union, Callable, List\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\n\nimport time\nimport datetime\nimport random\n\nimport pathlib, html\nfrom IPython.display import HTML\n\nimport openfwi_utils as utils\nimport openfwi_network as network\nimport openfwi_transforms as T\nfrom openfwi_scheduler import WarmupMultiStepLR","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:57:55.297023Z","iopub.execute_input":"2025-04-21T10:57:55.297564Z","iopub.status.idle":"2025-04-21T10:58:14.343895Z","shell.execute_reply.started":"2025-04-21T10:57:55.297534Z","shell.execute_reply":"2025-04-21T10:58:14.343023Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Root Paths for the train and test directory","metadata":{}},{"cell_type":"code","source":"# Root Paths\nBASE_DIR = '/kaggle/input/waveform-inversion'\nTRAIN_DIR = os.path.join(BASE_DIR, 'train_samples')\nTEST_DIR = os.path.join(BASE_DIR, 'test')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:16.510663Z","iopub.execute_input":"2025-04-21T10:58:16.511460Z","iopub.status.idle":"2025-04-21T10:58:16.515085Z","shell.execute_reply.started":"2025-04-21T10:58:16.511434Z","shell.execute_reply":"2025-04-21T10:58:16.514362Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Helping Functions","metadata":{}},{"cell_type":"code","source":"def collect_input_files(\n    data_dir: str = TRAIN_DIR\n) -> list:\n    \"\"\"\n    Recursively search for .npy files in data_dir that contain 'seis' or 'data' in their filename.\n\n    Parameters:\n    ----------\n    data_dir : str, default = TRAIN_DIR\n        The data_dir that need searching.\n\n    Returns:\n    -------\n    list\n        A list contains the paths of all input files in our data.\n    \"\"\"\n    return [\n        f \n        for f in Path(data_dir).rglob(\"*.npy\")\n        if (\"seis\" in f.stem) \n        or (\"data\" in f.stem)\n    ]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:17.169987Z","iopub.execute_input":"2025-04-21T10:58:17.170692Z","iopub.status.idle":"2025-04-21T10:58:17.174603Z","shell.execute_reply.started":"2025-04-21T10:58:17.170643Z","shell.execute_reply":"2025-04-21T10:58:17.173960Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def map_input_to_output(\n    input_files: list\n) -> list:\n    \"\"\"\n    Map each input file to its corresponding output file by replacing keywords.\n\n    Parameters:\n    ----------\n    input_files : list\n        The list that contains the paths of all input files in our data.\n\n    Returns:\n    -------\n    list\n        A list contains the paths of all output files in our data.\n    \"\"\"\n    return [\n        Path(str(f).replace(\"seis\", \"vel\").replace(\"data\", \"model\"))\n        for f in input_files\n    ]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:17.434012Z","iopub.execute_input":"2025-04-21T10:58:17.434648Z","iopub.status.idle":"2025-04-21T10:58:17.438412Z","shell.execute_reply.started":"2025-04-21T10:58:17.434625Z","shell.execute_reply":"2025-04-21T10:58:17.437713Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inspecting our Data","metadata":{}},{"cell_type":"code","source":"random_model = np.load('/kaggle/input/waveform-inversion/train_samples/FlatFault_A/seis4_1_0.npy')\nrandom_model.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:17.974419Z","iopub.execute_input":"2025-04-21T10:58:17.974696Z","iopub.status.idle":"2025-04-21T10:58:22.710602Z","shell.execute_reply.started":"2025-04-21T10:58:17.974656Z","shell.execute_reply":"2025-04-21T10:58:22.709939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"batch_size : \", random_model.shape[0])\nprint(\"num_sources : \", random_model.shape[1])\nprint(\"time_steps : \", random_model.shape[2])\nprint('num_receivers: ',random_model.shape[3])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:22.711710Z","iopub.execute_input":"2025-04-21T10:58:22.711932Z","iopub.status.idle":"2025-04-21T10:58:22.716700Z","shell.execute_reply.started":"2025-04-21T10:58:22.711915Z","shell.execute_reply":"2025-04-21T10:58:22.715844Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Example (batch) 0\n\n── Source 0\n\n── Time 0: [Receiver 0, Receiver 1, …, Receiver 69]\n\n── Time 1: [Receiver 0, Receiver 1, …, Receiver 69]\n\n── …\n\n── Time 999: [Receiver 0, …, Receiver 69]\n\n── Source 1\n\n── …\n\n── …\n\n── Source 4\n\nExample (batch) 1\n\n── …\n\n…\n\nExample (batch) 499\n\n── …\n\n…","metadata":{}},{"cell_type":"code","source":"random_velocity = np.load('/kaggle/input/waveform-inversion/train_samples/FlatFault_A/vel4_1_0.npy')\nrandom_velocity.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:22.717399Z","iopub.execute_input":"2025-04-21T10:58:22.717661Z","iopub.status.idle":"2025-04-21T10:58:22.789286Z","shell.execute_reply.started":"2025-04-21T10:58:22.717644Z","shell.execute_reply":"2025-04-21T10:58:22.788644Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"batch_size : \", random_velocity.shape[0])\nprint(\"num_sources : \", random_velocity.shape[1])\nprint(\"height : \", random_velocity.shape[2])\nprint('width: ',random_velocity.shape[3])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:22.790876Z","iopub.execute_input":"2025-04-21T10:58:22.791332Z","iopub.status.idle":"2025-04-21T10:58:22.795552Z","shell.execute_reply.started":"2025-04-21T10:58:22.791311Z","shell.execute_reply":"2025-04-21T10:58:22.794799Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Folder Explorer with Inline Explanations","metadata":{}},{"cell_type":"code","source":"ROOT = next(pathlib.Path(\"/kaggle/input\").rglob(\"train_samples\"), None)\nassert ROOT and ROOT.is_dir(), \"❌  /train_samples not found\"\n\n# ---------- helpers ----------\ndef human_bytes(n):\n    units = ['B','KB','MB','GB','TB']; i = 0\n    while n >= 1024 and i < len(units)-1: n /= 1024; i += 1\n    return f\"{n:.1f}{units[i]}\"\n\ndef npy_info(p):\n    arr = np.load(p, mmap_mode='r'); shp, dt = arr.shape, arr.dtype; arr._mmap.close()\n    return shp, dt\n\ndef folder_desc(name):\n    if name.startswith(\"FlatVel\"):   base = \"Vel—flat, gently layered\"; \n    elif name.startswith(\"CurveVel\"): base = \"Vel—curved/folded layers\";\n    elif name.startswith(\"FlatFault\"): base = \"Fault—flat layers with breaks\";\n    elif name.startswith(\"CurveFault\"):base = \"Fault—curved layers with breaks\";\n    elif name.startswith(\"Style\"):   base = \"Style—random texture pattern\";\n    else:                             base = \"Unknown pattern\"\n    level = \"A simpler\" if name.endswith(\"_A\") else \"B more complex\" if name.endswith(\"_B\") else \"\"\n    return f\"{base} ({level})\".strip()\n\nKIND_INFO = {\n    \"Data\":  \"4‑D seismic waveforms (sources × time × receivers)\",\n    \"Model\": \"2‑D velocity ground truth\",\n    \"Seis\":  \"Seismic recordings (prefix layout)\",\n    \"Vel\":   \"Velocity maps (prefix layout)\",\n    \"Files\": \"Misc. .npy files\"\n}\n\n# ---------- gather ----------\ntree = {}\nfor fld in sorted(ROOT.iterdir()):\n    if not fld.is_dir(): continue\n    kinds={}\n    def add(k, paths):\n        items=[(p.name, f\"shape={shp}, dtype={dt}, size={human_bytes(p.stat().st_size)}\")\n               for p in paths if (shp:=npy_info(p)[0]) or True for dt in [npy_info(p)[1]]]\n        if items: kinds[k]={\"count\":len(items),\"files\":items}\n    add(\"Data\",  (fld/\"data\").glob(\"*.npy\")  if (fld/\"data\").is_dir()  else [])\n    add(\"Model\", (fld/\"model\").glob(\"*.npy\") if (fld/\"model\").is_dir() else [])\n    add(\"Seis\", fld.glob(\"seis*.npy\"));  add(\"Vel\", fld.glob(\"vel*.npy\"))\n    if not kinds: add(\"Files\", fld.glob(\"*.npy\"))\n    tree[fld.name]={\"desc\":folder_desc(fld.name),\"kinds\":kinds}\n\n# ---------- html ----------\nhtml_parts=[\"\"\"\n<style>\n:root{--accent:#136efd;--bg:#fafafa;--border:#ddd;--font:system-ui,sans-serif}\n#exp{font-family:var(--font);background:var(--bg);padding:1rem;border:1px solid var(--border);\n     border-radius:8px;max-width:1050px;margin:auto}\n#exp h2{margin:0 0 1rem;text-align:center;font-size:1.35rem}\n.folder{border:1px solid var(--border);border-radius:6px;margin:.7rem 0;background:#fff}\n.folder summary{cursor:pointer;display:flex;gap:.6rem;align-items:center;padding:.5rem .75rem}\n.fname{font-weight:600}\n.fdesc{font-size:.8rem;color:#555}\n.kind-line{margin-left:1.3rem;margin-top:.45rem}\n.kind-tag{background:var(--accent);color:#fff;border-radius:4px;padding:.18rem .55rem;font-size:.75rem}\n.kind-info{font-size:.78rem;color:#222;margin-left:.45rem}\n.count{color:#555;font-size:.78rem;margin-left:.25rem}\n.file-list{margin:.25rem 0 .8rem 2.5rem;font-size:.85rem;color:#333}\n.file-list li{margin:.04rem 0;list-style-type:disc}\n</style>\n<div id=\"exp\">\n  <h2>Training .npy Explorer — Inline Explanations</h2>\n\"\"\"]\n\nfor name,info in tree.items():\n    kinds=info[\"kinds\"]; desc=html.escape(info[\"desc\"])\n    html_parts.append(\"<details class='folder'>\")\n    html_parts.append(f\"<summary><span class='fname'>{html.escape(name)}</span>\"\n                      f\"<span class='fdesc'>— {desc}</span></summary>\")\n    for kind,meta in kinds.items():\n        html_parts.append(f\"<div class='kind-line'><span class='kind-tag'>{kind}</span>\"\n                          f\"<span class='count'>({meta['count']})</span>\"\n                          f\"<span class='kind-info'>{html.escape(KIND_INFO.get(kind,''))}</span></div>\")\n        html_parts.append(\"<ul class='file-list'>\")\n        for fn,tip in meta[\"files\"]:\n            html_parts.append(f\"<li title='{html.escape(tip)}'>{html.escape(fn)}</li>\")\n        html_parts.append(\"</ul>\")\n    html_parts.append(\"</details>\")\n\nhtml_parts.append(\"</div>\")\nHTML(\"\".join(html_parts))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:22.796387Z","iopub.execute_input":"2025-04-21T10:58:22.796603Z","iopub.status.idle":"2025-04-21T10:58:23.890129Z","shell.execute_reply.started":"2025-04-21T10:58:22.796587Z","shell.execute_reply":"2025-04-21T10:58:23.889404Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualizing our Data","metadata":{}},{"cell_type":"code","source":"batch_size, num_sources, time_steps, num_receivers = random_model.shape\nprint(f\"Data form: {random_model.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:23.890934Z","iopub.execute_input":"2025-04-21T10:58:23.891136Z","iopub.status.idle":"2025-04-21T10:58:23.895132Z","shell.execute_reply.started":"2025-04-21T10:58:23.891119Z","shell.execute_reply":"2025-04-21T10:58:23.894534Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2D Format","metadata":{}},{"cell_type":"code","source":"selected_batch = 0\n\nfig, axes = plt.subplots(nrows=num_sources, figsize=(20, 25), sharex=True, sharey=True)\nfor source in range(num_sources):\n    data = random_model[selected_batch, source, :, :].T\n    im = axes[source].imshow(\n        data,\n        cmap='seismic',\n        aspect='auto',\n        extent=[0, time_steps, num_receivers, 0],\n        vmin=-np.abs(data).max(),\n        vmax=np.abs(data).max()\n    )\n    axes[source].set_title(f'Source {source}', pad=15)\n    axes[source].set_ylabel('Receiver number')\n    axes[source].set_xlabel('Time step' if source == num_sources-1 else '')\nplt.tight_layout()\nplt.show();","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:32.589761Z","iopub.execute_input":"2025-04-21T10:58:32.590580Z","iopub.status.idle":"2025-04-21T10:58:34.267934Z","shell.execute_reply.started":"2025-04-21T10:58:32.590549Z","shell.execute_reply":"2025-04-21T10:58:34.267157Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3D Format","metadata":{}},{"cell_type":"code","source":"selected_batch = 0\nselected_source = 2\n\ndata = random_model[selected_batch, selected_source, :, :]\nX, Y = np.meshgrid(np.arange(num_receivers), np.arange(time_steps))\nfig = plt.figure(figsize=(18, 10))\nax = fig.add_subplot(111, projection='3d')\nsurf = ax.plot_surface(X, Y, data,cmap='seismic', rstride=1, cstride=1)\nax.set_title(f'3D visualization of seismic data (batch={selected_batch}, source={selected_source})')\nax.set_xlabel('Receiver number')\nax.set_ylabel('Time step')\nax.set_zlabel('The amplitude')\nfig.colorbar(surf, ax=ax, shrink=0.5, label='The amplitude')\nplt.show();","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:34.269283Z","iopub.execute_input":"2025-04-21T10:58:34.269721Z","iopub.status.idle":"2025-04-21T10:58:37.721500Z","shell.execute_reply.started":"2025-04-21T10:58:34.269695Z","shell.execute_reply":"2025-04-21T10:58:37.720871Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del random_model, random_velocity\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:37.722471Z","iopub.execute_input":"2025-04-21T10:58:37.723015Z","iopub.status.idle":"2025-04-21T10:58:38.081240Z","shell.execute_reply.started":"2025-04-21T10:58:37.722991Z","shell.execute_reply":"2025-04-21T10:58:38.080542Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Building our model","metadata":{}},{"cell_type":"markdown","source":"## Arguments","metadata":{}},{"cell_type":"code","source":"args = {\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, # weight decay coefficient in AdamW Optimizer\n    \"batch_size\": 64, # Batch Size in DataLoader objects\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, # How many subprocesses to use in loading data\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}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:58:38.082772Z","iopub.execute_input":"2025-04-21T10:58:38.083019Z","iopub.status.idle":"2025-04-21T10:58:38.097598Z","shell.execute_reply.started":"2025-04-21T10:58:38.083000Z","shell.execute_reply":"2025-04-21T10:58:38.097013Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_torch(\n    seed_value: int\n) -> None:\n    \"\"\"\n    Controlling a unified seed value for Python, NumPy, and PyTorch (CPU, GPU).\n\n    Parameters:\n    ----------\n    seed_value : int\n        The unified random seed value.\n    \"\"\"\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-21T10:58:38.098446Z","iopub.execute_input":"2025-04-21T10:58:38.098727Z","iopub.status.idle":"2025-04-21T10:58:38.167733Z","shell.execute_reply.started":"2025-04-21T10:58:38.098703Z","shell.execute_reply":"2025-04-21T10:58:38.167215Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Managing Dataset","metadata":{}},{"cell_type":"code","source":"# Data config from OpenFWI Repository\n# https://github.com/lanl/OpenFWI/blob/main/dataset_config.json\ndata_config = {\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,"execution":{"iopub.status.busy":"2025-04-21T10:58:49.282088Z","iopub.execute_input":"2025-04-21T10:58:49.282694Z","iopub.status.idle":"2025-04-21T10:58:49.292963Z","shell.execute_reply.started":"2025-04-21T10:58:49.282656Z","shell.execute_reply":"2025-04-21T10:58:49.292334Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SeismicDataset(Dataset):\n    \"\"\"\n    PyTorch Dataset for handling seismic data with multiple examples per file.\n\n    This dataset supports loading seismic data from `.npy` files, optionally preloading\n    into memory, and applying transformations to both input data and labels.\n\n    Attributes\n    ----------\n    in_files : list of str\n        List of paths to input seismic data files (`.npy` format).\n    out_files : list of str or None\n        List of paths to corresponding label files (`.npy` format).\n        If `None`, labels are omitted.\n    preload : bool, optional (default=True)\n        If `True`, loads all data into memory during initialization for faster access.\n    sample_ratio : int, optional (default=1)\n        Downsampling ratio applied along the seismic time axis\n        (e.g., `2` for half resolution).\n    examples_per_file : int, optional (default=500)\n        Number of samples (examples) contained in each `.npy` file.\n    transform_data : torchvision.transforms.Compose, optional (default=None)\n        Transformations applied to the input seismic data.\n    transform_label : torchvision.transforms.Compose, optional (default=None)\n        Transformations applied to the labels.\n    \"\"\"\n    def __init__(\n        self,\n        in_files: list,\n        out_files: list,\n        preload: bool = True,\n        sample_ratio: int = 1,\n        examples_per_file: int = 500,\n        transform_data = None,\n        transform_label = None\n    ) -> None:\n        \"\"\"\n        Initialize the dataset.\n\n        Parameters\n        ----------\n        in_files : list of str\n            Paths to input seismic data files (`.npy` format).\n        out_files : list of str or None\n            Paths to label files (`.npy` format).\n            If `None`, labels are omitted.\n        preload : bool, optional (default=True)\n            Preload all data into memory for faster access during training.\n        sample_ratio : int, optional (default=1)\n            Downsampling ratio for the seismic time axis.\n        examples_per_file : int, optional (default=500)\n            Number of examples per `.npy` file.\n        transform_data : torchvision.transforms.Compose, optional (default=None)\n            Transforms for input data.\n        transform_label : torchvision.transforms.Compose, optional (default=None)\n            Transforms for labels.\n        \"\"\"\n        self.preload = preload\n        self.sample_ratio = sample_ratio\n        self.examples_per_file = examples_per_file\n        self.transform_data = transform_data\n        self.transform_label = transform_label\n        if out_files is not None:\n            self.batches = [\n                str(in_files[i])\n                +'&'\n                +str(out_files[i]) for i in range(len(in_files))\n            ]\n        else:\n            self.batches = [str(in_files[i]) for i in range(len(in_files))]\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    def load_every(\n        self,\n        batch: str\n    ) -> Tuple[np.ndarray, Optional[np.ndarray]]:\n        \"\"\"\n        Load a batch of data and labels from file paths.\n\n        Parameters\n        ----------\n        batch : str\n            String formatted as `\"data_path&label_path\"` or `\"data_path\"` (if no labels).\n\n        Returns\n        -------\n        data : np.ndarray\n            Loaded seismic data.\n        label : np.ndarray or None\n            Loaded labels (if available), otherwise `None`.\n        \"\"\"\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 __len__(self) -> int:\n        \"\"\"Total number of samples in the dataset.\"\"\"\n        return len(self.batches) * self.examples_per_file\n\n    def __getitem__(\n        self,\n        idx: int\n    ) -> Tuple[torch.Tensor, Union[torch.Tensor, np.ndarray]]:\n        \"\"\"\n        Retrieve a single sample by index.\n\n        Parameters\n        ----------\n        idx : int\n            Index of the sample to fetch.\n\n        Returns\n        -------\n        data : torch.Tensor\n            Transformed seismic data sample.\n        label : torch.Tensor or np.ndarray\n            Transformed label (if available), else an empty array.\n        \"\"\"\n        batch_idx, sample_idx = (\n            idx // self.examples_per_file,\n            idx % self.examples_per_file\n        )\n        if self.preload:\n            data = self.data_list[batch_idx][sample_idx]\n            label = self.label_list[batch_idx][sample_idx] if len(\n                self.label_list\n            ) != 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([])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:59:21.365508Z","iopub.execute_input":"2025-04-21T10:59:21.366101Z","iopub.status.idle":"2025-04-21T10:59:21.376766Z","shell.execute_reply.started":"2025-04-21T10:59:21.366080Z","shell.execute_reply":"2025-04-21T10:59:21.376150Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"inputs_all = collect_input_files(TRAIN_DIR)\noutputs_all = map_input_to_output(inputs_all)\n\n# Check all output files exist\nassert all(f.exists() for f in outputs_all)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:59:21.729971Z","iopub.execute_input":"2025-04-21T10:59:21.730276Z","iopub.status.idle":"2025-04-21T10:59:21.774301Z","shell.execute_reply.started":"2025-04-21T10:59:21.730246Z","shell.execute_reply":"2025-04-21T10:59:21.773584Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Split dataset into training and validation based on sampling frequency\ntrain_inputs = [inputs_all[i] for i in range(0, len(inputs_all), 2)]\nvalid_inputs = [f for f in inputs_all if f not in train_inputs]\ntrain_outputs = map_input_to_output(train_inputs)\nvalid_outputs = map_input_to_output(valid_inputs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:59:22.651570Z","iopub.execute_input":"2025-04-21T10:59:22.651866Z","iopub.status.idle":"2025-04-21T10:59:22.656343Z","shell.execute_reply.started":"2025-04-21T10:59:22.651845Z","shell.execute_reply":"2025-04-21T10:59:22.655724Z"}},"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])\n\ntransform_label = Compose([\n    T.MinMaxNormalize(2000, 6000)\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:59:23.106526Z","iopub.execute_input":"2025-04-21T10:59:23.107234Z","iopub.status.idle":"2025-04-21T10:59:23.110950Z","shell.execute_reply.started":"2025-04-21T10:59:23.107211Z","shell.execute_reply":"2025-04-21T10:59:23.110323Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = SeismicDataset(\n    train_inputs[:1],\n    train_outputs[:1],\n    transform_data = transform_data,\n    transform_label = transform_label,\n    examples_per_file = 500\n)\ndata, label = dataset[0]\nprint(data.shape)\nprint(label is None)\nprint(label.shape)\n\ndel dataset, data, label\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:59:23.530393Z","iopub.execute_input":"2025-04-21T10:59:23.530924Z","iopub.status.idle":"2025-04-21T10:59:24.422640Z","shell.execute_reply.started":"2025-04-21T10:59:23.530901Z","shell.execute_reply":"2025-04-21T10:59:24.422042Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ctx = data_config['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\"])\n\ntransform_data = Compose([\n    T.LogTransform(k=args[\"k\"]),\n    T.MinMaxNormalize(log_data_min, log_data_max)\n])\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 = SeismicDataset(\n        train_inputs,\n        train_outputs,\n        preload=True,\n        sample_ratio=args[\"sample_temporal\"],\n        examples_per_file=ctx['file_size'],\n        transform_data=transform_data,\n        transform_label=transform_label\n    )\n\ndataset_valid = SeismicDataset(\n    valid_inputs,\n    valid_outputs,\n    preload=True,\n    sample_ratio=args[\"sample_temporal\"],\n    examples_per_file=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, # Dataset from which to load the data.\n    batch_size=args[\"batch_size\"], # How many samples per batch to load.\n    sampler=train_sampler, # Defines the strategy to draw samples from the dataset.\n    num_workers=args[\"workers\"], # How many subprocesses to use for data loading.\n    pin_memory=True, # Copy Tensors into device/CUDA pinned memory before returning them.\n    drop_last=True, # Set to True to drop the last incomplete batch.\n    collate_fn=default_collate, # Merges a list of samples to form a mini-batch of Tensor(s).\n    persistent_workers = True\n)\n\ndataloader_valid = DataLoader(\n    dataset_valid, # Dataset from which to load the data.\n    batch_size=args[\"batch_size\"], # How many samples per batch to load.\n    sampler=valid_sampler, # Defines the strategy to draw samples from the dataset.\n    num_workers=args[\"workers\"], # How many subprocesses to use for data loading.\n    pin_memory=True, # Copy Tensors into device/CUDA pinned memory before returning them.\n    collate_fn=default_collate, # Merges a list of samples to form a mini-batch of Tensor(s).\n    persistent_workers = True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T10:59:24.423614Z","iopub.execute_input":"2025-04-21T10:59:24.423879Z","iopub.status.idle":"2025-04-21T11:01:07.580061Z","shell.execute_reply.started":"2025-04-21T10:59:24.423862Z","shell.execute_reply":"2025-04-21T11:01:07.579439Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Building","metadata":{}},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nmodel = network.model_dict[args[\"model\"]](\n    upsample_mode=args[\"up_mode\"],\n    sample_spatial=args[\"sample_spatial\"],\n    sample_temporal=args[\"sample_temporal\"]\n).to(device)\n\nif args[\"model_d\"] is not None:\n    model_d = network.model_dict[args[\"model_d\"]]().to(device)\nelse:\n    model_d = None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T11:01:07.581271Z","iopub.execute_input":"2025-04-21T11:01:07.581903Z","iopub.status.idle":"2025-04-21T11:01:08.100267Z","shell.execute_reply.started":"2025-04-21T11:01:07.581882Z","shell.execute_reply":"2025-04-21T11:01:08.099407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Scale lr according to effective batch size\nlr_g = args[\"lr_g\"]\noptimizer_g = torch.optim.AdamW(\n    model.parameters(),\n    lr=lr_g,\n    betas=(0, 0.9),\n    weight_decay=args[\"weight_decay\"]\n)\n\n# Conditionally create discriminator optimizer\noptimizer_d = None\nif model_d is not None:\n    lr_d = args[\"lr_d\"]\n    optimizer_d = torch.optim.AdamW(\n        model_d.parameters(),\n        lr=lr_d,\n        betas=(0, 0.9),\n        weight_decay=args[\"weight_decay\"]\n    )\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\n# Create schedulers only for existing optimizers\noptimizers = [optimizer_g]\nif model_d is not None:\n    optimizers.append(optimizer_d)\n\nlr_schedulers = [\n    WarmupMultiStepLR(\n        optimizer,\n        milestones=lr_milestones,\n        gamma=args[\"lr_gamma\"],\n        warmup_iters=warmup_iters,\n        warmup_factor=1e-5\n    ) for optimizer in optimizers\n]\n\nmodel_without_ddp = model\nmodel_d_without_ddp = model_d if model_d is not None else None\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    \n    # Only load discriminator components if they exist\n    if model_d is not None:\n        model_d_without_ddp.load_state_dict(network.replace_legacy(checkpoint['model_d']))\n        optimizer_d.load_state_dict(checkpoint['optimizer_d'])\n    \n    optimizer_g.load_state_dict(checkpoint['optimizer_g'])\n    args.start_epoch = checkpoint['epoch'] + 1\n    step = checkpoint['step']\n    \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    \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    \n    # Only load discriminator if it exists\n    if model_d is not None:\n        model_d_without_ddp.load_state_dict(network.replace_legacy(checkpoint['model_d']))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T11:01:08.101026Z","iopub.execute_input":"2025-04-21T11:01:08.101240Z","iopub.status.idle":"2025-04-21T11:01:11.271563Z","shell.execute_reply.started":"2025-04-21T11:01:08.101224Z","shell.execute_reply":"2025-04-21T11:01:11.270879Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"l1loss = nn.L1Loss()\nl2loss = nn.MSELoss()\n\ndef criterion_g(\n    pred: torch.Tensor,\n    gt: torch.Tensor,\n    model_d: Optional[torch.nn.Module] = None\n) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n    \"\"\"\n    Generator loss function combining L1, L2, and adversarial losses.\n\n    Computes a weighted sum of L1 (MAE) and L2 (MSE) losses between predictions and ground truth.\n    If a discriminator model (`model_d`) is provided, includes an adversarial loss term to improve\n    generative performance.\n\n    Parameters\n    ----------\n    pred : torch.Tensor\n        Predicted output tensor from the generator model.\n    gt : torch.Tensor\n        Ground truth tensor with the same shape as `pred`.\n    model_d : torch.nn.Module, optional (default=None)\n        Discriminator model used for adversarial loss computation.\n        If `None`, adversarial loss is omitted.\n\n    Returns\n    -------\n    loss : torch.Tensor\n        Total generator loss (weighted sum of all components).\n    loss_g1v : torch.Tensor\n        L1 loss term (MAE) between `pred` and `gt`.\n    loss_g2v : torch.Tensor\n        L2 loss term (MSE) between `pred` and `gt`.\n\n    Notes\n    -----\n    - Loss weights (`lambda_g1v`, `lambda_g2v`, `lambda_adv`) are read from a global `args` dictionary.\n    - Adversarial loss is computed as `-torch.mean(model_d(pred))` to encourage the generator to fool the discriminator.\n    \"\"\"\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\n\nif model_d is not None:\n    criterion_d = utils.Wasserstein_GP(device, args[\"lambda_gp\"])\nelse:\n    criterion_d = None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T11:01:11.273093Z","iopub.execute_input":"2025-04-21T11:01:11.273351Z","iopub.status.idle":"2025-04-21T11:01:11.280332Z","shell.execute_reply.started":"2025-04-21T11:01:11.273333Z","shell.execute_reply":"2025-04-21T11:01:11.279593Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Functions","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(\n    model: torch.nn.Module,\n    criterion_g: Callable,\n    optimizer_g: torch.optim.Optimizer,\n    lr_schedulers: List[torch.optim.lr_scheduler._LRScheduler],\n    dataloader: torch.utils.data.DataLoader,\n    device: torch.device,\n    epoch: int,\n    print_freq: int,\n    model_d: Optional[torch.nn.Module] = None,\n    criterion_d: Optional[Callable] = None,\n    optimizer_d: Optional[torch.optim.Optimizer] = None,\n    n_critic:int = 5\n) -> None:\n    \"\"\"\n    Train the model for one epoch with optional adversarial training.\n\n    Supports two modes of operation:\n    1. Standard training (when model_d is None): Only updates the main model\n    2. GAN training (when model_d provided): Alternates between discriminator\n       and generator updates according to n_critic schedule\n\n    Parameters\n    ----------\n    model : torch.nn.Module\n        Primary model (generator in GAN mode) to be trained.\n    model_d : torch.nn.Module, optional\n        Discriminator model for adversarial training. If None, runs in standard mode.\n    criterion_g : Callable\n        Loss function for the main model. Signature:\n        - GAN mode: f(pred, gt, model_d) -> (total_loss, l1_loss, l2_loss)\n        - Standard mode: f(pred, gt) -> (total_loss, l1_loss, l2_loss)\n    criterion_d : Callable, optional\n        Discriminator loss function. Required if model_d is provided.\n        Signature: f(gt, pred, model_d) -> (total_loss, diff_loss, gp_loss)\n    optimizer_g : torch.optim.Optimizer\n        Optimizer for the main model.\n    optimizer_d : torch.optim.Optimizer, optional\n        Optimizer for discriminator. Required if model_d is provided.\n    lr_schedulers : List[torch.optim.lr_scheduler._LRScheduler]\n        Learning rate schedulers to update after each batch.\n    dataloader : torch.utils.data.DataLoader\n        DataLoader providing (data, label) batches.\n    device : torch.device\n        Device to run training on (e.g., 'cuda' or 'cpu').\n    epoch : int\n        Current epoch number (for logging purposes).\n    print_freq : int\n        Frequency in batches to print training metrics.\n    n_critic : int, default=5\n        In GAN mode: number of discriminator updates per generator update.\n\n    Notes\n    -----\n    - Uses a global `step` counter which should be initialized outside this function.\n    - For GAN training (model_d provided), all GAN-related parameters must be specified.\n    - Learning rate schedulers are stepped after every batch.\n    - Metrics tracked:\n        * lr_g: Learning rate of main model\n        * lr_d: Learning rate of discriminator (GAN mode only)\n        * samples/s: Processing speed\n        * loss_g1v: L1 loss component\n        * loss_g2v: L2 loss component\n        * loss_diff: Discriminator real/fake loss (GAN mode only)\n        * loss_gp: Gradient penalty loss (GAN mode only)\n\n    Examples\n    --------\n    >>> # Standard training\n    >>> train_one_epoch(\n    ...     model=generator,\n    ...     criterion_g=simple_loss,\n    ...     optimizer_g=opt_g,\n    ...     lr_schedulers=[scheduler_g],\n    ...     dataloader=train_loader,\n    ...     device='cuda',\n    ...     epoch=0,\n    ...     print_freq=100\n    ... )\n\n    >>> # GAN training\n    >>> train_one_epoch(\n    ...     model=generator,\n    ...     model_d=discriminator,\n    ...     criterion_g=gan_g_loss,\n    ...     criterion_d=gan_d_loss,\n    ...     optimizer_g=opt_g,\n    ...     optimizer_d=opt_d,\n    ...     lr_schedulers=[scheduler_g, scheduler_d],\n    ...     dataloader=train_loader,\n    ...     device='cuda',\n    ...     epoch=0,\n    ...     print_freq=100,\n    ...     n_critic=5\n    ... )\n    \"\"\"\n    global step\n    model.train()\n\n    # Logger setup\n    metric_logger = utils.MetricLogger(delimiter='  ')\n    metric_logger.add_meter(\n        'lr_g',\n        utils.SmoothedValue(window_size=1, fmt='{value}')\n    )\n    if model_d is not None:\n        model_d.train()\n        \n        # Validate GAN mode parameters\n        assert criterion_d is not None, \"criterion_d required for GAN training\"\n        assert optimizer_d is not None, \"optimizer_d required for GAN training\"\n        \n        metric_logger.add_meter(\n            'lr_d',\n            utils.SmoothedValue(window_size=1, fmt='{value}')\n        )\n    metric_logger.add_meter(\n        'samples/s', utils.SmoothedValue(window_size=10, fmt='{value:.3f}')\n    )\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 tqdm(\n        metric_logger.log_every(dataloader, print_freq, header),\n        desc = f\"Train Epoch {epoch}\",\n        leave = False\n    ):\n        start_time = time.time()\n        data, label = data.to(device), label.to(device)\n\n        if model_d is not None:\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(\n                lr_g=optimizer_g.param_groups[0]['lr'],\n                lr_d=optimizer_d.param_groups[0]['lr']\n            )\n            metric_logger.meters['samples/s'].update(\n                batch_size / (time.time() - start_time)\n            )\n\n        else:\n            optimizer_g.zero_grad()\n            pred = model(data)\n            loss_g, loss_g1v, loss_g2v = criterion_g(pred, label)\n            loss_g.backward()\n            optimizer_g.step()\n\n            metric_logger.update(loss_g1v=loss_g1v, loss_g2v=loss_g2v)\n\n            batch_size = data.shape[0]\n            metric_logger.update(\n                lr_g=optimizer_g.param_groups[0]['lr'],\n            )\n            metric_logger.meters['samples/s'].update(\n                batch_size / (time.time() - start_time)\n            )\n        step += 1\n        itr += 1\n        for lr_scheduler in lr_schedulers:\n            lr_scheduler.step()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T11:01:11.281098Z","iopub.execute_input":"2025-04-21T11:01:11.281599Z","iopub.status.idle":"2025-04-21T11:01:11.302578Z","shell.execute_reply.started":"2025-04-21T11:01:11.281580Z","shell.execute_reply":"2025-04-21T11:01:11.301963Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate(\n    model: torch.nn.Module,\n    criterion: Callable,\n    dataloader: torch.utils.data.DataLoader,\n    device: torch.device,\n    epoch: int\n) -> Tuple[float, float]:\n    \"\"\"\n    Evaluate model performance on validation/test data.\n\n    Computes evaluation metrics, denormalizes outputs, and optionally visualizes\n    results. Supports distributed evaluation with metric synchronization.\n\n    Parameters\n    ----------\n    model : torch.nn.Module\n        Model to be evaluated.\n    criterion : Callable\n        Loss function that returns (total_loss, l1_loss, l2_loss).\n        Signature: f(pred, label) -> (loss, loss_g1v, loss_g2v).\n    dataloader : torch.utils.data.DataLoader\n        DataLoader providing (data, label) batches.\n    device : torch.device\n        Device to run evaluation on (e.g., 'cuda' or 'cpu').\n    epoch : int\n        Current epoch number (used for visualization and logging).\n\n    Returns\n    -------\n    avg_loss : float\n        Average total loss across all batches.\n    l1loss_eval : float\n        L1 loss computed on denormalized full dataset.\n\n    Notes\n    -----\n    - Performs min-max denormalization using global `ctx['label_min/max']`.\n    - Visualizes first sample every 4 epochs (matplotlib required).\n    - Handles distributed evaluation with `synchronize_between_processes()`.\n    - All metrics are computed on denormalized data for final reporting.\n\n    Examples\n    --------\n    >>> model = MyModel()\n    >>> criterion = MyLoss()\n    >>> val_loader = DataLoader(val_dataset, batch_size=32)\n    >>> avg_loss, l1_loss = evaluate(model, criterion, val_loader, 'cuda', epoch=10)\n    \"\"\"\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 tqdm(\n            metric_logger.log_every(dataloader, 20, header),\n            desc=\"Validating\",\n            leave=False\n        ):\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(\n                loss=loss.item(),\n                loss_g1v=loss_g1v.item(),\n                loss_g2v=loss_g2v.item()\n            )\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(\n        all_output,\n        ctx['label_min'],\n        ctx['label_max']\n    )\n    all_label = T.minmax_denormalize(\n        all_label,\n        ctx['label_min'],\n        ctx['label_max']\n    )\n    l1loss_eval = l1loss(all_output, all_label)\n    \n    # Gather the stats from all processes\n    metric_logger.synchronize_between_processes()\n    print(\n        ' * Loss {loss.global_avg:.8f}, L1_loss {l1loss_eval:.8f} \\n'\n        .format(\n            loss=metric_logger.loss,\n            l1loss_eval=l1loss_eval\n        )\n    )\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-21T11:01:11.303350Z","iopub.execute_input":"2025-04-21T11:01:11.303697Z","iopub.status.idle":"2025-04-21T11:01:11.322182Z","shell.execute_reply.started":"2025-04-21T11:01:11.303657Z","shell.execute_reply":"2025-04-21T11:01:11.321507Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"step = 0\n\nif 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    with tqdm(\n        range(args[\"start_epoch\"], args[\"epochs\"]),\n        desc=\"Training Progress\",\n        unit=\"epoch\"\n    ) as epoch_pbar:\n        \n        for epoch in epoch_pbar:\n            train_one_epoch(\n                model,\n                criterion_g,\n                optimizer_g,\n                lr_schedulers,\n                dataloader_train,\n                device,\n                epoch,\n                args[\"print_freq\"],\n                model_d,\n                criterion_d if model_d else None,\n                optimizer_d if model_d else None,\n                args[\"n_critic\"] if model_d else 1\n            )\n        \n            loss_global_avg, l1loss_eval = evaluate(\n                model,\n                criterion_g,\n                dataloader_valid,\n                device,\n                epoch\n            )\n            checkpoint = {\n                'model': model_without_ddp.state_dict(),\n                'optimizer_g': optimizer_g.state_dict(),\n                'lr_schedulers': [scheduler.state_dict() for scheduler in lr_schedulers],\n                'epoch': epoch,\n                'step': step,\n                'args': args\n            }\n        \n            # Only include GAN components if they exist\n            if model_d is not None:\n                checkpoint.update({\n                    'model_d': model_d_without_ddp.state_dict(),\n                    'optimizer_d': optimizer_d.state_dict()\n                })\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        \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    \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-21T11:01:11.322878Z","iopub.execute_input":"2025-04-21T11:01:11.323160Z","iopub.status.idle":"2025-04-21T13:00:01.588935Z","shell.execute_reply.started":"2025-04-21T11:01:11.323144Z","shell.execute_reply":"2025-04-21T13:00:01.587856Z"},"_kg_hide-input":true,"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluating","metadata":{}},{"cell_type":"code","source":"%%time\ntest_files = list(Path('/kaggle/input/waveform-inversion/test').glob('*.npy'))\nlen(test_files)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-21T13:31:09.246379Z","iopub.execute_input":"2025-04-21T13:31:09.246943Z","iopub.status.idle":"2025-04-21T13:31:11.678628Z","shell.execute_reply.started":"2025-04-21T13:31:09.246902Z","shell.execute_reply":"2025-04-21T13:31:11.678065Z"}},"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":{"iopub.status.busy":"2025-04-20T20:56:02.988117Z","iopub.execute_input":"2025-04-20T20:56:02.988390Z","iopub.status.idle":"2025-04-20T20:56:02.992446Z","shell.execute_reply.started":"2025-04-20T20:56:02.988370Z","shell.execute_reply":"2025-04-20T20:56:02.991699Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TestDataset(Dataset):\n    \"\"\"\n    PyTorch Dataset for loading test/validation seismic data files.\n\n    This dataset loads individual seismic data files during inference/prediction,\n    optionally applying transformations, and returns both the data and original\n    filename (without extension) for tracking purposes.\n\n    Attributes\n    ----------\n    test_files : list\n        Stores the list of input file paths.\n    transform_data : callable or None\n        Stores the data transformation function.\n\n    Examples\n    --------\n    >>> from pathlib import Path\n    >>> test_files = list(Path('data/test').glob('*.npy'))\n    >>> dataset = TestDataset(test_files)\n    >>> len(dataset)  # Number of test files\n    50\n    >>> sample, fname = dataset[0]  # First sample and its filename\n    \"\"\"\n    def __init__(\n        self,\n        test_files: List[Union[str, Path]],\n        transform_data: Optional[Callable] = None\n    ) -> None:\n        \"\"\"\n        Initialize the TestDataset with file paths and optional transforms.\n\n        Parameters\n        ----------\n        test_files : list of PathLike\n            List of paths to seismic data files (typically .npy format).\n        transform_data : callable, optional\n            Transformations to apply to each loaded sample. If None, no transforms are applied.\n            Expected signature: transform(data: np.ndarray) -> transformed_data.\n        \"\"\"\n        self.test_files = test_files\n        self.transform_data = transform_data\n\n\n    def __len__(self) -> int:\n        \"\"\"\n        Return the number of test files in the dataset.\n\n        Returns\n        -------\n        int\n            Number of samples/files in the dataset.\n        \"\"\"\n        return len(self.test_files)\n\n\n    def __getitem__(\n        self,\n        i: int\n    ) -> Tuple[np.ndarray, str]:\n        \"\"\"\n        Load and return the i-th sample from the dataset.\n\n        Parameters\n        ----------\n        i : int\n            Index of the sample to retrieve.\n\n        Returns\n        -------\n        tuple\n            Contains:\n            - data : np.ndarray\n                Loaded seismic data array\n            - str\n                Base filename (without extension) of the loaded file\n\n        Notes\n        -----\n        - Automatically handles pathlib.Path or string file paths\n        - Applies transforms if transform_data was specified\n        - Uses numpy.load() for loading .npy files\n        \"\"\"\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":{"iopub.status.busy":"2025-04-20T20:56:05.696292Z","iopub.execute_input":"2025-04-20T20:56:05.696896Z","iopub.status.idle":"2025-04-20T20:56:05.703014Z","shell.execute_reply.started":"2025-04-20T20:56:05.696872Z","shell.execute_reply":"2025-04-20T20:56:05.702219Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ctx_test = data_config['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(\n    ds,\n    batch_size=8,\n    num_workers=4,\n    pin_memory=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T20:56:11.037254Z","iopub.execute_input":"2025-04-20T20:56:11.037857Z","iopub.status.idle":"2025-04-20T20:56:11.042656Z","shell.execute_reply.started":"2025-04-20T20:56:11.037834Z","shell.execute_reply":"2025-04-20T20:56:11.041949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"checkpoint = torch.load('/kaggle/input/openfwi-gans/pytorch/default/2/best_model(1).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(\n            y_preds,\n            ctx_test['label_min'],\n            ctx_test['label_max']\n        )\n        \n        for y_pred, oid_test in zip(y_preds, oids_test):\n            for y_pos in range(70):\n                row = dict(\n                    zip(\n                        x_cols,\n                        [y_pred[y_pos, x_pos] for x_pos in range(1, 70, 2)]\n                    )\n                )\n                row['oid_ypos'] = f\"{oid_test}_y_{y_pos}\"\n            \n                writer.writerow(row)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T20:56:13.920136Z","iopub.execute_input":"2025-04-20T20:56:13.920833Z","iopub.status.idle":"2025-04-20T21:03:05.668131Z","shell.execute_reply.started":"2025-04-20T20:56:13.920813Z","shell.execute_reply":"2025-04-20T21:03:05.667223Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#checkpoint = torch.load('/kaggle/input/openfwi-gans/pytorch/default/1/best_model.pth')\n#model_without_ddp.load_state_dict(network.replace_legacy(checkpoint['model']))\n\n# Ensure the model is in evaluation mode\n#model.eval()\n\n# Export the model to ONNX\n# Get a sample input batch from the DataLoader\n#sample_inputs, _ = next(iter(dl))\n#sample_inputs = sample_inputs.to(device)\n\n# Export the model\n#torch.onnx.export(\n#    model,  # Model to export\n#    sample_inputs,  # Example input\n#    \"model.onnx\",  # Output file name\n#    export_params=True,  # Include model parameters\n#    opset_version=12,  # ONNX opset version\n#    do_constant_folding=True,  # Optimize constants\n#    input_names=['input'],  # Input tensor name\n#    output_names=['output'],  # Output tensor name\n#    dynamic_axes={  # Allow dynamic batch size\n#        'input': {0: 'batch_size'},\n#        'output': {0: 'batch_size'}\n#    }\n#)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T21:06:43.380778Z","iopub.execute_input":"2025-04-20T21:06:43.381718Z","iopub.status.idle":"2025-04-20T21:07:07.025075Z","shell.execute_reply.started":"2025-04-20T21:06:43.381679Z","shell.execute_reply":"2025-04-20T21:07:07.024433Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}