{"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":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":13333,"databundleVersionId":862146}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n#for dirname, _, filenames in os.walk('/kaggle/input'):\n #   for filename in filenames:\n        #print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"b4b79dd4-fe7a-4296-b168-cbf9a55cbacd","_cell_guid":"24eca766-4af4-4ae9-b360-eb8a6da268ef","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T01:43:38.923851Z","iopub.execute_input":"2026-07-06T01:43:38.924133Z","iopub.status.idle":"2026-07-06T01:43:40.851528Z","shell.execute_reply.started":"2026-07-06T01:43:38.924099Z","shell.execute_reply":"2026-07-06T01:43:40.850691Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nimport math\nfrom pathlib import Path\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.model_selection import train_test_split\nimport albumentations as A","metadata":{"_uuid":"ffef5185-5880-49a2-8c25-7e642501778e","_cell_guid":"77b28114-2470-4582-95ec-9dc001f13885","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:32:19.238572Z","iopub.execute_input":"2026-07-06T03:32:19.238872Z","iopub.status.idle":"2026-07-06T03:32:19.243367Z","shell.execute_reply.started":"2026-07-06T03:32:19.238835Z","shell.execute_reply":"2026-07-06T03:32:19.242629Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = Path(\"/kaggle/input/competitions/understanding_cloud_organization\")\nTRAIN_IMG_DIR = DATA_DIR / \"train_images\"\nTEST_IMG_DIR = DATA_DIR / \"test_images\"\nTRAIN_CSV = DATA_DIR / \"train.csv\"\nSAMPLE_SUB_CSV = DATA_DIR / \"sample_submission.csv\"\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(DEVICE)\n\nIMAGE_H = 350\nIMAGE_W = 525\nBATCH_SIZE = 10\nNUM_EPOCHS = 30\nLR = 1e-3\n\nCLASS_NAMES = [\"Fish\", \"Flower\", \"Gravel\", \"Sugar\"]\nCLASS_TO_IDX = {name: i for i, name in enumerate(CLASS_NAMES)}","metadata":{"_uuid":"5f872f2a-fd7e-4bee-9786-2de483502407","_cell_guid":"cbd00064-be11-4554-bf73-a2c8e41d643b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:32:22.168537Z","iopub.execute_input":"2026-07-06T03:32:22.168800Z","iopub.status.idle":"2026-07-06T03:32:22.174660Z","shell.execute_reply.started":"2026-07-06T03:32:22.168777Z","shell.execute_reply":"2026-07-06T03:32:22.173829Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rle_decode(mask_rle, shape):\n    \"\"\"\n    mask_rle: string like 'start length start length ...'\n    shape: (height, width)\n    return: np.array of shape (height, width), values in {0,1}\n    \"\"\"\n    s = mask_rle.split()\n    starts = np.asarray(s[0::2], dtype=int) - 1\n    lengths = np.asarray(s[1::2], dtype=int)\n    ends = starts + lengths\n\n    img = np.zeros(shape[0] * shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n\n    # Kaggle 常见是按列优先展开，需要 reshape 后转置\n    img = img.reshape((shape[1], shape[0])).T\n    return img\nprint(\"done\")","metadata":{"_uuid":"31e6c20d-5d20-4bd9-8a33-b96b3d5a2e2d","_cell_guid":"10a4ab19-601e-43a8-8a1f-3372bb849ba0","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:32:25.575003Z","iopub.execute_input":"2026-07-06T03:32:25.575517Z","iopub.status.idle":"2026-07-06T03:32:25.581752Z","shell.execute_reply.started":"2026-07-06T03:32:25.575486Z","shell.execute_reply":"2026-07-06T03:32:25.580938Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rle_encode(mask):\n    \"\"\"\n    mask: np.array of shape (H, W), values in {0,1}\n    returns RLE string\n    \"\"\"\n    pixels = mask.T.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[0::2]\n    return \" \".join(str(x) for x in runs)\nprint(\"done\")","metadata":{"_uuid":"83f1abd2-073b-4a0f-acca-6381eb0f03aa","_cell_guid":"9251e51c-d308-4804-9c55-da43ef1dd2d1","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:32:27.842536Z","iopub.execute_input":"2026-07-06T03:32:27.843319Z","iopub.status.idle":"2026-07-06T03:32:27.848672Z","shell.execute_reply.started":"2026-07-06T03:32:27.843287Z","shell.execute_reply":"2026-07-06T03:32:27.847762Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(TRAIN_CSV)\ndf[[\"image\", \"label\"]] = df[\"Image_Label\"].str.rsplit(\"_\", n=1, expand=True)\n\n# 每张图整理成一行，四类分别一列\nmask_df = df.pivot(index=\"image\", columns=\"label\", values=\"EncodedPixels\")\nmask_df = mask_df.reset_index()\n\nfor c in CLASS_NAMES:\n    if c not in mask_df.columns:\n        mask_df[c] = np.nan\n\nmask_df.head()","metadata":{"_uuid":"e0756c22-7041-439b-944a-e2c94185ba75","_cell_guid":"99add6e9-6a2c-4f85-96dc-e32a96534827","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:32:30.142852Z","iopub.execute_input":"2026-07-06T03:32:30.143682Z","iopub.status.idle":"2026-07-06T03:32:32.372141Z","shell.execute_reply.started":"2026-07-06T03:32:30.143650Z","shell.execute_reply":"2026-07-06T03:32:32.371453Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df, val_df = train_test_split(\n    mask_df,\n    test_size=0.2,\n    random_state=42\n)\n\ntrain_df = train_df.reset_index(drop=True)\nval_df = val_df.reset_index(drop=True)\nprint(\"done\")","metadata":{"_uuid":"b8320ac5-e33a-4c29-a060-1101eca58cbc","_cell_guid":"bbeab563-f64e-4314-afd7-a72b42abd181","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:32:35.468810Z","iopub.execute_input":"2026-07-06T03:32:35.469666Z","iopub.status.idle":"2026-07-06T03:32:35.481179Z","shell.execute_reply.started":"2026-07-06T03:32:35.469633Z","shell.execute_reply":"2026-07-06T03:32:35.480346Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMAGENET_MEAN = (0.485, 0.456, 0.406)\nIMAGENET_STD = (0.229, 0.224, 0.225)\n\ntrain_transform = A.Compose([\n    A.Resize(IMAGE_H, IMAGE_W),\n\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n\n    A.ShiftScaleRotate(\n        shift_limit=0.05,\n        scale_limit=0.10,\n        rotate_limit=10,\n        border_mode=cv2.BORDER_CONSTANT,\n        value=0,\n        mask_value=0,\n        p=0.5\n    ),\n\n    A.RandomBrightnessContrast(\n        brightness_limit=0.15,\n        contrast_limit=0.15,\n        p=0.4\n    ),\n\n    A.GaussNoise(p=0.2),\n\n    A.Normalize(\n        mean=IMAGENET_MEAN,\n        std=IMAGENET_STD,\n        max_pixel_value=255.0\n    ),\n])\n\nval_transform = A.Compose([\n    A.Resize(IMAGE_H, IMAGE_W),\n    A.Normalize(\n        mean=IMAGENET_MEAN,\n        std=IMAGENET_STD,\n        max_pixel_value=255.0\n    ),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-06T03:32:38.380268Z","iopub.execute_input":"2026-07-06T03:32:38.380957Z","iopub.status.idle":"2026-07-06T03:32:38.391296Z","shell.execute_reply.started":"2026-07-06T03:32:38.380926Z","shell.execute_reply":"2026-07-06T03:32:38.390549Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CloudDataset(Dataset):\n    def __init__(self, df, image_dir, transform=None):\n        self.df = df\n        self.image_dir = Path(image_dir)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_name = row[\"image\"]\n\n        image_path = self.image_dir / image_name\n        image = cv2.imread(str(image_path))\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n        h0, w0 = image.shape[:2]\n\n        mask = np.zeros((h0, w0, 4), dtype=np.float32)\n\n        for class_name in CLASS_NAMES:\n            rle = row[class_name]\n            if isinstance(rle, str):\n                class_idx = CLASS_TO_IDX[class_name]\n                mask[:, :, class_idx] = rle_decode(rle, (h0, w0))\n\n        if self.transform is not None:\n            transformed = self.transform(image=image, mask=mask)\n            image = transformed[\"image\"]\n            mask = transformed[\"mask\"]\n\n        # HWC -> CHW\n        image = np.transpose(image, (2, 0, 1)).astype(np.float32)\n        mask = np.transpose(mask, (2, 0, 1)).astype(np.float32)\n\n        image = torch.tensor(image, dtype=torch.float32)\n        mask = torch.tensor(mask, dtype=torch.float32)\n\n        return image, mask, image_name\n\nprint(\"done\")","metadata":{"_uuid":"c2205a5d-e288-46cf-8b74-fc424519e8a9","_cell_guid":"19c03b59-e097-450a-bae6-d1839434cbbf","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:32:43.004237Z","iopub.execute_input":"2026-07-06T03:32:43.004574Z","iopub.status.idle":"2026-07-06T03:32:43.013145Z","shell.execute_reply.started":"2026-07-06T03:32:43.004545Z","shell.execute_reply":"2026-07-06T03:32:43.012262Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = CloudDataset(\n    train_df,\n    TRAIN_IMG_DIR,\n    transform=train_transform\n)\n\nval_dataset = CloudDataset(\n    val_df,\n    TRAIN_IMG_DIR,\n    transform=val_transform\n)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=4,\n    pin_memory=True,\n    persistent_workers=True,\n    prefetch_factor=2\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=4,\n    pin_memory=True,\n    persistent_workers=True,\n    prefetch_factor=2\n)\nprint(\"done\")","metadata":{"_uuid":"8412d67f-0bcf-4b3f-bb26-deb2a9147c01","_cell_guid":"0e03b749-b9d3-458f-8eff-f52b4864dbc6","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:32:46.380804Z","iopub.execute_input":"2026-07-06T03:32:46.381547Z","iopub.status.idle":"2026-07-06T03:32:46.489007Z","shell.execute_reply.started":"2026-07-06T03:32:46.381516Z","shell.execute_reply":"2026-07-06T03:32:46.488136Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DiceLoss(nn.Module):\n    def __init__(self, smooth=1.0):\n        super().__init__()\n        self.smooth = smooth\n\n    def forward(self, logits, targets):\n        probs = torch.sigmoid(logits)\n        probs = probs.contiguous().view(probs.size(0), probs.size(1), -1)\n        targets = targets.contiguous().view(targets.size(0), targets.size(1), -1)\n\n        intersection = (probs * targets).sum(dim=2)\n        denominator = probs.sum(dim=2) + targets.sum(dim=2)\n\n        dice = (2.0 * intersection + self.smooth) / (denominator + self.smooth)\n        loss = 1.0 - dice.mean()\n        return loss\n\n\nclass BCEDiceLoss(nn.Module):\n    def __init__(self, bce_weight=0.5):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss()\n        self.dice = DiceLoss()\n        self.bce_weight = bce_weight\n\n    def forward(self, logits, targets):\n        bce = self.bce(logits, targets)\n        dice = self.dice(logits, targets)\n        return self.bce_weight * bce + (1 - self.bce_weight) * dice\n\nprint(\"done\")","metadata":{"_uuid":"d160eba3-27fd-42c9-8310-34ecab6e3e28","_cell_guid":"6d2e9e6a-c9b9-4c39-9b45-9c3b96c2e8b2","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:32:49.645992Z","iopub.execute_input":"2026-07-06T03:32:49.646903Z","iopub.status.idle":"2026-07-06T03:32:49.654330Z","shell.execute_reply.started":"2026-07-06T03:32:49.646871Z","shell.execute_reply":"2026-07-06T03:32:49.653752Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef dice_score(logits, targets, threshold=0.5, eps=1e-7):\n    probs = torch.sigmoid(logits)\n    preds = (probs > threshold).float()\n\n    preds = preds.view(preds.size(0), preds.size(1), -1)\n    targets = targets.view(targets.size(0), targets.size(1), -1)\n\n    intersection = (preds * targets).sum(dim=2)\n    denominator = preds.sum(dim=2) + targets.sum(dim=2)\n\n    dice = (2 * intersection + eps) / (denominator + eps)\n    return dice.mean().item()\nprint(\"done\")","metadata":{"_uuid":"2db6128e-b20f-4430-9d5d-0d839fb44022","_cell_guid":"72071838-8dee-4082-a8a6-0a86a663858f","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:32:54.189003Z","iopub.execute_input":"2026-07-06T03:32:54.189274Z","iopub.status.idle":"2026-07-06T03:32:54.195336Z","shell.execute_reply.started":"2026-07-06T03:32:54.189250Z","shell.execute_reply":"2026-07-06T03:32:54.194654Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef find_best_thresholds(model, loader, thresholds=None, device=DEVICE, eps=1e-7):\n    if thresholds is None:\n        thresholds = np.arange(0.10, 0.91, 0.05)\n\n    model.eval()\n    thresholds = np.asarray(thresholds, dtype=np.float32)\n    intersections = torch.zeros((len(thresholds), len(CLASS_NAMES)), device=device)\n    denominators = torch.zeros((len(thresholds), len(CLASS_NAMES)), device=device)\n\n    for images, masks, _ in loader:\n        images = images.to(device)\n        masks = masks.to(device)\n\n        probs = torch.sigmoid(model(images))\n        probs = probs.view(probs.size(0), probs.size(1), -1)\n        masks = masks.view(masks.size(0), masks.size(1), -1)\n\n        for t_idx, threshold in enumerate(thresholds):\n            preds = (probs > float(threshold)).float()\n            intersections[t_idx] += (preds * masks).sum(dim=(0, 2))\n            denominators[t_idx] += preds.sum(dim=(0, 2)) + masks.sum(dim=(0, 2))\n\n    scores = ((2 * intersections + eps) / (denominators + eps)).detach().cpu().numpy()\n\n    best_indices = scores.argmax(axis=0)\n    best_thresholds = thresholds[best_indices]\n    best_scores = scores[best_indices, np.arange(len(CLASS_NAMES))]\n\n    return best_thresholds, best_scores, thresholds, scores\nprint(\"done\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-06T01:44:50.918574Z","iopub.execute_input":"2026-07-06T01:44:50.919395Z","iopub.status.idle":"2026-07-06T01:44:50.927393Z","shell.execute_reply.started":"2026-07-06T01:44:50.919347Z","shell.execute_reply":"2026-07-06T01:44:50.926453Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q segmentation-models-pytorch","metadata":{"_uuid":"7ac36780-6d58-4d22-812a-374857cfe39b","_cell_guid":"cf356223-bbd9-4465-83a7-cbec617f5dd0","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T01:44:56.282297Z","iopub.execute_input":"2026-07-06T01:44:56.282596Z","iopub.status.idle":"2026-07-06T01:45:02.042822Z","shell.execute_reply.started":"2026-07-06T01:44:56.282572Z","shell.execute_reply":"2026-07-06T01:45:02.041764Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import segmentation_models_pytorch as smp\n\nmodel = smp.Unet(\n    encoder_name=\"resnet18\",\n    encoder_weights=\"imagenet\",\n    in_channels=3,\n    classes=4,\n    activation=None\n).to(DEVICE)\n\ncriterion = BCEDiceLoss(bce_weight=0.5)\n\noptimizer = torch.optim.AdamW(\n    [\n        {\"params\": model.encoder.parameters(), \"lr\": 1e-4},\n        {\"params\": model.decoder.parameters(), \"lr\": 3e-4},\n        {\"params\": model.segmentation_head.parameters(), \"lr\": 3e-4},\n    ],\n    weight_decay=1e-4\n)","metadata":{"_uuid":"dd5a7ebf-5002-4c89-9f4c-c4543c103823","_cell_guid":"0e3245b9-f342-4c75-a5a2-e85fb6c09b97","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:33:06.676574Z","iopub.execute_input":"2026-07-06T03:33:06.676840Z","iopub.status.idle":"2026-07-06T03:33:06.898351Z","shell.execute_reply.started":"2026-07-06T03:33:06.676816Z","shell.execute_reply":"2026-07-06T03:33:06.897720Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"scaler = torch.amp.GradScaler(\"cuda\")\n\ndef train_one_epoch(model, loader, optimizer, criterion, device):\n    model.train()\n    total_loss = 0.0\n\n    for images, masks, _ in loader:\n        images = images.to(device,non_blocking=True)\n        masks = masks.to(device,non_blocking=True)\n\n        optimizer.zero_grad()\n        with torch.autocast(device_type=\"cuda\", dtype=torch.float16):\n            logits = model(images)\n            loss = criterion(logits, masks)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        total_loss += loss.item() * images.size(0)\n\n    return total_loss / len(loader.dataset)\n\n\n@torch.no_grad()\ndef validate_one_epoch(model, loader, criterion, device):\n    model.eval()\n    total_loss = 0.0\n    total_dice = 0.0\n\n    for images, masks, _ in loader:\n        images = images.to(device)\n        masks = masks.to(device)\n\n        logits = model(images)\n        loss = criterion(logits, masks)\n        dsc = dice_score(logits, masks)\n\n        total_loss += loss.item() * images.size(0)\n        total_dice += dsc * images.size(0)\n\n    val_loss = total_loss / len(loader.dataset)\n    val_dice = total_dice / len(loader.dataset)\n    return val_loss, val_dice\n\nprint(\"done\")","metadata":{"_uuid":"7762e78a-3154-4f87-bf6c-3b637032b545","_cell_guid":"d251b299-aeb6-4374-9f89-fe83d20b3ac5","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:33:11.697217Z","iopub.execute_input":"2026-07-06T03:33:11.697881Z","iopub.status.idle":"2026-07-06T03:33:11.705870Z","shell.execute_reply.started":"2026-07-06T03:33:11.697843Z","shell.execute_reply":"2026-07-06T03:33:11.705011Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"max\",\n    factor=0.5,\n    patience=2\n)\n\nbest_score = -1\n\npatience = 6\ncounter = 0\nmin_delta = 0.001\n\nfor epoch in range(NUM_EPOCHS):\n    t0=time.time()\n    train_loss = train_one_epoch(model, train_loader, optimizer, criterion, DEVICE)\n    val_loss, val_dice = validate_one_epoch(model, val_loader, criterion, DEVICE)\n\n    scheduler.step(val_dice)\n    current_lrs = [group[\"lr\"] for group in optimizer.param_groups]\n    dt = time.time()-t0\n\n    print(f\"Epoch {epoch+1}/{NUM_EPOCHS} | time: {dt:.1f}s | \"\n          f\"lr={[round(lr, 6) for lr in current_lrs]} | \"\n          f\"train_loss={train_loss:.4f} | \"\n          f\"val_loss={val_loss:.4f} | \"\n          f\"val_dice={val_dice:.4f}\")\n\n    # 每轮保存 latest\n    torch.save({\n        \"epoch\": epoch,\n        \"model_state_dict\": model.state_dict(),\n        \"optimizer_state_dict\": optimizer.state_dict(),\n        \"best_score\": best_score,\n    }, \"latest_checkpoint.pth\")\n\n    # 按 best score 保存\n    if val_dice > best_score + min_delta:\n        best_score = val_dice\n        counter = 0\n\n        torch.save(model.state_dict(), \"best_unet_cloud.pth\")\n        torch.save({\n            \"epoch\": epoch,\n            \"model_state_dict\": model.state_dict(),\n            \"optimizer_state_dict\": optimizer.state_dict(),\n            \"best_score\": best_score,\n        }, \"best_checkpoint.pth\")\n    else:\n        counter += 1\n\n    if counter >= patience:\n        print(\"Early stopping triggered\")\n        break","metadata":{"_uuid":"5c65ee88-0f16-4c10-ab7a-c54acaee1f69","_cell_guid":"77e68589-de31-4f31-9795-e4593cfec250","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-07-06T03:33:30.057964Z","iopub.execute_input":"2026-07-06T03:33:30.058286Z","iopub.status.idle":"2026-07-06T04:41:13.021360Z","shell.execute_reply.started":"2026-07-06T03:33:30.058259Z","shell.execute_reply":"2026-07-06T04:41:13.020172Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def dice_binary_batch(preds, targets, eps=1e-7):\n    \"\"\"\n    preds:   [B, H, W], bool or uint8\n    targets: [B, H, W], bool or uint8\n    return:  [B]\n    \"\"\"\n    B = preds.shape[0]\n\n    preds = preds.reshape(B, -1).astype(np.uint8)\n    targets = targets.reshape(B, -1).astype(np.uint8)\n\n    intersection = (preds * targets).sum(axis=1)\n    pred_sum = preds.sum(axis=1)\n    target_sum = targets.sum(axis=1)\n\n    dice = (2.0 * intersection + eps) / (pred_sum + target_sum + eps)\n\n    return dice\n\n\n@torch.no_grad()\ndef search_threshold_min_area(\n    model,\n    loader,\n    threshold_grid=None,\n    min_area_grid=None,\n    device=DEVICE\n):\n    if threshold_grid is None:\n        threshold_grid = [0.30, 0.35, 0.40, 0.45, 0.50, 0.55, 0.60, 0.65, 0.70]\n\n    if min_area_grid is None:\n        min_area_grid = [0, 500, 1000, 2000, 3000, 5000, 8000, 10000]\n\n    threshold_grid = np.array(threshold_grid, dtype=np.float32)\n    min_area_grid = np.array(min_area_grid, dtype=np.int32)\n\n    n_cls = len(CLASS_NAMES)\n    n_thr = len(threshold_grid)\n    n_area = len(min_area_grid)\n\n    score_sum = np.zeros((n_cls, n_thr, n_area), dtype=np.float64)\n    score_count = np.zeros((n_cls, n_thr, n_area), dtype=np.float64)\n\n    model.eval()\n\n    for images, masks, _ in loader:\n        images = images.to(device, non_blocking=True)\n\n        logits = model(images)\n        probs = torch.sigmoid(logits).cpu().numpy()  # [B, 4, H, W]\n        targets = masks.cpu().numpy()                # [B, 4, H, W]\n\n        B = probs.shape[0]\n\n        for c in range(n_cls):\n            gt = targets[:, c] > 0.5\n\n            for ti, thr in enumerate(threshold_grid):\n                pred_base = probs[:, c] > thr\n                pred_area = pred_base.reshape(B, -1).sum(axis=1)\n\n                for ai, min_area in enumerate(min_area_grid):\n                    pred = pred_base.copy()\n\n                    # 如果某张图该类预测面积太小，就直接置空\n                    small_mask = pred_area < min_area\n                    pred[small_mask] = False\n\n                    dice = dice_binary_batch(pred, gt)\n\n                    score_sum[c, ti, ai] += dice.sum()\n                    score_count[c, ti, ai] += len(dice)\n\n    mean_scores = score_sum / score_count\n\n    best_thresholds = []\n    best_min_areas = []\n    best_scores = []\n    rows = []\n\n    for c, class_name in enumerate(CLASS_NAMES):\n        best_idx = np.unravel_index(\n            np.argmax(mean_scores[c]),\n            mean_scores[c].shape\n        )\n\n        ti, ai = best_idx\n\n        best_thr = float(threshold_grid[ti])\n        best_area = int(min_area_grid[ai])\n        best_score = float(mean_scores[c, ti, ai])\n\n        best_thresholds.append(best_thr)\n        best_min_areas.append(best_area)\n        best_scores.append(best_score)\n\n        rows.append({\n            \"class\": class_name,\n            \"best_threshold\": best_thr,\n            \"best_min_area\": best_area,\n            \"best_dice\": best_score\n        })\n\n    result_df = pd.DataFrame(rows)\n\n    print(result_df)\n    print()\n    print(\"Estimated overall val dice:\", np.mean(best_scores))\n\n    return result_df, best_thresholds, best_min_areas, mean_scores\n\n\n# 加载训练过程中保存的最佳模型\nmodel.load_state_dict(torch.load(\"best_unet_cloud.pth\", map_location=DEVICE))\nmodel.eval()\n\nbest_param_df, best_thresholds, best_min_areas, all_scores = search_threshold_min_area(\n    model,\n    val_loader\n)\n\nprint(\"best_thresholds =\", best_thresholds)\nprint(\"best_min_areas =\", best_min_areas)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-06T04:45:01.729694Z","iopub.execute_input":"2026-07-06T04:45:01.730578Z","iopub.status.idle":"2026-07-06T04:47:35.503786Z","shell.execute_reply.started":"2026-07-06T04:45:01.730531Z","shell.execute_reply":"2026-07-06T04:47:35.502816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, image_dir, transform=None):\n        self.image_dir = Path(image_dir)\n        self.image_paths = sorted(list(self.image_dir.glob(\"*.jpg\")))\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        image_path = self.image_paths[idx]\n        image_name = image_path.name\n\n        image = cv2.imread(str(image_path))\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        orig_h, orig_w = image.shape[:2]\n\n        if self.transform is not None:\n            transformed = self.transform(image=image)\n            image = transformed[\"image\"]\n\n        image = np.transpose(image, (2, 0, 1)).astype(np.float32)\n        image = torch.tensor(image, dtype=torch.float32)\n\n        return image, image_name, orig_h, orig_w","metadata":{"_uuid":"91e895aa-1e96-4d93-8ae0-e8094db6fa44","_cell_guid":"b95887c7-61d8-4529-b835-18ee479eb670","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-07-06T04:48:25.422268Z","iopub.execute_input":"2026-07-06T04:48:25.422912Z","iopub.status.idle":"2026-07-06T04:48:25.430217Z","shell.execute_reply.started":"2026-07-06T04:48:25.422878Z","shell.execute_reply":"2026-07-06T04:48:25.429270Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset = TestDataset(\n    TEST_IMG_DIR,\n    transform=val_transform\n)\ntest_loader = DataLoader(test_dataset, batch_size=4, shuffle=False)\n\nmodel.load_state_dict(torch.load(\"best_unet_cloud.pth\", map_location=DEVICE))\nmodel.eval()","metadata":{"_uuid":"a5d97cf2-c80d-4a8a-bd88-58460ac0fe21","_cell_guid":"7a9ebe8a-cbbc-43f1-909f-85af02a54a47","trusted":true,"collapsed":true,"jupyter":{"outputs_hidden":true},"execution":{"iopub.status.busy":"2026-07-06T04:48:28.407362Z","iopub.execute_input":"2026-07-06T04:48:28.408223Z","iopub.status.idle":"2026-07-06T04:48:28.534753Z","shell.execute_reply.started":"2026-07-06T04:48:28.408192Z","shell.execute_reply":"2026-07-06T04:48:28.534108Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SUB_H, SUB_W = 350, 525\n\n# 使用 validation set 搜索出来的结果\nclass_thresholds = best_thresholds\nclass_min_areas = best_min_areas\n\nprint(\"class_thresholds:\", class_thresholds)\nprint(\"class_min_areas:\", class_min_areas)\n\n@torch.no_grad()\ndef predict_and_build_submission(model, loader):\n    model.eval()\n    submission_rows = []\n\n    for images, image_names, orig_hs, orig_ws in loader:\n        images = images.to(DEVICE)\n\n        logits = model(images)\n        probs = torch.sigmoid(logits).cpu().numpy()  # [B, 4, H, W]\n\n        for b in range(len(image_names)):\n            image_name = image_names[b]\n            pred = probs[b]\n\n            for class_idx, class_name in enumerate(CLASS_NAMES):\n                mask = pred[class_idx]\n\n                # 固定 resize 到 submission 要求尺寸\n                mask = cv2.resize(\n                    mask,\n                    (SUB_W, SUB_H),\n                    interpolation=cv2.INTER_LINEAR\n                )\n\n                threshold = class_thresholds[class_idx]\n                min_area = class_min_areas[class_idx]\n\n                mask_bin = (mask > threshold).astype(np.uint8)\n\n                if mask_bin.sum() < min_area:\n                    mask_bin[:] = 0\n\n                rle = \"\" if mask_bin.sum() == 0 else rle_encode(mask_bin)\n\n                submission_rows.append({\n                    \"Image_Label\": f\"{image_name}_{class_name}\",\n                    \"EncodedPixels\": rle\n                })\n\n    return pd.DataFrame(submission_rows)\n\n\nsubmission = predict_and_build_submission(model, test_loader)\nsubmission.to_csv(\"submission.csv\", index=False)\n\nprint(submission.shape)\nprint(submission[\"Image_Label\"].nunique())\nprint(submission[\"Image_Label\"].duplicated().sum())\nprint(submission.head())","metadata":{"_uuid":"d5e4e1e3-da0d-4df6-a8f8-d8aef45a23fd","_cell_guid":"1df64458-a325-4c47-a3f6-2c2a80f3c310","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-07-06T04:48:48.541840Z","iopub.execute_input":"2026-07-06T04:48:48.542703Z","iopub.status.idle":"2026-07-06T04:51:10.147517Z","shell.execute_reply.started":"2026-07-06T04:48:48.542669Z","shell.execute_reply":"2026-07-06T04:51:10.146780Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission[\"rle_len\"] = submission[\"EncodedPixels\"].fillna(\"\").apply(len)\nsubmission[\"is_non_empty\"] = submission[\"EncodedPixels\"].fillna(\"\").apply(lambda x: len(x) > 0)\nsubmission[\"class\"] = submission[\"Image_Label\"].apply(lambda x: x.split(\"_\")[-1])\n\nprint(\"总行数:\", len(submission))\nprint(\"非空比例:\", submission[\"is_non_empty\"].mean())\nprint(\"平均 RLE 长度:\", submission[\"rle_len\"].mean())\nprint(\"最大 RLE 长度:\", submission[\"rle_len\"].max())\n\nprint(submission.groupby(\"class\")[[\"is_non_empty\", \"rle_len\"]].mean())\nprint(submission.groupby(\"class\")[\"rle_len\"].max())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-06T04:51:35.122384Z","iopub.execute_input":"2026-07-06T04:51:35.122674Z","iopub.status.idle":"2026-07-06T04:51:35.159626Z","shell.execute_reply.started":"2026-07-06T04:51:35.122649Z","shell.execute_reply":"2026-07-06T04:51:35.158996Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null}]}