{"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":11511096,"sourceType":"datasetVersion","datasetId":7218098},{"sourceId":11568812,"sourceType":"datasetVersion","datasetId":7253205},{"sourceId":11569667,"sourceType":"datasetVersion","datasetId":7253605},{"sourceId":11569755,"sourceType":"datasetVersion","datasetId":7253661},{"sourceId":12014343,"sourceType":"datasetVersion","datasetId":7558517}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"%%writefile config.yaml\nKAGGLE_TRAIN_DIR1 : \"/kaggle/input/open-wfi-1/openfwi_float16_1\"\nKAGGLE_TRAIN_DIR2 : \"/kaggle/input/open-wfi-2/openfwi_float16_2\"\nKAGGLE_TEST_DIR : \"/kaggle/input/open-wfi-test/test\"\nWORKING_DIR : \"/kaggle/working\"\nTEST_SIZE : 0.1\nBATCH_SIZE : 256\nMAX_EPOCHS : 1\nLEARNING_RATE : 1e-5\nWEIGHT_DECAY : 1e-6\nPLOT_EVERY_STEPS : 1000\nREAD_WEIGHTS : \"/kaggle/input/classify-best0531-loss022/classify_best0531_loss022.pt\"\nTRAIN : \"True\"\nTEST_WEIGHTS : \"/kaggle/input/classify-best0531-loss022/classify_best0531_loss022.pt\"\nFACTOR : 0.8\nPATIENCE : 0\nES_EPOCHS : 20\nSEED : 99","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T01:18:24.399030Z","iopub.execute_input":"2025-06-03T01:18:24.399209Z","iopub.status.idle":"2025-06-03T01:18:24.407481Z","shell.execute_reply.started":"2025-06-03T01:18:24.399194Z","shell.execute_reply":"2025-06-03T01:18:24.406680Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import yaml\n\nwith open(\"config.yaml\", \"r\") as file_obj:\n    cfg = yaml.safe_load(file_obj)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T01:18:24.408760Z","iopub.execute_input":"2025-06-03T01:18:24.408960Z","iopub.status.idle":"2025-06-03T01:18:24.443734Z","shell.execute_reply.started":"2025-06-03T01:18:24.408940Z","shell.execute_reply":"2025-06-03T01:18:24.443227Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nfrom pathlib import Path\nimport datetime\nimport random\nimport time\n\n\nimport torch\nimport torch.amp\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.distributed as dist\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm.auto import tqdm\n\n# This is classification case.\nLabelsMap = {\n    0: \"CurveFault_A\",\n    1: \"CurveFault_B\",\n    2: \"CurveVel_A\",\n    3: \"CurveVel_B\",\n    4: \"FlatFault_A\",\n    5: \"FlatFault_B\",\n    6: \"FlatVel_A\",\n    7: \"FlatVel_B\",\n    8: \"Style_A\",\n    9: \"Style_B\",\n}\nLabelToNum = {v: k for k, v in LabelsMap.items()}\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\ndef get_train_files(data_path):\n    all_inputs = [\n        f for f in Path(data_path).rglob(\"*.npy\")\n        if ('seis' in f.stem) or ('data' in f.stem)\n    ]\n\n    assert all(f.exists() for f in all_inputs)\n\n    return all_inputs\n\nclass SeismicDataset(Dataset):\n    def __init__(self, inputs_files, mode, n_examples_per_file=500):\n        self.inputs_files = inputs_files\n        self.n_examples_per_file = n_examples_per_file\n        self.mode = mode\n\n    def __len__(self):\n        return len(self.inputs_files) * self.n_examples_per_file\n\n    def __getitem__(self, idx):\n        file_idx = idx // self.n_examples_per_file\n        sample_idx = idx % self.n_examples_per_file\n\n        X = np.load(self.inputs_files[file_idx], mmap_mode='r')\n        y = os.path.basename(os.path.dirname(self.inputs_files[file_idx]))\n        if y == 'data' or y == 'model':\n            y = os.path.basename(os.path.dirname(os.path.dirname(self.inputs_files[file_idx])))\n        y = LabelToNum[y]\n\n        if self.mode == 'train': \n            if np.random.random() < 0.5:\n                X = X[::-1, :, ::-1]\n            \n        try:\n            return X[sample_idx].copy(), y\n        finally:\n            del X, y\n\n\nclass TestDataset(Dataset):\n    def __init__(self, test_files):\n        self.test_files = test_files\n\n    def __len__(self):\n        return len(self.test_files)\n\n    def __getitem__(self, i):\n        test_file = self.test_files[i]\n        return np.load(test_file), test_file.stem","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T01:18:24.444298Z","iopub.execute_input":"2025-06-03T01:18:24.444516Z","iopub.status.idle":"2025-06-03T01:18:28.801595Z","shell.execute_reply.started":"2025-06-03T01:18:24.444501Z","shell.execute_reply":"2025-06-03T01:18:28.800861Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"import datetime\nimport random\nimport torch\nimport numpy as np\n\ndef format_time(elapsed):\n    elapsed_rounded = int(round((elapsed)))\n    return str(datetime.timedelta(seconds=elapsed_rounded))\n\n\ndef seed_everything(\n    seed_value: int\n) -> None:\n    random.seed(seed_value)\n    np.random.seed(seed_value)\n    torch.manual_seed(seed_value)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed_value)\n        torch.cuda.manual_seed_all(seed_value)\n    if torch.backends.cudnn.is_available:\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T01:18:28.802995Z","iopub.execute_input":"2025-06-03T01:18:28.803296Z","iopub.status.idle":"2025-06-03T01:18:28.808628Z","shell.execute_reply.started":"2025-06-03T01:18:28.803280Z","shell.execute_reply":"2025-06-03T01:18:28.807763Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class SimpleNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.stem = nn.Sequential(\n            nn.ReflectionPad2d((2,2,80,80)),\n            nn.Conv2d(5,5,kernel_size=(4,4),stride=(4,1),padding=(0,1)),\n            nn.Conv2d(5,5,kernel_size=(4,4),stride=(4,1),padding=(0,1)),\n            nn.InstanceNorm2d(5)\n        )\n        self.features = nn.Sequential(\n            nn.Conv2d(5, 32, kernel_size=3, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            nn.MaxPool2d(2),#(B,32,36,36)\n            \n            nn.Conv2d(32, 64, kernel_size=3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.MaxPool2d(2),#(B,64,18,18)\n\n            nn.Conv2d(64, 128, kernel_size=3, padding=1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.MaxPool2d(2),#(B,128,9,9)\n        )\n\n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(128*81, 128),\n            nn.Dropout(0.5),\n            nn.Linear(128, 10),\n        )\n\n    def forward(self, x):\n        x = self.stem(x)\n        x = self.features(x)\n        x = self.classifier(x)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T01:18:28.809341Z","iopub.execute_input":"2025-06-03T01:18:28.809603Z","iopub.status.idle":"2025-06-03T01:18:28.829569Z","shell.execute_reply.started":"2025-06-03T01:18:28.809580Z","shell.execute_reply":"2025-06-03T01:18:28.828804Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train ","metadata":{}},{"cell_type":"code","source":"import sys\nimport os\nimport time\nimport gc\nimport numpy as np\n\n\nimport torch\nimport torch.nn as nn\n\nimport torch.distributed as dist\nfrom torch.nn.parallel import DistributedDataParallel\nfrom torch.utils.data import DataLoader, DistributedSampler\nfrom sklearn.model_selection import train_test_split\nfrom torchinfo import summary\n\n\n\nLabelsMap = {\n    0: \"CurveFault_A\",\n    1: \"CurveFault_B\",\n    2: \"CurveVel_A\",\n    3: \"CurveVel_B\",\n    4: \"FlatFault_A\",\n    5: \"FlatFault_B\",\n    6: \"FlatVel_A\",\n    7: \"FlatVel_B\",\n    8: \"Style_A\",\n    9: \"Style_B\",\n}\nLabelToNum = {v: k for k, v in LabelsMap.items()}\n\n\ndef train(cfg):\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    # Get Dataset\n    all_inputs = []\n    for x in [cfg[\"KAGGLE_TRAIN_DIR1\"], cfg[\"KAGGLE_TRAIN_DIR2\"]]:\n        all_inputs_tmp = get_train_files(x)\n        all_inputs.extend(all_inputs_tmp)\n    print(\"Total number of input files:\", len(all_inputs))\n\n\n    train_inputs, valid_inputs = train_test_split(all_inputs, test_size=cfg[\"TEST_SIZE\"], random_state=cfg[\"SEED\"])\n    print(f\"Num of train files: {len(train_inputs)}\")\n    print(f\"Num of valid files: {len(valid_inputs)}\")\n\n    dstrain = SeismicDataset(train_inputs, 'train')\n\n    dltrain = DataLoader(\n        dstrain,\n        batch_size=cfg[\"BATCH_SIZE\"],\n        shuffle=True,\n        pin_memory=False,\n        drop_last=True,\n        num_workers=4,\n        persistent_workers=True,\n    )\n    \n    dsvalid = SeismicDataset(valid_inputs, 'valid')\n\n    dlvalid = DataLoader(\n        dsvalid,\n        batch_size=cfg[\"BATCH_SIZE\"],\n        shuffle=True,\n        pin_memory=False,\n        drop_last=False,\n        num_workers=4,\n        persistent_workers=True,\n    )\n\n    # Define model\n    if cfg[\"READ_WEIGHTS\"] != \"None\":\n        print(\"Reading weights from:\", cfg[\"READ_WEIGHTS\"])\n        model = SimpleNet()\n        model.load_state_dict(torch.load(cfg[\"READ_WEIGHTS\"], map_location='cuda', weights_only=True))\n        model = model.to(device)\n    else:\n        model = SimpleNet().to(device)\n\n\n    # Define training params\n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=float(cfg[\"LEARNING_RATE\"]), weight_decay=float(cfg[\"WEIGHT_DECAY\"]))\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', factor=float(cfg[\"FACTOR\"]), patience=int(cfg[\"PATIENCE\"]))\n\n    best_val_loss = 1000.0\n    epochs_wo_improvement = 0\n    t0 = time.time()\n\n    for epoch in range(1, int(cfg[\"MAX_EPOCHS\"])+1):\n        # Train\n        model.train()\n        train_losses = []\n        correct = 0\n        for step, (inputs, targets) in enumerate(dltrain):\n            inputs = inputs.to(device)\n            targets = targets.to(device)\n            optimizer.zero_grad()\n            with torch.amp.autocast('cuda', dtype=torch.bfloat16):\n                outputs = model(inputs)\n                loss = criterion(outputs, targets)\n\n            preds = outputs.argmax(dim=1).cpu().numpy()\n            correct += (outputs.argmax(1) == targets).type(torch.float).sum().item()\n            loss.backward()\n            optimizer.step()\n            train_losses.append(loss.item())\n        \n        accuracy = (correct / len(dltrain.dataset))*100\n        if step % int(cfg[\"PLOT_EVERY_STEPS\"]) == 1 or step == len(dltrain) - 1:\n            trn_loss = np.mean(train_losses)\n            t1 = format_time(time.time() - t0)\n            lr = optimizer.param_groups[-1]['lr']\n            print(\n                    f\"Epoch: {epoch:02d} Step {step+1}/{len(dltrain)}  Trn Loss: {trn_loss:.2f} Accuracy: {accuracy:.2f}% LR: {lr:.2e} Elapsed Time: {t1}\",\n                    flush=True,\n                )\n\n        \n        # Valid\n        model.eval()\n        valid_losses = []\n        correct = 0\n        for inputs, targets in dlvalid:\n            inputs = inputs.to(device)\n            targets = targets.to(device)\n\n            with torch.inference_mode():\n                with torch.amp.autocast('cuda', dtype=torch.bfloat16):\n                    outputs = model(inputs)\n                    loss = criterion(outputs, targets)\n\n            preds = outputs.argmax(dim=1).cpu().numpy()\n            correct += (outputs.argmax(1) == targets).type(torch.float).sum().item()\n            valid_losses.append(loss.item())\n        accuracy = (correct / len(dlvalid.dataset))*100\n\n\n        # Gater loss on the same device\n        t1 = format_time(time.time() - t0)\n        trn_loss = np.mean(train_losses)\n        val_loss = np.mean(valid_losses)\n\n        free, total = torch.cuda.mem_get_info(device=0)\n        mem_used = (total - free) / 1024**3\n\n        # Log\n        print(\n            f\"\\nEpoch: {epoch:02d}  Trn Loss: {trn_loss:.2f}  Val Loss: {val_loss:.2f}  Accuracy: {accuracy:.2f}%  GPU Usage: {mem_used:.2f}GB  Elapsed Time: {t1}\",\n            flush=True,\n        )\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            epochs_wo_improvement = 0\n            torch.save(model.state_dict(), \"my_best_model.pt\")\n            print(f\"\\nNew best val_loss: {val_loss:.2f}\\n\", flush=True)\n\n        elif epoch == cfg[\"MAX_EPOCHS\"]:\n            torch.save(model.state_dict(), \"max_epoch.pt\")\n        else:\n            epochs_wo_improvement += 1\n            print(f\"\\nEpochs without improvement: {epochs_wo_improvement}\\n\", flush=True)\n\n        if epochs_wo_improvement == cfg[\"ES_EPOCHS\"]:\n            break\n\n        scheduler.step(val_loss)\n\n\n    # Cleanup\n    del model, optimizer, scheduler\n    del dltrain, dlvalid, dstrain, dsvalid\n    gc.collect()\n    torch.cuda.empty_cache()\n\n    return\n    \nseed_everything(cfg[\"SEED\"])\nif cfg[\"TRAIN\"] == \"True\":\n    train(cfg)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T02:30:19.772794Z","iopub.execute_input":"2025-06-03T02:30:19.773147Z","iopub.status.idle":"2025-06-03T02:34:30.585214Z","shell.execute_reply.started":"2025-06-03T02:30:19.773107Z","shell.execute_reply":"2025-06-03T02:34:30.580510Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"==================== Output Classification Result ==================\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T01:41:54.229364Z","iopub.status.idle":"2025-06-03T01:41:54.229604Z","shell.execute_reply.started":"2025-06-03T01:41:54.229492Z","shell.execute_reply":"2025-06-03T01:41:54.229501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\ntest_files = list(Path(\"/kaggle/input/open-wfi-test/test\").glob(\"*.npy\"))\nds = TestDataset(test_files)\ndl = DataLoader(ds, batch_size=cfg[\"BATCH_SIZE\"], num_workers=4, pin_memory=False)\n\n\nprint(\"Reading weights from:\", cfg[\"TEST_WEIGHTS\"])\nmodel = SimpleNet()\nmodel.load_state_dict(torch.load(cfg[\"TEST_WEIGHTS\"], map_location='cuda', weights_only=True))\nmodel = model.to(device)\nmodel.eval()\n\nresult = []\nfor i, (inputs, oid) in enumerate(dl):\n    inputs = inputs.to(device)\n    with torch.inference_mode():\n        with torch.autocast(device_type=\"cuda\", dtype=torch.bfloat16):\n            outputs = model(inputs)\n\n    preds = outputs.argmax(dim=1).cpu().numpy()\n    result.extend(preds)\n    n = (i+1)*cfg[\"BATCH_SIZE\"]\n    if n%4096 == 0:\n        print(f\"processing {n} files\")\n\nfilename = \"test_type.txt\"\nwith open(filename, 'w', encoding='utf-8') as f:\n    for item in result:\n        f.write(f\"{item}\\n\")\n\nprint(\"output complete\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T01:41:54.230625Z","iopub.status.idle":"2025-06-03T01:41:54.230940Z","shell.execute_reply.started":"2025-06-03T01:41:54.230778Z","shell.execute_reply":"2025-06-03T01:41:54.230792Z"}},"outputs":[],"execution_count":null}]}