{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":51753,"databundleVersionId":5692552,"sourceType":"competition"},{"sourceId":707566,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":537277,"modelId":550661}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# =========================================================\n# [1] Imports + Config (with tqdm)\n# =========================================================\nimport os\nimport random\nimport time\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport matplotlib.pyplot as plt\n\nfrom tqdm.auto import tqdm  # <-- progress bars\n\n# Try albumentations; fallback if missing\ntry:\n    import albumentations as A\n    _HAS_ALB = True\nexcept Exception:\n    _HAS_ALB = False\n    A = None\n\nCFG = {\n    \"seed\": 42,\n    \"data_dir\": \"/kaggle/input/google-research-identify-contrails-reduce-global-warming\",\n\n    \"batch_size\": 2,\n    \"num_workers\": 2,\n\n    # ============================\n    # 【改动 1】继续训练相关配置\n    # ============================\n    # 你的第8轮权重路径（你需要改成你自己的文件名/路径）\n    # 例如: \"/kaggle/working/epoch8.pth\" 或 \"best_3dunet_temporalattn.pth\"\n    \"resume_path\": \"/kaggle/input/model1/pytorch/default/1/best_3dunet_temporalattn.pth\",\n    \"resume_epoch\": 8,      # 你目前已经跑完8轮，所以从8之后继续\n    \"extra_epochs\": 3,      # 继续再跑3轮 -> 9/10/11\n\n    \"lr\": 1e-4,\n    \"weight_decay\": 1e-2,\n    \"grad_clip\": 1.0,\n\n    # augmentation\n    \"augment\": \"d4\",          # safer than \"rotation\" for speed\n    \"augment_prob\": 0.95,\n\n    # losses / threshold\n    \"lambda_cons\": 0.10,\n    \"th\": 0.40,\n\n    # inference TTA\n    \"use_tta\": False,\n    \"tta_mode\": \"d4prob\",     # 8x inference. set use_tta=False if time tight\n\n    # visualization\n    \"vis_every\": 1,           # 0 means only visualize at epoch 1 and last\n    \"vis_samples\": 2,\n}\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nCFG[\"amp\"] = (DEVICE == \"cuda\")  # <-- only enable AMP on GPU\nprint(\"DEVICE:\", DEVICE, \"| AMP:\", CFG[\"amp\"], \"| Albumentations:\", _HAS_ALB)\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# [2] Reproducibility\n# =========================================================\ndef seed_everything(seed=42):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything(CFG[\"seed\"])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# [3] First-place-style y_sym trick: create_grid(offset=0.5) + augmentation\n# =========================================================\ndef create_grid(nc: int, offset: float = 0.5) -> torch.Tensor:\n    \"\"\"\n    Grid for torch.grid_sample in [-1,1].\n    offset=0.5 implements the \"0.5 pixel shift\" label alignment trick.\n    Returns: (1, nc, nc, 2) float32 on CPU.\n    \"\"\"\n    grid = np.zeros((nc, nc, 2), dtype=np.float32)\n    for ix in range(nc):\n        for iy in range(nc):\n            grid[ix, iy, 1] = -1 + 2 * (ix + 0.5) / nc + offset / 128\n            grid[ix, iy, 0] = -1 + 2 * (iy + 0.5) / nc + offset / 128\n    return torch.from_numpy(grid).unsqueeze(0)  # (1, nc, nc, 2)\n\ndef get_augmenter(mode: str):\n    \"\"\"\n    Returns an augmentation function that takes (image(H,W,CT), mask(H,W,1))\n    and returns augmented (image, mask).\n    If Albumentations is not available, uses numpy rot90/flip fallback.\n    \"\"\"\n    if mode not in [\"d4\", \"rotation\"]:\n        raise ValueError(\"augment must be 'd4' or 'rotation'\")\n\n    if _HAS_ALB:\n        if mode == \"d4\":\n            aug = A.Compose([\n                A.RandomRotate90(p=1.0),\n                A.HorizontalFlip(p=0.5),\n            ])\n        else:\n            aug = A.Compose([\n                A.RandomRotate90(p=1.0),\n                A.HorizontalFlip(p=0.5),\n                A.ShiftScaleRotate(rotate_limit=30, scale_limit=0.2, p=0.75),\n            ])\n\n        def _apply(image, mask):\n            out = aug(image=image, mask=mask)\n            return out[\"image\"], out[\"mask\"]\n        return _apply\n\n    # Fallback: only rot90 + flip (safe)\n    def _apply(image, mask):\n        k = np.random.randint(0, 4)\n        img = np.rot90(image, k, axes=(0, 1)).copy()\n        msk = np.rot90(mask, k, axes=(0, 1)).copy()\n        if np.random.rand() < 0.5:\n            img = np.flip(img, axis=1).copy()\n            msk = np.flip(msk, axis=1).copy()\n        return img, msk\n\n    return _apply\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# [4] Ash false color\n# =========================================================\ndef rescale_range(x, f_min, f_max):\n    return (x - f_min) / (f_max - f_min)\n\ndef ash_color_np(b11, b14, b15):\n    r = rescale_range(b15 - b14, -4, 2)\n    g = rescale_range(b14 - b11, -4, 5)\n    b = rescale_range(b14, 243, 303)\n    out = np.stack([r, g, b], axis=0).astype(np.float32)\n    out = 1.0 - out\n    out = np.clip(out, 0.0, 1.0)\n    return out\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# [5] Dataset: returns x(3,8,256,256), y_sym(1,256,256), y(1,256,256), label(1,256,256), w, record_id\n# =========================================================\nclass Contrails3DDataset(Dataset):\n    def __init__(self, root_dir: str, record_ids, train_mode: bool, cfg: dict):\n        self.root = Path(root_dir)\n        self.record_ids = list(record_ids)\n        self.train_mode = train_mode\n        self.cfg = cfg\n\n        self.grid = create_grid(256, offset=0.5)  # CPU grid\n        self.aug_prob = cfg[\"augment_prob\"] if train_mode else 0.0\n        self.aug_apply = get_augmenter(cfg[\"augment\"]) if train_mode else None\n\n    def __len__(self):\n        return len(self.record_ids)\n\n    def _safe_load(self, path: Path):\n        if not path.exists():\n            return None\n        return np.load(path).astype(np.float32)\n\n    def _load_bands(self, rid_path: Path):\n        b11 = self._safe_load(rid_path / \"band_11.npy\")\n        b14 = self._safe_load(rid_path / \"band_14.npy\")\n        b15 = self._safe_load(rid_path / \"band_15.npy\")\n        return b11, b14, b15\n\n    def _build_ash_sequence(self, b11, b14, b15):\n        # b11/b14/b15: (256,256,8)\n        frames = []\n        T = b11.shape[-1]\n        for t in range(T):\n            frames.append(ash_color_np(b11[..., t], b14[..., t], b15[..., t]))\n        x = np.stack(frames, axis=1)  # (3,T,H,W)\n        return x.astype(np.float32)\n\n    def _load_targets(self, rid_path: Path):\n        \"\"\"\n        y: soft label if available (mean of individuals), else hard\n        label: hard GT for metric\n        \"\"\"\n        label = None\n        hard = self._safe_load(rid_path / \"human_pixel_masks.npy\")\n        if hard is not None:\n            # could be (256,256,1) or (256,256) depending on version\n            hard = hard.reshape(-1)\n            hard = hard.reshape(256, 256, 1).astype(np.float32)\n            label = np.transpose(hard, (2, 0, 1))  # (1,256,256)\n\n        y = None\n        ind = self._safe_load(rid_path / \"human_individual_masks.npy\")\n        if ind is not None:\n            # expect (256,256,1,A)\n            if ind.ndim == 4:\n                y_soft = np.mean(ind, axis=3)  # (256,256,1)\n                y_soft = np.transpose(y_soft, (2, 0, 1))  # (1,256,256)\n                y = y_soft.astype(np.float32)\n\n        if y is None and label is not None:\n            y = label.copy()\n\n        return y, label\n\n    def __getitem__(self, idx):\n        rid = self.record_ids[idx]\n        rid_path = self.root / rid\n\n        b11, b14, b15 = self._load_bands(rid_path)\n        if b11 is None or b14 is None or b15 is None:\n            # Fail-safe: return a zero sample (won't crash loader)\n            x = torch.zeros((3, 8, 256, 256), dtype=torch.float32)\n            return {\"record_id\": rid, \"x\": x}\n\n        x = self._build_ash_sequence(b11, b14, b15)  # (3,8,256,256)\n        x = torch.from_numpy(x)\n\n        ret = {\"record_id\": rid, \"x\": x}\n\n        y, label = self._load_targets(rid_path)\n        if y is None:\n            # test sample: no targets\n            return ret\n\n        y = torch.from_numpy(y)  # (1,256,256)\n\n        # y_sym via grid_sample on CPU\n        y_sym = F.grid_sample(\n            y.unsqueeze(0),                 # (1,1,256,256)\n            self.grid,                      # (1,256,256,2)\n            mode=\"bilinear\",\n            padding_mode=\"border\",\n            align_corners=False\n        ).squeeze(0)  # (1,256,256)\n\n        # w=0 if augmented, w=1 if not augmented (First-place style)\n        w_original = 1.0\n\n        # Apply SAME spatial augmentation to ALL frames packed as channels, and to y_sym\n        if self.train_mode and (np.random.random() < self.aug_prob):\n            w_original = 0.0\n\n            # pack (C,T) -> channels: (H,W,C*T)\n            x_np = x.permute(2, 3, 0, 1).contiguous().numpy()  # (H,W,C,T)\n            H, W, C, T = x_np.shape\n            x_np = x_np.reshape(H, W, C * T).astype(np.float32)\n\n            y_np = y_sym.permute(1, 2, 0).contiguous().numpy().astype(np.float32)  # (H,W,1)\n\n            try:\n                x_aug, y_aug = self.aug_apply(x_np, y_np)\n            except Exception:\n                # Safety fallback: do nothing if aug fails\n                x_aug, y_aug = x_np, y_np\n\n            x_aug = x_aug.reshape(H, W, C, T).transpose(2, 3, 0, 1)  # (C,T,H,W)\n            x = torch.from_numpy(x_aug.astype(np.float32))\n            y_sym = torch.from_numpy(y_aug.transpose(2, 0, 1).astype(np.float32))  # (1,H,W)\n\n        ret.update({\n            \"x\": x,\n            \"y\": y,\n            \"y_sym\": y_sym,\n            \"w\": torch.tensor(w_original, dtype=torch.float32),\n        })\n\n        if label is not None:\n            ret[\"label\"] = torch.from_numpy(label.astype(np.float32))  # (1,256,256)\n\n        return ret","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# [6] Model: 3D U-Net (no T downsample) + Bottleneck Temporal Attention + asym_conv\n# =========================================================\nclass Conv3DBlock(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.Conv3d(in_ch, out_ch, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(out_ch, out_ch, kernel_size=3, padding=1),\n            nn.BatchNorm3d(out_ch),\n            nn.ReLU(inplace=True),\n        )\n    def forward(self, x):\n        return self.block(x)\n\nclass TemporalAttentionBlock(nn.Module):\n    \"\"\"\n    Attention only along time dimension T at each spatial location.\n    Input : (B,C,T,H,W) -> reshape to (B*H*W, T, C) -> MHA -> reshape back.\n    \"\"\"\n    def __init__(self, channels: int, num_heads: int = 4, dropout: float = 0.1):\n        super().__init__()\n        assert channels % num_heads == 0, \"channels must be divisible by num_heads\"\n        self.ln1 = nn.LayerNorm(channels)\n        self.attn = nn.MultiheadAttention(embed_dim=channels, num_heads=num_heads,\n                                          dropout=dropout, batch_first=True)\n        self.drop1 = nn.Dropout(dropout)\n\n        self.ln2 = nn.LayerNorm(channels)\n        self.ffn = nn.Sequential(\n            nn.Linear(channels, channels * 4),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(channels * 4, channels),\n        )\n        self.drop2 = nn.Dropout(dropout)\n\n    def forward(self, x):\n        B, C, T, H, W = x.shape\n        seq = x.permute(0, 3, 4, 2, 1).contiguous().view(B * H * W, T, C)  # (BHW,T,C)\n\n        seq_n = self.ln1(seq)\n        attn_out, _ = self.attn(seq_n, seq_n, seq_n, need_weights=False)\n        seq = seq + self.drop1(attn_out)\n\n        seq_n2 = self.ln2(seq)\n        seq = seq + self.drop2(self.ffn(seq_n2))\n\n        out = seq.view(B, H, W, T, C).permute(0, 4, 3, 1, 2).contiguous()  # (B,C,T,H,W)\n        return out\n\ndef get_asym_conv_256(hidden=9):\n    # Tiny 2D conv head to map y_sym_pred -> y_pred (original label space)\n    return nn.Sequential(\n        nn.Conv2d(1, hidden, kernel_size=3, padding=1, padding_mode=\"replicate\"),\n        nn.ReLU(inplace=True),\n        nn.Conv2d(hidden, 1, kernel_size=1),\n    )\n\nclass UNet3D_TemporalAttn(nn.Module):\n    \"\"\"\n    Input : (B,3,8,256,256)\n    Outputs:\n      y_sym_pred (B,1,256,256) from out_3d at t=4\n      y_pred     (B,1,256,256) after asym_conv\n      out_3d     (B,1,8,256,256) logits\n    \"\"\"\n    def __init__(self, in_channels=3, base_channels=16, attn_heads=4):\n        super().__init__()\n        # Encoder (no temporal downsample, pool=(1,2,2))\n        self.enc1 = Conv3DBlock(in_channels, base_channels)\n        self.pool1 = nn.MaxPool3d(kernel_size=(1, 2, 2))\n\n        self.enc2 = Conv3DBlock(base_channels, base_channels * 2)\n        self.pool2 = nn.MaxPool3d(kernel_size=(1, 2, 2))\n\n        self.enc3 = Conv3DBlock(base_channels * 2, base_channels * 4)\n        self.pool3 = nn.MaxPool3d(kernel_size=(1, 2, 2))\n\n        # Bottleneck\n        self.proj = nn.Conv3d(base_channels * 4, base_channels * 4, kernel_size=1)\n        self.temporal_attn = TemporalAttentionBlock(base_channels * 4, num_heads=attn_heads, dropout=0.1)\n\n        # Decoder\n        self.up3 = nn.ConvTranspose3d(base_channels * 4, base_channels * 4,\n                                      kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec3 = Conv3DBlock(base_channels * 8, base_channels * 4)\n\n        self.up2 = nn.ConvTranspose3d(base_channels * 4, base_channels * 2,\n                                      kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec2 = Conv3DBlock(base_channels * 4, base_channels * 2)\n\n        self.up1 = nn.ConvTranspose3d(base_channels * 2, base_channels,\n                                      kernel_size=(1, 2, 2), stride=(1, 2, 2))\n        self.dec1 = Conv3DBlock(base_channels * 2, base_channels)\n\n        self.final3d = nn.Conv3d(base_channels, 1, kernel_size=1)\n        self.asym_conv = get_asym_conv_256()\n\n    def forward(self, x):\n        e1 = self.enc1(x)         # (B,base,T,256,256)\n        p1 = self.pool1(e1)       # (B,base,T,128,128)\n\n        e2 = self.enc2(p1)        # (B,2b,T,128,128)\n        p2 = self.pool2(e2)       # (B,2b,T,64,64)\n\n        e3 = self.enc3(p2)        # (B,4b,T,64,64)\n        p3 = self.pool3(e3)       # (B,4b,T,32,32)\n\n        z = self.proj(p3)\n        z = self.temporal_attn(z) # (B,4b,T,32,32)\n\n        u3 = self.up3(z)\n        d3 = self.dec3(torch.cat([u3, e3], dim=1))\n\n        u2 = self.up2(d3)\n        d2 = self.dec2(torch.cat([u2, e2], dim=1))\n\n        u1 = self.up1(d2)\n        d1 = self.dec1(torch.cat([u1, e1], dim=1))\n\n        out_3d = self.final3d(d1)          # (B,1,T,256,256)\n        y_sym_pred = out_3d[:, :, 4, :, :] # (B,1,256,256) (5th frame)\n        y_pred = self.asym_conv(y_sym_pred)\n\n        return y_sym_pred, y_pred, out_3d","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# [7] Losses: First-place-style main loss + Neighbor consistency loss\n# =========================================================\nclass BCELossSymOriginal(nn.Module):\n    \"\"\"\n    L_main = BCE(y_sym_pred, y_sym) + w * BCE(y_pred, y)\n    where w=0 for augmented samples, w=1 for non-aug samples.\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss(reduction=\"none\")\n\n    def forward(self, y_sym_pred, y_sym, y_pred, y, w):\n        # per-sample mean over pixels\n        loss_sym = self.bce(y_sym_pred, y_sym).mean(dim=(1, 2, 3))\n        loss_org = self.bce(y_pred, y).mean(dim=(1, 2, 3))\n        loss = loss_sym + w * loss_org\n        return loss.mean()\n\ndef neighbor_consistency_loss(out_3d, center_t=4):\n    \"\"\"\n    out_3d logits: (B,1,T,H,W)\n    L_cons = MSE(sigmoid(t-1), sigmoid(t).detach()) + MSE(sigmoid(t+1), sigmoid(t).detach())\n    \"\"\"\n    T = out_3d.shape[2]\n    if center_t - 1 < 0 or center_t + 1 >= T:\n        return out_3d.new_tensor(0.0)\n\n    p_c = torch.sigmoid(out_3d[:, :, center_t, :, :])\n    p_l = torch.sigmoid(out_3d[:, :, center_t - 1, :, :])\n    p_r = torch.sigmoid(out_3d[:, :, center_t + 1, :, :])\n\n    return F.mse_loss(p_l, p_c.detach()) + F.mse_loss(p_r, p_c.detach())\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# [8] Train/Validate with tqdm progress bars\n# =========================================================\n@torch.no_grad()\ndef validate_loss_and_global_dice(model, loader, criterion, w_val, threshold=0.45, lambda_cons=0.1):\n    model.eval()\n    total_loss, total_n = 0.0, 0\n    tp, pp, tt = 0.0, 0.0, 0.0\n\n    pbar = tqdm(loader, desc=\"val\", leave=False)\n    for batch in pbar:\n        if \"y\" not in batch or \"y_sym\" not in batch or \"label\" not in batch:\n            continue\n\n        x = batch[\"x\"].to(DEVICE, non_blocking=True)\n        y = batch[\"y\"].to(DEVICE, non_blocking=True)\n        y_sym = batch[\"y_sym\"].to(DEVICE, non_blocking=True)\n        label = batch[\"label\"].to(DEVICE, non_blocking=True)\n\n        y_sym_pred, y_pred, out_3d = model(x)\n        loss_main = criterion(y_sym_pred, y_sym, y_pred, y, w_val)\n        loss_cons = neighbor_consistency_loss(out_3d, center_t=4)\n        loss = loss_main + lambda_cons * loss_cons\n\n        bs = x.size(0)\n        total_loss += loss.item() * bs\n        total_n += bs\n\n        pred_bin = (torch.sigmoid(y_pred) > threshold).float()\n        tp += (pred_bin * label).sum().item()\n        pp += pred_bin.sum().item()\n        tt += label.sum().item()\n\n        pbar.set_postfix({\"loss\": f\"{(total_loss/max(total_n,1)):.4f}\"})\n\n    val_loss = total_loss / max(total_n, 1)\n    val_dice = (2.0 * tp) / (pp + tt + 1e-6)\n    return val_loss, float(val_dice)\n\ndef train_one_epoch(model, loader, optimizer, scaler, criterion, lambda_cons=0.1):\n    model.train()\n    total_loss, total_n = 0.0, 0\n\n    pbar = tqdm(loader, desc=\"train\", leave=False)\n    for step, batch in enumerate(pbar):\n        if \"y\" not in batch or \"y_sym\" not in batch or \"w\" not in batch:\n            continue\n\n        x = batch[\"x\"].to(DEVICE, non_blocking=True)\n        y = batch[\"y\"].to(DEVICE, non_blocking=True)\n        y_sym = batch[\"y_sym\"].to(DEVICE, non_blocking=True)\n        w = batch[\"w\"].to(DEVICE, non_blocking=True)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        if CFG[\"amp\"]:\n            with autocast():\n                y_sym_pred, y_pred, out_3d = model(x)\n                loss_main = criterion(y_sym_pred, y_sym, y_pred, y, w)\n                loss_cons = neighbor_consistency_loss(out_3d, center_t=4)\n                loss = loss_main + lambda_cons * loss_cons\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_norm_(model.parameters(), CFG[\"grad_clip\"])\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            y_sym_pred, y_pred, out_3d = model(x)\n            loss_main = criterion(y_sym_pred, y_sym, y_pred, y, w)\n            loss_cons = neighbor_consistency_loss(out_3d, center_t=4)\n            loss = loss_main + lambda_cons * loss_cons\n            loss.backward()\n            nn.utils.clip_grad_norm_(model.parameters(), CFG[\"grad_clip\"])\n            optimizer.step()\n\n        bs = x.size(0)\n        total_loss += loss.item() * bs\n        total_n += bs\n\n        pbar.set_postfix({\"loss\": f\"{(total_loss/max(total_n,1)):.4f}\"})\n\n        if step == 0:\n            pbar.write(f\"[sanity] first batch ok. x={tuple(x.shape)} | DEVICE={DEVICE}\")\n\n    return total_loss / max(total_n, 1)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# [9] TTA for inference (d4prob) on 3D input -> average 2D logits\n# =========================================================\n_TTA_CONFIG = {\n    \"none\": (1, True),\n    \"d4prob\": (8, True),\n}\n\ndef _tta_stack_3d(x, n):\n    # x: (B,C,T,H,W) -> (nB,C,T,H,W)\n    if n == 1:\n        return x\n    stack = []\n    for k in range(4):\n        xr = torch.rot90(x, k, dims=[3, 4])\n        stack.append(xr)\n        if n == 8:\n            stack.append(torch.flip(xr, dims=[4]))  # horizontal flip\n    return torch.cat(stack, dim=0)\n\ndef _tta_average_2d_logits(y_logits, n, prob_average=True):\n    # y_logits: (nB,1,H,W) -> (B,1,H,W)\n    if n == 1:\n        return y_logits\n    N, C, H, W = y_logits.shape\n    B = N // n\n    y = y_logits.view(n, B, C, H, W)\n\n    if prob_average:\n        y = torch.sigmoid(y)\n\n    y_avg = torch.zeros((B, C, H, W), device=y_logits.device, dtype=y_logits.dtype)\n\n    if n == 8:\n        for k in range(4):\n            y_avg += (1.0 / n) * torch.rot90(y[2 * k], -k, dims=[2, 3])\n            y_avg += (1.0 / n) * torch.rot90(torch.flip(y[2 * k + 1], dims=[3]), -k, dims=[2, 3])\n    elif n == 4:\n        for k in range(4):\n            y_avg += (1.0 / n) * torch.rot90(y[k], -k, dims=[2, 3])\n    else:\n        raise ValueError(\"n must be 1,4,8\")\n\n    if prob_average:\n        y_avg = y_avg.clamp(1e-6, 1 - 1e-6)\n        return torch.log(y_avg) - torch.log(1 - y_avg)\n    return y_avg\n\n@torch.no_grad()\ndef predict_logits(model, x, tta_mode=\"none\"):\n    n, _ = _TTA_CONFIG.get(tta_mode, (1, True))\n    if n == 1:\n        _, y_pred, _ = model(x)\n        return y_pred\n    x_stack = _tta_stack_3d(x, n)\n    _, y_pred_stack, _ = model(x_stack)\n    y_pred = _tta_average_2d_logits(y_pred_stack, n=n, prob_average=True)\n    return y_pred\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# [10] Visualization: training curves + qualitative predictions\n# =========================================================\ndef plot_training_curves(history):\n    if not history or len(history[\"epoch\"]) == 0:\n        print(\"No history to plot.\")\n        return\n\n    epochs = history[\"epoch\"]\n\n    plt.figure()\n    plt.plot(epochs, history[\"train_loss\"], marker=\"o\", label=\"train_loss\")\n    plt.plot(epochs, history[\"val_loss\"], marker=\"o\", label=\"val_loss\")\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"Loss\")\n    plt.title(\"Train/Val Loss\")\n    plt.grid(True)\n    plt.legend()\n    plt.show()\n\n    plt.figure()\n    plt.plot(epochs, history[\"val_dice\"], marker=\"o\", label=\"val_global_dice\")\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"Global Dice\")\n    plt.title(\"Validation Global Dice\")\n    plt.grid(True)\n    plt.legend()\n    plt.show()\n\ndef dice_np(pred, gt, eps=1e-6):\n    pred = pred.astype(np.float32)\n    gt = gt.astype(np.float32)\n    inter = (pred * gt).sum()\n    return (2.0 * inter) / (pred.sum() + gt.sum() + eps)\n\n@torch.no_grad()\ndef visualize_val_predictions(model, dataset, num_samples=3, threshold=0.45):\n    if len(dataset) == 0:\n        print(\"Empty dataset.\")\n        return\n\n    model.eval()\n    idxs = np.random.choice(len(dataset), size=min(num_samples, len(dataset)), replace=False)\n\n    for idx in idxs:\n        sample = dataset[idx]\n        if \"label\" not in sample or \"x\" not in sample:\n            print(f\"[skip] idx={idx} no label/x\")\n            continue\n\n        x = sample[\"x\"].unsqueeze(0).to(DEVICE)  # (1,3,8,256,256)\n        gt = sample[\"label\"].cpu().numpy()[0]    # (256,256)\n\n        logits = predict_logits(model, x, tta_mode=(CFG[\"tta_mode\"] if CFG[\"use_tta\"] else \"none\"))\n        prob = torch.sigmoid(logits).cpu().numpy()[0, 0]  # (256,256)\n        pred = (prob > threshold).astype(np.uint8)\n\n        # show input (t=5) Ash RGB\n        x_cpu = sample[\"x\"].cpu().numpy()  # (3,8,256,256)\n        rgb = np.transpose(x_cpu[:, 4], (1, 2, 0))  # (256,256,3)\n\n        d = dice_np(pred, gt)\n\n        plt.figure(figsize=(12, 3))\n        plt.subplot(1, 4, 1)\n        plt.imshow(rgb)\n        plt.title(f\"Ash RGB t=5 (idx={idx})\")\n        plt.axis(\"off\")\n\n        plt.subplot(1, 4, 2)\n        plt.imshow(gt, cmap=\"gray\")\n        plt.title(\"GT\")\n        plt.axis(\"off\")\n\n        plt.subplot(1, 4, 3)\n        plt.imshow(prob, cmap=\"gray\")\n        plt.title(\"Pred prob\")\n        plt.axis(\"off\")\n\n        plt.subplot(1, 4, 4)\n        plt.imshow(pred, cmap=\"gray\")\n        plt.title(f\"Pred mask\\nDice={d:.3f}\")\n        plt.axis(\"off\")\n\n        plt.tight_layout()\n        plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# [11] RLE + Submission writer (align with sample_submission order)\n# =========================================================\ndef rle_encode(mask):\n    \"\"\"\n    mask: (H,W) uint8 {0,1}\n    Return run-length list [start, length, start, length, ...] (Fortran order)\n    \"\"\"\n    dots = np.where(mask.T.flatten() == 1)[0]\n    if len(dots) == 0:\n        return []\n    runs = []\n    prev = -2\n    for b in dots:\n        if b > prev + 1:\n            runs.extend((b + 1, 0))\n        runs[-1] += 1\n        prev = b\n    return runs\n\ndef rle_to_string(rle):\n    if len(rle) == 0:\n        return \"\"\n    return \" \".join(map(str, rle))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# [12] Build loaders + Resume-Train (继续跑3轮) + Curves + Visuals\n# =========================================================\nDATA_DIR = Path(CFG[\"data_dir\"])\nTRAIN_DIR = DATA_DIR / \"train\"\nVAL_DIR   = DATA_DIR / \"validation\"\nTEST_DIR  = DATA_DIR / \"test\"\n\ntrain_ids = sorted(os.listdir(TRAIN_DIR))\nval_ids   = sorted(os.listdir(VAL_DIR))\ntest_ids  = sorted(os.listdir(TEST_DIR))\n\ntrain_ds = Contrails3DDataset(TRAIN_DIR, train_ids, train_mode=True, cfg=CFG)\nval_ds   = Contrails3DDataset(VAL_DIR,   val_ids,   train_mode=False, cfg=CFG)\ntest_ds  = Contrails3DDataset(TEST_DIR,  test_ids,  train_mode=False, cfg=CFG)\n\ntrain_loader = DataLoader(\n    train_ds, batch_size=CFG[\"batch_size\"], shuffle=True,\n    num_workers=CFG[\"num_workers\"], pin_memory=(DEVICE==\"cuda\"), drop_last=True\n)\nval_loader = DataLoader(\n    val_ds, batch_size=CFG[\"batch_size\"], shuffle=False,\n    num_workers=CFG[\"num_workers\"], pin_memory=(DEVICE==\"cuda\"), drop_last=False\n)\ntest_loader = DataLoader(\n    test_ds, batch_size=CFG[\"batch_size\"], shuffle=False,\n    num_workers=CFG[\"num_workers\"], pin_memory=(DEVICE==\"cuda\"), drop_last=False\n)\n\nprint(\"train/val/test:\", len(train_ds), len(val_ds), len(test_ds))\n\nmodel = UNet3D_TemporalAttn(in_channels=3, base_channels=16, attn_heads=4).to(DEVICE)\ncriterion = BCELossSymOriginal()\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG[\"lr\"], weight_decay=CFG[\"weight_decay\"])\nscaler = GradScaler(enabled=CFG[\"amp\"])\n\nw_val = torch.tensor(1.0 - CFG[\"augment_prob\"], device=DEVICE, dtype=torch.float32)\n\nbest_score = -1.0\nbest_path = \"best_3dunet_temporalattn.pth\"\n\nhistory = {\"epoch\": [], \"train_loss\": [], \"val_loss\": [], \"val_dice\": []}\n\n# ============================\n# 【改动 2】继续训练：加载第8轮权重（容错，不会报错退出）\n# ============================\nstart_epoch = int(CFG[\"resume_epoch\"])\nresume_path = CFG[\"resume_path\"]\n\nif resume_path and os.path.exists(resume_path):\n    print(f\"[resume] Loading model weights from: {resume_path}\")\n    state = torch.load(resume_path, map_location=DEVICE)\n    try:\n        model.load_state_dict(state, strict=True)\n    except Exception as e:\n        # 兼容某些人保存的是 {\"model\": state_dict} 这种结构\n        if isinstance(state, dict) and (\"model\" in state):\n            model.load_state_dict(state[\"model\"], strict=True)\n        else:\n            raise e\n    print(f\"[resume] OK. Will continue training: epoch {start_epoch+1} -> {start_epoch + CFG['extra_epochs']}\")\nelse:\n    print(f\"[resume] WARNING: resume_path not found: {resume_path}\")\n    print(\"[resume] Will train from scratch for extra_epochs (still safe).\")\n    start_epoch = 0  # 让逻辑保持一致：从1开始跑 extra_epochs\n\n# ============================\n# 【改动 3】训练循环：只跑“继续的3轮”，而不是写死8轮\n# ============================\nend_epoch = start_epoch + int(CFG[\"extra_epochs\"])\n\nprint(f\"Start training (resume) for {CFG['extra_epochs']} epochs: [{start_epoch+1}..{end_epoch}] ...\")\nt0 = time.time()\n\nfor epoch in range(start_epoch + 1, end_epoch + 1):\n    ep0 = time.time()\n\n    tr_loss = train_one_epoch(\n        model, train_loader, optimizer, scaler, criterion,\n        lambda_cons=CFG[\"lambda_cons\"]\n    )\n    val_loss, val_dice = validate_loss_and_global_dice(\n        model, val_loader, criterion, w_val,\n        threshold=CFG[\"th\"], lambda_cons=CFG[\"lambda_cons\"]\n    )\n\n    history[\"epoch\"].append(epoch)\n    history[\"train_loss\"].append(tr_loss)\n    history[\"val_loss\"].append(val_loss)\n    history[\"val_dice\"].append(val_dice)\n\n    print(f\"Epoch {epoch:02d} | train_loss={tr_loss:.5f} | val_loss={val_loss:.5f} | val_global_dice={val_dice:.5f} | epoch_time={(time.time()-ep0):.1f}s\")\n\n    # 额外：每一轮都存一个“当前epoch权重”，方便你继续接着跑\n    torch.save(model.state_dict(), f\"epoch_{epoch:02d}.pth\")\n\n    if val_dice > best_score:\n        best_score = val_dice\n        torch.save(model.state_dict(), best_path)\n        print(f\"  -> saved best to {best_path} (best dice={best_score:.5f})\")\n\n    if epoch == (start_epoch + 1) or epoch == end_epoch or (CFG[\"vis_every\"] > 0 and epoch % CFG[\"vis_every\"] == 0):\n        try:\n            visualize_val_predictions(model, val_ds, num_samples=CFG[\"vis_samples\"], threshold=CFG[\"th\"])\n        except Exception as e:\n            print(\"[warn] visualization failed:\", e)\n\nprint(f\"Done. Best val dice={best_score:.5f} | total {(time.time()-t0)/60:.1f} min\")\n\nplot_training_curves(history)\n\n# Load best and visualize more\nif os.path.exists(best_path):\n    model.load_state_dict(torch.load(best_path, map_location=DEVICE))\nmodel.eval()\nvisualize_val_predictions(model, val_ds, num_samples=6, threshold=CFG[\"th\"])\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# [13] Inference + submission.csv (with tqdm)\n# =========================================================\nif os.path.exists(best_path):\n    model.load_state_dict(torch.load(best_path, map_location=DEVICE))\nmodel.eval()\n\nsample_path = DATA_DIR / \"sample_submission.csv\"\nsample_sub = pd.read_csv(sample_path)\n\npred_map = {}\n\nwith torch.no_grad():\n    for batch in tqdm(test_loader, desc=\"infer\", leave=False):\n        x = batch[\"x\"].to(DEVICE, non_blocking=True)\n        rids = batch[\"record_id\"]\n\n        logits = predict_logits(model, x, tta_mode=(CFG[\"tta_mode\"] if CFG[\"use_tta\"] else \"none\"))\n        prob = torch.sigmoid(logits).float().cpu().numpy()\n        masks = (prob > CFG[\"th\"]).astype(np.uint8)[:, 0]\n\n        for rid, m in zip(rids, masks):\n            pred_map[rid] = rle_to_string(rle_encode(m))\n\nsample_sub[\"encoded_pixels\"] = sample_sub[\"record_id\"].map(lambda r: pred_map.get(r, \"\"))\nsample_sub.to_csv(\"submission.csv\", index=False)\n\nprint(\"Saved submission.csv\")\nprint(sample_sub.head())\nprint(\"Non-empty masks:\", (sample_sub[\"encoded_pixels\"].str.len() > 0).sum())","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}