{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# PyTorch + UNet + 5-fold Ensemble","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import KFold\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom tqdm.notebook import tqdm\nimport pandas as pd\nimport csv\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False\n\nCOMP_PATH = '/kaggle/input/waveform-inversion'\ntrain_dir = os.path.join(COMP_PATH, \"train_samples\")\ntest_dir = os.path.join(COMP_PATH, \"test\")\nBATCH_SIZE = 16\nN_FOLDS = 5\nNUM_EPOCHS = 30\nPATIENCE = 7\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def find_label_file(casedir):\n    # velocity.npy, vel.npy, model.npy, parent含む\n    for vname in [\"velocity.npy\", \"vel.npy\", \"model.npy\"]:\n        for search_dir in [casedir, os.path.join(casedir, \"data\")]:\n            vpath = os.path.join(search_dir, vname)\n            if os.path.isfile(vpath):\n                return vpath\n    for root, dirs, files in os.walk(casedir):\n        for fname in files:\n            if any(x in fname for x in [\"velocity\", \"vel\", \"model\"]) and fname.endswith(\".npy\"):\n                return os.path.join(root, fname)\n    return None\n\ndef get_train_triples(train_dir, verbose=True):\n    triples = []\n    max_chan = 0\n    for case in sorted(os.listdir(train_dir)):\n        case_dir = os.path.join(train_dir, case)\n        label_file = find_label_file(case_dir)\n        if not label_file:\n            if verbose:\n                print(f\"警告: ラベルファイル見つからずスキップ: {case_dir}\")\n            continue\n        for root, dirs, files in os.walk(case_dir):\n            for fname in sorted(files):\n                if fname.startswith(\"data\") and fname.endswith(\".npy\"):\n                    data_path = os.path.join(root, fname)\n                    arr = np.load(data_path, mmap_mode='r')\n                    shape = arr.shape\n                    # shape: (N, C, 1000, 70) or (C, 1000, 70)\n                    if len(shape) == 4:\n                        n_sample, n_chan, tlen, wlen = shape\n                    elif len(shape) == 3:\n                        n_sample, n_chan, tlen, wlen = 1, *shape\n                    else:\n                        raise ValueError(f\"data shapeが異常: {data_path} {shape}\")\n                    max_chan = max(max_chan, n_chan)\n                    for i in range(n_sample):\n                        triples.append({\n                            \"data_path\": data_path,\n                            \"label_path\": label_file,\n                            \"case\": case,\n                            \"file\": fname,\n                            \"idx\": i,\n                            \"n_chan\": n_chan,\n                            \"shape\": (n_chan, tlen, wlen)\n                        })\n    if verbose:\n        print(f\"Train triples: {len(triples)}, Max channel: {max_chan}\")\n    return triples, max_chan\n\ndef get_test_triples(test_dir, verbose=True):\n    triples = []\n    max_chan = 0\n    for fname in sorted(os.listdir(test_dir)):\n        if not fname.endswith(\".npy\"):\n            continue\n        fpath = os.path.join(test_dir, fname)\n        arr = np.load(fpath, mmap_mode='r')\n        shape = arr.shape\n        # shape: (C, 1000, 70)\n        if len(shape) == 3:\n            n_chan, tlen, wlen = shape\n        else:\n            raise ValueError(f\"test shape異常: {fpath} {shape}\")\n        max_chan = max(max_chan, n_chan)\n        triples.append({\n            \"data_path\": fpath,\n            \"label_path\": None,\n            \"case\": None,\n            \"file\": fname,\n            \"idx\": 0,\n            \"n_chan\": n_chan,\n            \"shape\": (n_chan, tlen, wlen)\n        })\n    if verbose:\n        print(f\"Test triples: {len(triples)}, Max channel: {max_chan}\")\n    return triples, max_chan\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class WaveformDataset(Dataset):\n    def __init__(self, triples, normalize=True, desired_channels=5):\n        self.triples = triples\n        self.normalize = normalize\n        self.desired_channels = desired_channels\n        arr = np.load(self.triples[0][\"data_path\"], mmap_mode='r')\n        if len(arr.shape) == 4:\n            ex = arr[0]\n        else:\n            ex = arr\n        self.global_mean = ex.mean()\n        self.global_std = ex.std() + 1e-8\n\n    def __len__(self):\n        return len(self.triples)\n\n    def __getitem__(self, idx):\n        triple = self.triples[idx]\n        arr = np.load(triple[\"data_path\"], mmap_mode='r')\n        if len(arr.shape) == 4:\n            waves = arr[triple[\"idx\"]]\n        else:\n            waves = arr\n        n_chan = waves.shape[0]\n        if n_chan < self.desired_channels:\n            pad = ((0, self.desired_channels - n_chan), (0,0), (0,0))\n            waves = np.pad(waves, pad)\n        elif n_chan > self.desired_channels:\n            waves = waves[:self.desired_channels]\n        if self.normalize:\n            waves = (waves - self.global_mean) / self.global_std\n        waves = np.ascontiguousarray(waves, dtype=np.float32)\n        # ラベル\n        if triple[\"label_path\"] is None:\n            return torch.from_numpy(waves), torch.zeros(1, 70, 70)\n        label_arr = np.load(triple[\"label_path\"], mmap_mode='r')\n        if label_arr.ndim == 4:\n            velocity = label_arr[triple[\"idx\"]]\n        elif label_arr.ndim == 3:\n            velocity = label_arr\n        elif label_arr.ndim == 2:\n            velocity = np.expand_dims(label_arr, axis=0)\n        else:\n            raise RuntimeError(f\"label shape異常: {triple['label_path']} shape={label_arr.shape}\")\n        velocity = np.ascontiguousarray(velocity, dtype=np.float32)\n        assert velocity.shape[-2:] == (70, 70), f\"velocity shape不正: {velocity.shape}\"\n        return torch.from_numpy(waves), torch.from_numpy(velocity)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_triples, max_train_chan = get_train_triples(train_dir)\ntest_triples, max_test_chan = get_test_triples(test_dir)\nDESIRED_CHANNELS = min(max_train_chan, max_test_chan, 8)\nprint(f\"Train max channels: {max_train_chan} | Test max channels: {max_test_chan} | Using: {DESIRED_CHANNELS}\")\n\nkf = KFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\nfolds = list(kf.split(train_triples))\nfold_loaders = []\nfor train_idx, val_idx in folds:\n    train_list = [train_triples[i] for i in train_idx]\n    val_list = [train_triples[i] for i in val_idx]\n    train_dataset = WaveformDataset(train_list, normalize=True, desired_channels=DESIRED_CHANNELS)\n    val_dataset = WaveformDataset(val_list, normalize=True, desired_channels=DESIRED_CHANNELS)\n    train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=2, pin_memory=True, persistent_workers=True)\n    val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2, pin_memory=True, persistent_workers=True)\n    fold_loaders.append((train_loader, val_loader))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class WaveformInversionUNet(nn.Module):\n    def __init__(self, in_channels):\n        super().__init__()\n        self.conv1 = nn.Sequential(nn.Conv2d(in_channels, 32, 3, padding=1), nn.ReLU(), nn.Conv2d(32, 32, 3, padding=1), nn.ReLU())\n        self.conv2 = nn.Sequential(nn.Conv2d(32, 64, 3, padding=1), nn.ReLU(), nn.Conv2d(64, 64, 3, padding=1), nn.ReLU())\n        self.conv3 = nn.Sequential(nn.Conv2d(64, 128, 3, padding=1), nn.ReLU(), nn.Conv2d(128, 128, 3, padding=1), nn.ReLU())\n        self.conv4 = nn.Sequential(nn.Conv2d(128, 256, 3, padding=1), nn.ReLU(), nn.Conv2d(256, 256, 3, padding=1), nn.ReLU())\n        self.conv5 = nn.Sequential(nn.Conv2d(256, 256, 3, padding=1), nn.ReLU(), nn.Conv2d(256, 256, 3, padding=1), nn.ReLU())\n        self.up4_conv = nn.Sequential(nn.Conv2d(256+256, 128, 3, padding=1), nn.ReLU(), nn.Conv2d(128, 128, 3, padding=1), nn.ReLU())\n        self.up3_conv = nn.Sequential(nn.Conv2d(128+128, 128, 3, padding=1), nn.ReLU(), nn.Conv2d(128, 128, 3, padding=1), nn.ReLU())\n        self.final_conv = nn.Conv2d(128, 1, 1)\n    def forward(self, x):\n        d1 = self.conv1(x)\n        p1 = F.max_pool2d(d1, (2,2))\n        d2 = self.conv2(p1)\n        p2 = F.max_pool2d(d2, (2,1))\n        d3 = self.conv3(p2)\n        if d3.shape[2] % 2 == 1:\n            d3 = F.pad(d3, (0,0,0,1))\n        skip3 = d3\n        p3 = F.max_pool2d(d3, (2,1))\n        d4 = self.conv4(p3)\n        if d4.shape[2] % 2 == 1:\n            d4 = F.pad(d4, (0,0,0,1))\n        skip4 = d4\n        p4 = F.max_pool2d(d4, (2,1))\n        d5 = self.conv5(p4)\n        up4 = F.interpolate(d5, size=(skip4.shape[2], skip4.shape[3]), mode='bilinear', align_corners=False)\n        u4 = self.up4_conv(torch.cat([up4, skip4], 1))\n        up3 = F.interpolate(u4, size=(skip3.shape[2], skip3.shape[3]), mode='bilinear', align_corners=False)\n        u3 = self.up3_conv(torch.cat([up3, skip3], 1))\n        out = self.final_conv(u3)\n        out = F.interpolate(out, size=(70, 70), mode='bilinear', align_corners=False)\n        out = 1500.0 + F.relu(out - 1500.0)\n        return out\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def total_variation_loss(img):\n    diff_i = torch.abs(img[:, :, 1:, :] - img[:, :, :-1, :]).mean()\n    diff_j = torch.abs(img[:, :, :, 1:] - img[:, :, :, :-1]).mean()\n    return diff_i + diff_j\n\ncriterion_l1 = nn.L1Loss()\ncriterion_l2 = nn.MSELoss()\n\nfold_models = []\nfor fold_idx, (train_loader, val_loader) in enumerate(fold_loaders):\n    print(f\"\\n===== Training Fold {fold_idx} =====\")\n    model = WaveformInversionUNet(in_channels=DESIRED_CHANNELS).to(DEVICE)\n    optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', factor=0.5, patience=5, verbose=True)\n    best_val_loss = float('inf')\n    patience_counter = 0\n    for epoch in range(1, NUM_EPOCHS+1):\n        model.train()\n        train_loss = 0.0\n        for batch_wave, batch_vel in tqdm(train_loader, desc=f\"Fold{fold_idx+1} Epoch{epoch} [Train]\", leave=False):\n            batch_wave = batch_wave.to(DEVICE, dtype=torch.float, non_blocking=True)\n            batch_vel = batch_vel.to(DEVICE, dtype=torch.float, non_blocking=True)\n            optimizer.zero_grad()\n            pred_vel = model(batch_wave)\n            loss_val = criterion_l1(pred_vel, batch_vel) + criterion_l2(pred_vel, batch_vel)\n            loss_tv = total_variation_loss(pred_vel)\n            loss = loss_val + 1e-4 * loss_tv\n            loss.backward()\n            optimizer.step()\n            train_loss += loss.item()\n        train_loss /= len(train_loader)\n        model.eval()\n        val_loss = 0.0\n        with torch.no_grad():\n            for batch_wave, batch_vel in val_loader:\n                batch_wave = batch_wave.to(DEVICE, dtype=torch.float, non_blocking=True)\n                batch_vel = batch_vel.to(DEVICE, dtype=torch.float, non_blocking=True)\n                pred_vel = model(batch_wave)\n                loss_val = criterion_l1(pred_vel, batch_vel) + criterion_l2(pred_vel, batch_vel)\n                loss_tv = total_variation_loss(pred_vel)\n                loss = loss_val + 1e-4 * loss_tv\n                val_loss += loss.item()\n        val_loss /= len(val_loader)\n        print(f\"Epoch {epoch}: Train Loss={train_loss:.4f}, Val Loss={val_loss:.4f}\")\n        scheduler.step(val_loss)\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            patience_counter = 0\n            best_state = model.state_dict()\n        else:\n            patience_counter += 1\n            if patience_counter >= PATIENCE:\n                print(f\"Early stopping at epoch {epoch}\")\n                break\n    model.load_state_dict(best_state)\n    fold_models.append(model.cpu())\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.backends.cudnn.benchmark = True  # 推論高速化（重要）\ntest_loader = DataLoader(\n    WaveformDataset(test_triples, desired_channels=DESIRED_CHANNELS),\n    batch_size=8, shuffle=False, num_workers=4, pin_memory=True\n)\n\nsubmit_path = '/kaggle/working/submission.csv'\nrows_buffer = []\nbuffer_size = 50000  # 適度なバッファリングサイズ\n\nwith open(submit_path, 'w', newline='') as f:\n    writer = csv.writer(f)\n    writer.writerow([\"Id\", \"Predicted\"])\n\nfor fold_idx, model in enumerate(fold_models):\n    model.to(DEVICE).eval()\n\n# 推論はバッチ単位でfoldごとにensembleを即時計算しRAM節約\nwith torch.inference_mode():\n    for batch_idx, (batch_wave, _) in enumerate(tqdm(test_loader, desc=\"Fast Inference\")):\n        batch_wave = batch_wave.to(DEVICE, dtype=torch.float, non_blocking=True)\n        \n        preds = []\n        for model in fold_models:\n            pred_vel = model(batch_wave).cpu().numpy()  # (B, 1, 70, 70)\n            preds.append(pred_vel)\n        ensemble_pred = np.mean(preds, axis=0)[:, 0]  # (B, 70, 70)\n\n        # CSV用行作成\n        for i, pred_ens in enumerate(ensemble_pred):\n            triple_idx = batch_idx * test_loader.batch_size + i\n            triple = test_triples[triple_idx]\n            base = os.path.basename(triple[\"data_path\"]).replace('.npy', '')\n            case = triple.get(\"case\") or \"test\"\n            sample_idx = triple[\"idx\"]\n            \n            for x in range(70):\n                for y in range(70):\n                    rows_buffer.append([\n                        f\"{case}_{base}_{sample_idx}_{x}_{y}\",\n                        pred_ens[x, y]\n                    ])\n\n            # バッファが一定数を超えたら書き込み\n            if len(rows_buffer) >= buffer_size:\n                with open(submit_path, 'a', newline='') as f:\n                    csv.writer(f).writerows(rows_buffer)\n                rows_buffer = []\n\n# 残りデータをすべて書き込み\nif rows_buffer:\n    with open(submit_path, 'a', newline='') as f:\n        csv.writer(f).writerows(rows_buffer)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}