{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install webdataset","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:32:54.925787Z","iopub.execute_input":"2025-04-15T13:32:54.926040Z","iopub.status.idle":"2025-04-15T13:33:01.148479Z","shell.execute_reply.started":"2025-04-15T13:32:54.926000Z","shell.execute_reply":"2025-04-15T13:33:01.147654Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### The data size of this compeition is very large (~670 GB) (https://openfwi-lanl.github.io/docs/data.html).  \n\n##### I introduce an implementation using WebDataset to handle all data.","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\n\nimport numpy as np\nimport webdataset as wds\nfrom torch.utils.data import DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom tqdm.auto import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:33:01.149924Z","iopub.execute_input":"2025-04-15T13:33:01.150286Z","iopub.status.idle":"2025-04-15T13:33:09.326554Z","shell.execute_reply.started":"2025-04-15T13:33:01.150249Z","shell.execute_reply":"2025-04-15T13:33:09.325655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class cfg:\n    data_dir     = \"/kaggle/input/waveform-inversion\"\n    output_dir   = \"/kaggle/working\"\n    dataset_name = \"fwi_dataset\"\n    stage        = \"train\"    # train|test\n    target_dirs  = [\n      \"FlatVel_A\",\n      \"FlatVel_B\",\n      \"CurveVel_A\",\n      \"CurveVel_B\",\n      \"FlatFault_A\",\n      \"FlatFault_B\",\n      # \"CurveFault_A\",\n      # \"CurveFault_B\",\n      # \"Style_A\",\n      # \"Style_B\",\n    ]\n    maxsize         = 0.6e9    # 600MB (default: 3GB). Maximum size of each shard.\n    num_used_shards = None\n    test_size       = 0.2      # represent the proportion of the dataset to include in the test split\n    batch_size      = 2\n    seed            = 42\n    debug           = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:33:09.328731Z","iopub.execute_input":"2025-04-15T13:33:09.329101Z","iopub.status.idle":"2025-04-15T13:33:09.334361Z","shell.execute_reply.started":"2025-04-15T13:33:09.329080Z","shell.execute_reply":"2025-04-15T13:33:09.333639Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def search_data_path(target_dirs, root_dir, shuffle=True):\n    files = []\n    for target_dir in target_dirs:\n        data_dir = Path(root_dir, target_dir)\n        assert data_dir.is_dir(), f\"{data_dir} is not found\"\n        if Path(data_dir, \"data\").is_dir():\n            in_files = sorted(Path(data_dir, \"data\").glob(\"*.npy\"))\n            out_files = sorted(Path(data_dir, \"model\").glob(\"*.npy\"))\n        else:\n            in_files = sorted(data_dir.glob(\"seis*.npy\"))\n            out_files = sorted(data_dir.glob(\"vel*.npy\"))\n        assert len(in_files) == len(out_files)\n        files += list(zip(in_files, out_files))\n    if shuffle:\n        np.random.shuffle(files)\n    return files\n\n\ndef generate_sample(in_file, out_file=None):\n    if out_file is None:  # test data\n        seis = np.load(in_file)\n        assert seis.shape == (5, 1000, 70)\n        seis = seis.astype(np.float16)\n        data = [{\n            \"__key__\": in_file.stem,\n            \"sample_id.txt\": in_file.stem,\n            \"seis.npy\": seis,\n        }]\n    else:  # train/val data\n        seis = np.load(in_file)\n        assert seis.shape == (500, 5, 1000, 70)\n        seis = seis.astype(np.float16)\n        vel = np.load(out_file)\n        assert vel.shape == (500, 1, 70, 70)\n        vel = vel.astype(np.float16)\n        for i, j in zip(in_file.parents, out_file.parents):\n            if i == j:\n                common_path = str(i)\n                break\n        data = []\n        for i in range(len(seis)):\n            sample_id = f\"{in_file.stem}_{out_file.stem}_{i}\"\n            data.append({\n                \"__key__\": common_path + \"_\" + sample_id,\n                \"sample_id.txt\": sample_id,\n                \"seis.npy\": seis[i],\n                \"vel.npy\": vel[i],\n            })\n    return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:33:09.335076Z","iopub.execute_input":"2025-04-15T13:33:09.335295Z","iopub.status.idle":"2025-04-15T13:33:09.349604Z","shell.execute_reply.started":"2025-04-15T13:33:09.335277Z","shell.execute_reply":"2025-04-15T13:33:09.348931Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cfg.stage == \"train\":\n    data_paths = search_data_path(cfg.target_dirs, Path(cfg.data_dir, \"train_samples\"), shuffle=True)\nelse:\n    data_paths = sorted(Path(cfg.data_dir, \"test\").glob(\"*.npy\"))\n    data_paths = [(i, None) for i in data_paths]\n    if cfg.debug:\n        data_paths = data_paths[:100]\ndata_paths[:3]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:33:09.350660Z","iopub.execute_input":"2025-04-15T13:33:09.351001Z","iopub.status.idle":"2025-04-15T13:33:10.687441Z","shell.execute_reply.started":"2025-04-15T13:33:09.350966Z","shell.execute_reply":"2025-04-15T13:33:10.686634Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !rm -fr /kaggle/working/*","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:33:10.688305Z","iopub.execute_input":"2025-04-15T13:33:10.688655Z","iopub.status.idle":"2025-04-15T13:33:10.692719Z","shell.execute_reply.started":"2025-04-15T13:33:10.688630Z","shell.execute_reply":"2025-04-15T13:33:10.691848Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset_name = f\"{cfg.stage}_{cfg.dataset_name}\"\noutput_dir = Path(cfg.output_dir, dataset_name)\noutput_dir.mkdir(parents=True, exist_ok=True)\noutput_dir","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:33:10.693634Z","iopub.execute_input":"2025-04-15T13:33:10.693944Z","iopub.status.idle":"2025-04-15T13:33:10.710949Z","shell.execute_reply.started":"2025-04-15T13:33:10.693917Z","shell.execute_reply":"2025-04-15T13:33:10.710060Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if len(list(output_dir.glob(\"*.tar\"))) > 0:\n    print(f\"already exists tar files in {output_dir}, skip\")\nelse:\n    writer = wds.ShardWriter(str(Path(output_dir, \"%04d.tar\")), maxsize=cfg.maxsize)\n    for in_file, out_file in tqdm(data_paths):\n        data = generate_sample(in_file, out_file)\n        for d in data:\n            writer.write(d)\n    writer.close()\ndata[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:33:10.711864Z","iopub.execute_input":"2025-04-15T13:33:10.712213Z","iopub.status.idle":"2025-04-15T13:33:13.837256Z","shell.execute_reply.started":"2025-04-15T13:33:10.712176Z","shell.execute_reply":"2025-04-15T13:33:13.836530Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_shard_paths(root_dir, dataset_name, stage, num_shards=None, test_size=0.2, seed=42):\n    assert stage in [\"train\", \"test\"]\n    assert num_shards is None or num_shards > 1\n    assert 0 < test_size < 1\n    dataset_dir = Path(root_dir, f\"{stage}_{dataset_name}\")\n    shard_paths = np.array(sorted(map(str, dataset_dir.glob(\"*.tar\"))))\n    if stage == \"train\":\n        if num_shards is not None:\n            rng = np.random.default_rng(seed)\n            shard_paths = rng.choice(shard_paths, size=min(num_shards, len(shard_paths)), replace=False)\n            shard_paths.sort()\n        trn_idx, val_idx = train_test_split(np.arange(len(shard_paths)), test_size=test_size, random_state=seed, shuffle=True)\n        trn_shard_paths = sorted(shard_paths[trn_idx])\n        val_shard_paths = sorted(shard_paths[val_idx])\n        print(f\"# of train shards: {len(trn_shard_paths)}, # of val shards: {len(val_shard_paths)}\")\n        return trn_shard_paths, val_shard_paths\n    else:\n        print(f\"# of test shards: {len(shard_paths)}\")\n        return sorted(shard_paths)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:33:13.839311Z","iopub.execute_input":"2025-04-15T13:33:13.839557Z","iopub.status.idle":"2025-04-15T13:33:13.846145Z","shell.execute_reply.started":"2025-04-15T13:33:13.839538Z","shell.execute_reply":"2025-04-15T13:33:13.845421Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"paths = get_shard_paths(\n    cfg.output_dir,\n    cfg.dataset_name,\n    cfg.stage,\n    num_shards=cfg.num_used_shards,\n    test_size=cfg.test_size,\n    seed=cfg.seed,\n)\nif cfg.stage == \"train\":\n    train_paths, val_paths = paths\n    print(\"example of train files:\", train_paths[:3])\n    print(\"example of val files:\", val_paths[:3])\nelse:\n    test_paths = paths\n    print(\"example of test files:\", test_paths[:3])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:33:13.846996Z","iopub.execute_input":"2025-04-15T13:33:13.847193Z","iopub.status.idle":"2025-04-15T13:33:13.890811Z","shell.execute_reply.started":"2025-04-15T13:33:13.847178Z","shell.execute_reply":"2025-04-15T13:33:13.889906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_dataset(paths, stage, seed=42):\n    dataset = wds.WebDataset(paths, seed=seed).decode()\n    if stage != \"test\":\n        dataset = (\n            dataset\n            .to_tuple(\"sample_id.txt\", \"seis.npy\", \"vel.npy\")\n            .map(\n                lambda x: {\n                    \"sample_id\": x[0],\n                    \"seis\": x[1],\n                    \"vel\": x[2],\n                }\n            )\n        )\n        if stage == \"train\":\n            dataset = dataset.shuffle(10)\n    elif stage == \"test\":\n        dataset = (\n            dataset\n            .to_tuple(\"sample_id.txt\", \"seis.npy\")\n            .map(\n                lambda x: {\n                    \"sample_id\": x[0],\n                    \"seis\": x[1],\n                }\n            )\n        )\n    return dataset","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:33:13.891602Z","iopub.execute_input":"2025-04-15T13:33:13.892013Z","iopub.status.idle":"2025-04-15T13:33:13.910492Z","shell.execute_reply.started":"2025-04-15T13:33:13.891978Z","shell.execute_reply":"2025-04-15T13:33:13.909774Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cfg.stage == \"train\":\n    train_dataset = get_dataset(train_paths, \"train\", seed=cfg.seed)\n    val_dataset = get_dataset(val_paths, \"val\", seed=cfg.seed)\n    train_dataloader = DataLoader(train_dataset, batch_size=cfg.batch_size, num_workers=1, drop_last=True, pin_memory=True)\n    val_dataloader = DataLoader(val_dataset, batch_size=cfg.batch_size, num_workers=1, drop_last=False, pin_memory=True)\nelse:\n    test_dataset = get_dataset(test_paths, \"test\", seed=cfg.seed)\n    test_dataloader = DataLoader(test_dataset, batch_size=cfg.batch_size, num_workers=1, drop_last=False, pin_memory=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:33:13.911458Z","iopub.execute_input":"2025-04-15T13:33:13.911797Z","iopub.status.idle":"2025-04-15T13:33:13.926943Z","shell.execute_reply.started":"2025-04-15T13:33:13.911777Z","shell.execute_reply":"2025-04-15T13:33:13.925966Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cfg.stage == \"train\":\n    print(\"train dataloader\")\n    for batch in train_dataloader:\n        print(batch)\n        break\n    print(\"val dataloader\")\n    for batch in val_dataloader:\n        print(batch)\n        break\nelse:\n    print(\"test dataloader\")\n    for batch in test_dataloader:\n        print(batch)\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T13:33:13.927955Z","iopub.execute_input":"2025-04-15T13:33:13.928265Z","iopub.status.idle":"2025-04-15T13:33:14.228685Z","shell.execute_reply.started":"2025-04-15T13:33:13.928206Z","shell.execute_reply":"2025-04-15T13:33:14.227651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}