{"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":9988,"databundleVersionId":868324}],"dockerImageVersionId":31328,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\n\ndata_root: str = \"/kaggle/input/competitions/airbus-ship-detection\"\nprint(f\"The root of the data: '{data_root}'\")\nprint(f\"Items of the root dir: {os.listdir(data_root)}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:12:32.698437Z","iopub.execute_input":"2026-03-31T16:12:32.698971Z","iopub.status.idle":"2026-03-31T16:12:32.703865Z","shell.execute_reply.started":"2026-03-31T16:12:32.698937Z","shell.execute_reply":"2026-03-31T16:12:32.703114Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nimport albumentations\nimport numpy as np\nimport pandas as pd\nfrom torch.utils.data import Dataset\nimport tqdm\n\nclass csv: ... # place holder for type annotation\nclass AirbusShipDataset(Dataset):\n\n    def __init__(self, imgs_root: str, masks_root: csv):\n        \n        super().__init__()\n        self.imgs_root = imgs_root\n        \n        new_masks_root: str = \"/kaggle/working/_inter_mask\"\n        os.makedirs(new_masks_root, exist_ok = True)\n        self.masks_root = new_masks_root\n        self.mask_infos = self._prepMasks(masks_root, new_masks_root)\n        \n    def _prepMasks(self, masks_root: str, to: str) -> pd.DataFrame:\n      \n        masks = pd.read_csv(masks_root)\n        os.makedirs(to, exist_ok=True)\n    \n        data_list = []\n        groups = masks.groupby(\"ImageId\")\n        \n        for image_id, group in tqdm.tqdm(groups):\n            rles = group[\"EncodedPixels\"].dropna()\n            ship_count = len(rles)\n        \n            if ship_count > 0:\n                mask = np.zeros(768 * 768, dtype=np.uint8)\n                \n                for rle in rles:\n                    starts_lens = rle.split()\n                    starts = np.asarray(starts_lens[::2], dtype=int) - 1\n                    lens = np.asarray(starts_lens[1::2], dtype=int)\n                    \n                    for start, length in zip(starts, lens):\n                        mask[start: start+length] = 1 \n\n                mask_path = os.path.join(to, f\"{image_id}.npy\")\n                np.save(mask_path, np.packbits(mask))\n            \n                data_list.append({\n                    \"ImageId\": image_id,\n                    \"ShipCounts\": ship_count,\n                    \"MaskPath\": mask_path\n                })\n            else:\n                data_list.append({\n                    \"ImageId\": image_id,\n                    \"ShipCounts\": 0,\n                    \"MaskPath\": None\n                })\n            \n        return pd.DataFrame(data_list)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:12:32.705285Z","iopub.execute_input":"2026-03-31T16:12:32.705760Z","iopub.status.idle":"2026-03-31T16:12:32.718679Z","shell.execute_reply.started":"2026-03-31T16:12:32.705724Z","shell.execute_reply":"2026-03-31T16:12:32.717936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = AirbusShipDataset(imgs_root = os.path.join(data_root, \"train_v2\"),\n                            masks_root = os.path.join(data_root, \"train_ship_segmentations_v2.csv\"))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:12:32.719553Z","iopub.execute_input":"2026-03-31T16:12:32.719843Z","iopub.status.idle":"2026-03-31T16:13:22.228848Z","shell.execute_reply.started":"2026-03-31T16:12:32.719812Z","shell.execute_reply":"2026-03-31T16:13:22.228192Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\n\nclass ImageDataset(Dataset):\n\n    def __init__(self, src: AirbusShipDataset, indices: list[int], transform = None):\n        super().__init__()\n        \n        self.imgs_root = src.imgs_root \n        self.masks_root = src.masks_root\n        \n        self.imgids = src.mask_infos[\"ImageId\"].values\n        self.mask_paths = src.mask_infos[\"MaskPath\"].values\n        self.ship_counts  = src.mask_infos[\"ShipCounts\"].values\n        \n        self.indices = indices\n        self.transform = transform\n        \n    def __len__(self) -> int:\n        return len(self.indices)\n    \n    def __getitem__(self, index: int) -> tuple:\n\n        index = self.indices[index]\n        image_path = os.path.join(self.imgs_root ,  self.imgids[index])\n        mask_path = self.mask_paths[index]\n            \n        image = cv2.imread(image_path, cv2.IMREAD_COLOR_RGB)\n\n        mask = None\n        if mask_path is None:\n            mask = np.zeros(768 * 768)\n        else:\n            mask = np.unpackbits(np.load(mask_path))[: 768 * 768]\n        mask = mask.reshape((768, 768), order = \"F\")\n        \n        ship_counts = self.ship_counts[index]\n        \n        if self.transform:\n            argmt = self.transform(image = image, mask = mask)\n            image = argmt[\"image\"]\n            mask = argmt[\"mask\"]\n            mask = mask.float()\n            \n        n_ship_pixels = mask.sum().item()\n\n        n_ship_pixels = torch.tensor(n_ship_pixels, dtype = torch.float32)\n        ship_counts = torch.tensor(ship_counts, dtype = torch.float32)\n        return image, mask, ship_counts, n_ship_pixels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:22.229820Z","iopub.execute_input":"2026-03-31T16:13:22.230383Z","iopub.status.idle":"2026-03-31T16:13:22.238235Z","shell.execute_reply.started":"2026-03-31T16:13:22.230357Z","shell.execute_reply":"2026-03-31T16:13:22.237473Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import albumentations as A\n\ntrain_tf: A.Compose = A.Compose([\n    A.RandomResizedCrop((512, 512), interpolation= cv2.INTER_LINEAR, scale = (0.6 , 1.0)),\n    A.HorizontalFlip(p = 0.5),\n    A.VerticalFlip(p = 0.5),\n    A.RandomRotate90(p = 0.4),\n    A.ShiftScaleRotate((-0.1, 0.1), (-0.1, 0.2), (-30, 30)),\n    A.GridDistortion(num_steps = 10),\n    A.ElasticTransform(alpha = 2, sigma = 30),\n    A.ColorJitter((0.7 , 1.0), (0.7 , 1.0), (0.8 , 1.2), hue = (-0.1, 0.1), p = 0.3),\n    A.OneOf([\n         A.GaussNoise( std_range = (0.03, 0.07), p = 0.2),\n         A.ISONoise( p = 0.4),\n    ], p = 1.0),\n    A.OneOf([\n        A.GaussianBlur((3, 5), p = 0.5),\n        A.MotionBlur((3, 5), p = 0.5),\n        A.MedianBlur((3, 5), p = 0.5)\n    ], p = 1.0),\n    A.OneOf([\n         A.RandomFog(fog_coef_range  = (0.1, 0.3), p = 0.4),\n         A.RandomShadow(num_shadows_limit = (1,3), p = 0.1),\n         A.Solarize(threshold_range = (0.8, 0.9), p = 0.1),\n    ], p = 1.0),\n    \n    A.CoarseDropout((4, 8), (8 , 16), (8 , 16), p = 0.07),\n    A.Normalize(),\n    A.pytorch.transforms.ToTensorV2()\n])\n\ntest_tf: A.Compose = A.Compose([\n    A.CenterCrop(height = 512, width = 512, border_mode = cv2.BORDER_REFLECT_101),\n    A.Normalize(),\n    A.pytorch.transforms.ToTensorV2()\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:22.240118Z","iopub.execute_input":"2026-03-31T16:13:22.240347Z","iopub.status.idle":"2026-03-31T16:13:22.625045Z","shell.execute_reply.started":"2026-03-31T16:13:22.240320Z","shell.execute_reply":"2026-03-31T16:13:22.624198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n\nall_indices = list(range(len(dataset.mask_infos)))\nhas_ships = dataset.mask_infos['ShipCounts'] > 0\n\ntrain_indices, val_indices, _ , _ = train_test_split(\n    all_indices,\n    all_indices,\n    train_size = 0.7,\n    test_size = 0.3,\n    random_state = 13,\n    shuffle = True,\n    stratify = has_ships\n)\n\nprint(f\"Train indices peak: {train_indices[:5]}, size = {len(train_indices)}\")\nprint(f\"Validation indices peak: {val_indices[:5]}, size = {len(val_indices)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:22.625876Z","iopub.execute_input":"2026-03-31T16:13:22.626079Z","iopub.status.idle":"2026-03-31T16:13:22.869807Z","shell.execute_reply.started":"2026-03-31T16:13:22.626059Z","shell.execute_reply":"2026-03-31T16:13:22.869029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TrainingDataset = ImageDataset(dataset, train_indices)\nValidationDataset = ImageDataset(dataset, val_indices)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:22.870819Z","iopub.execute_input":"2026-03-31T16:13:22.871405Z","iopub.status.idle":"2026-03-31T16:13:22.874593Z","shell.execute_reply.started":"2026-03-31T16:13:22.871380Z","shell.execute_reply":"2026-03-31T16:13:22.874038Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport random\nimport torch\n\nshape: tuple[int, int] = (2, 4)\nfig, axes = plt.subplots(shape[0], shape[1])\n\nfor j in range(shape[1]):\n    random_index = random.randint(0, len(TrainingDataset))\n    image, mask, sc, nsp = TrainingDataset[random_index]\n    axes[0, j].imshow(image, cmap = \"viridis\")\n    axes[1, j].imshow(mask, cmap = \"grey\")\n    axes[0, j].set_title(f\"Image no: {random_index},\\nsc: {sc},\\nnsp: {nsp}\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:22.875502Z","iopub.execute_input":"2026-03-31T16:13:22.875841Z","iopub.status.idle":"2026-03-31T16:13:23.880603Z","shell.execute_reply.started":"2026-03-31T16:13:22.875812Z","shell.execute_reply":"2026-03-31T16:13:23.879906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\n\nfig_2, axes_2 = plt.subplots(2, 2)\nmeans = torch.Tensor([[[0.485, 0.456, 0.406]]])\nstds = torch.Tensor([[[0.229, 0.224, 0.225]]])\n\ndef Denormalize(image: torch.Tensor) -> torch.Tensor:\n    image = image * stds.permute(2, 0, 1) + means.permute(2, 0, 1)\n    return image\n\nrandom_index = random.randint(0 , len(TrainingDataset))\nimg, mask, sc, nps = TrainingDataset[random_index]\naxes_2[0, 0].imshow(img)\naxes_2[1, 0].imshow(mask, cmap = \"grey\")\n\nargmt = train_tf(image = img, mask = mask)\naxes_2[0, 1].imshow(Denormalize(argmt[\"image\"]).permute(1, 2, 0))\naxes_2[1, 1].imshow(argmt[\"mask\"], cmap = \"grey\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:23.881527Z","iopub.execute_input":"2026-03-31T16:13:23.881840Z","iopub.status.idle":"2026-03-31T16:13:24.471090Z","shell.execute_reply.started":"2026-03-31T16:13:23.881800Z","shell.execute_reply":"2026-03-31T16:13:24.470227Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TrainingDataset.transform = train_tf\nValidationDataset.transform = test_tf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:24.472144Z","iopub.execute_input":"2026-03-31T16:13:24.472432Z","iopub.status.idle":"2026-03-31T16:13:24.476128Z","shell.execute_reply.started":"2026-03-31T16:13:24.472369Z","shell.execute_reply":"2026-03-31T16:13:24.475544Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchvision.models as models\nmodels.resnext50_32x4d()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:24.476970Z","iopub.execute_input":"2026-03-31T16:13:24.477209Z","iopub.status.idle":"2026-03-31T16:13:24.773059Z","shell.execute_reply.started":"2026-03-31T16:13:24.477188Z","shell.execute_reply":"2026-03-31T16:13:24.772445Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:24.773906Z","iopub.execute_input":"2026-03-31T16:13:24.774168Z","iopub.status.idle":"2026-03-31T16:13:24.777759Z","shell.execute_reply.started":"2026-03-31T16:13:24.774146Z","shell.execute_reply":"2026-03-31T16:13:24.777065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch import nn\n\n\n\nclass AttentionGate(nn.Module):\n\n    def __init__(self, nxic: int, nsic: int, niic: int):\n        super().__init__()\n        self.xConv = nn.Sequential(\n            nn.Conv2d(nxic, niic, kernel_size = 1, stride = 1, padding = 0, bias = True),\n            nn.BatchNorm2d(niic)\n        )\n        \n        self.skipConv = nn.Sequential(\n            nn.Conv2d(nsic, niic, kernel_size = 1, stride = 1, padding = 0, bias = True),\n            nn.BatchNorm2d(niic)\n        )\n        \n        self.final_conv = nn.Sequential(\n            nn.Conv2d(niic, 1, kernel_size = 1, stride = 1, padding = 0, bias = True),\n            nn.BatchNorm2d(1),\n            nn.Sigmoid()\n        )\n        self.relu = nn.ReLU(inplace = True)\n        \n    def forward(self, x: torch.Tensor, skip: torch.Tensor) -> torch.Tensor:\n        \n        x_out = self.xConv(x)\n        skip_out = self.skipConv(skip)\n        weight = self.final_conv(self.relu(x_out + skip_out))\n        return skip * weight\n        \nclass DecoderBlock(nn.Module):\n\n    def __init__(self, nxic: int , nsic: int, noc: int):\n        super().__init__()\n        \n        self.attention = AttentionGate(nxic , nsic, (nxic + nsic)//2)\n        self.convblock = nn.Sequential(\n            nn.Conv2d(nxic + nsic , noc , kernel_size = 3, stride = 1, padding = 1, bias = False),\n            nn.BatchNorm2d(noc),\n            nn.ELU(inplace = True),\n\n            nn.Conv2d(noc, noc, kernel_size = 3, stride = 1, padding = 1, bias= False),\n            nn.BatchNorm2d(noc),\n            nn.ELU(inplace = True)\n        )\n    def forward(self, x: torch.Tensor , skip: torch.Tensor) -> torch.Tensor:\n        \n        skip = self.attention(x, skip)\n        x = torch.cat([x, skip], dim = 1)\n        x = self.convblock(x)\n        \n        return x\n\nclass AttentionUnet(nn.Module):\n\n    def __init__(self):\n        \n        super().__init__()\n        weights = models.ResNeXt50_32X4D_Weights.DEFAULT\n        self.backbone = models.resnext50_32x4d(weights = weights)\n        \n        self.backbone.fc = nn.Identity()\n        self.backbone.layer4 = nn.Identity()\n        \n        for param in self.backbone.parameters():\n            param.requires_grad = False\n        \n        for param in self.backbone.layer3.parameters():\n            param.requires_grad = True\n            \n        self.bottleneck = nn.Sequential(\n            nn.Conv2d(1024, 512, kernel_size = 3, stride = 1, padding = 1, bias = False),\n            nn.BatchNorm2d(512),\n            nn.ELU(inplace = True),\n\n            nn.Conv2d(512, 1024, kernel_size = 3, stride = 1, padding = 1, bias = False),\n            nn.BatchNorm2d(1024),\n            nn.ELU(inplace = True)\n        )\n        \n        self.decoder_l3 = DecoderBlock(1024, 1024, 512)\n        self.decoder_l2 = DecoderBlock(512, 512, 256)\n        self.decoder_l1 = DecoderBlock(256, 256, 64)\n        self.decoder_l0 = DecoderBlock(64, 64, 16)\n        self.final_conv = nn.Conv2d(16, 1, kernel_size = 3, stride = 1, padding = 1, bias = True)\n        self.upsample = nn.Upsample(scale_factor = 2, mode = \"bilinear\", align_corners = True)\n        \n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n\n        b1 = self.backbone.conv1(x)     # 256 x 256\n        b1 = self.backbone.bn1(b1)      # 256 x 256\n        b1 = self.backbone.relu(b1)     # 256 x 256\n        p1 = self.backbone.maxpool(b1)  # 128 x 128\n\n        b2 = self.backbone.layer1(p1)   # 128 x 128\n        b3 = self.backbone.layer2(b2)   # 64 x 64\n        b4 = self.backbone.layer3(b3)   # 32 x 32\n        \n        x  = self.bottleneck(b4)        # 32 x 32\n        \n        x  = self.decoder_l3(x, b4)   \n        x  = self.upsample(x)          # 64 x 64\n        x  = self.decoder_l2(x, b3)    \n        x  = self.upsample(x)          # 128 x 128\n        x  = self.decoder_l1(x, b2)    \n        x  = self.upsample(x)          # 256 x 256\n        x  = self.decoder_l0(x, b1)    \n        x  = self.upsample(x)          # 512 x 512\n        x  = self.final_conv(x)\n        return x\n\ndef DefaultWeighter(ns: torch.Tensor, nsp: torch.Tensor, positive_weight: float = 1.0) -> torch.Tensor:\n\n    with torch.no_grad():\n        mask = ns > 0\n        weights = torch.ones_like(ns, device = device).to(torch.float32)\n        weights[mask] = positive_weight\n        weights[mask] += torch.log( (512 * 512 * ns[mask]) / nsp[mask])\n        weights[mask] = torch.clamp(weights[mask], min = 1.0, max = 20.0)\n    return weights\n    \nclass SegmentLoss(nn.Module):\n\n    def __init__(self, \n                 gamma: float = 2.0,\n                 smoothing: float = 1.0,\n                 focal_weight: float = 2.0,\n                 iou_weight: float= 1.0,\n                 pos_weight: float = 4.0,\n                 positive_weight: float = 10.0,\n                 sample_weighter = DefaultWeighter):\n        \n        super().__init__()\n        \n        self.gamma = gamma\n        self.smoothing = smoothing\n        self.sample_weighter = sample_weighter\n        self.focal_weight  = focal_weight\n        self.iou_weight = iou_weight\n        self.pos_weight = pos_weight\n        self.positive_weight = positive_weight\n        \n        self.flatten = nn.Flatten()\n        self.bce = nn.BCEWithLogitsLoss(reduction = \"none\")\n        self.sig = nn.Sigmoid()\n        \n    def forward(self, preds: torch.Tensor, \n                      targets: torch.Tensor, \n                      ns: torch.Tensor, \n                      nsp: torch.Tensor) -> tuple[torch.Tensor, float, float]:\n        \n        with torch.amp.autocast('cuda', dtype = torch.float32):\n\n            # This will flatten the preds and targets (B, 1, W, H) => (B, H * W)\n            preds_flat = self.flatten(preds)\n            targets_flat = self.flatten(targets)\n\n            # Apply bce with raw logits and find focal loss\n            bce = self.bce(preds_flat, targets_flat)\n            pt = torch.exp(-bce)\n            preds_flat = self.sig(preds_flat)\n            alpha = torch.abs((preds_flat > 0.5).to(torch.float32) - targets_flat).sum(dim = 1)\n            alpha = torch.clamp(alpha, min = 1, max = 10.0)\n            focal_loss = ((1 - pt)**self.gamma * bce).mean(dim = 1)\n            \n            # Track the raw loss\n            raw_focal    = focal_loss.mean().item()\n            # Then weighten the loss\n            focal_loss = alpha * focal_loss\n            \n            # Find the iou loss\n            intersection = (preds_flat * targets_flat).sum(dim = 1)\n            total        = (preds_flat + targets_flat).sum(dim = 1)\n            union        = total - intersection\n            iou_loss     = 1 - (intersection + self.smoothing )/(union + self.smoothing)\n\n            # Track the raw loss , necessary for later visualization\n            raw_iou      = iou_loss.mean().item()\n\n            # Weighten the focal loss and iou loss seperately\n            focal_loss   *= self.focal_weight\n            iou_loss     *= self.iou_weight\n\n            # Find the weights based on num_ships and num_ship_pixels\n            if self.sample_weighter is not None:\n                weights: torch.Tensor = self.sample_weighter(ns, nsp, self.positive_weight)\n            else:\n                weights: torch.Tensor = nn.ones_like(preds, device = device)\n                \n            # Finally weighten and return!\n        return (weights * (iou_loss + focal_loss)).mean(), raw_iou, raw_focal","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:24.778624Z","iopub.execute_input":"2026-03-31T16:13:24.778982Z","iopub.status.idle":"2026-03-31T16:13:24.803367Z","shell.execute_reply.started":"2026-03-31T16:13:24.778940Z","shell.execute_reply":"2026-03-31T16:13:24.802667Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.init as init\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torch.amp import GradScaler\n\nnum_epochs = 3\ndef init_layer(layer: nn.Module) -> None:\n\n    if isinstance(layer, nn.Conv2d):\n        init.xavier_uniform_(layer.weight)\n        if layer.bias is not None:\n            init.zeros_(layer.bias)\n\nmodel: AttentionUnet = AttentionUnet()\nparam_group: list[dict] = [\n    {\"params\": model.backbone.layer3.parameters(), \"lr\": 3e-5},\n    {\"params\": model.bottleneck.parameters(), \"lr\": 1e-4},\n    {\"params\": model.decoder_l3.parameters(), \"lr\": 1e-4},\n    {\"params\": model.decoder_l2.parameters(), \"lr\": 1e-4},\n    {\"params\": model.decoder_l1.parameters(), \"lr\": 1e-4},\n    {\"params\": model.decoder_l0.parameters(), \"lr\": 1e-4},\n    {\"params\": model.final_conv.parameters(), \"lr\": 2e-4},\n]\n\nnum_pos_samples: int = (dataset.mask_infos[\"ShipCounts\"] > 0).sum()\nnum_samples: int = len(dataset.mask_infos)\n\noptimizer: AdamW     = AdamW(param_group, lr = 2e-4, weight_decay = 5e-3)\nscheduler: CosineAnnealingLR = CosineAnnealingLR(optimizer, T_max = num_epochs * 1.5 , eta_min = 2e-5)\n\npos_weight = min(num_samples/num_pos_samples, 10.0)\ncriterion: SegmentLoss = SegmentLoss(positive_weight = pos_weight)\nevaluator: SegmentLoss = SegmentLoss(sample_weighter = None, positive_weight = pos_weight)\ngrad_scaler: GradScaler = GradScaler('cuda')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:24.805543Z","iopub.execute_input":"2026-03-31T16:13:24.805838Z","iopub.status.idle":"2026-03-31T16:13:25.381395Z","shell.execute_reply.started":"2026-03-31T16:13:24.805814Z","shell.execute_reply":"2026-03-31T16:13:25.380787Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\nbatch_size = 16\nTrainingLoader: DataLoader = DataLoader(TrainingDataset, \n                                        batch_size = batch_size,  \n                                        num_workers = 4, \n                                        pin_memory = True,\n                                        shuffle    = True,\n                                        persistent_workers = True)\n\nValidationLoader: DataLoader = DataLoader(ValidationDataset,\n                                          batch_size = batch_size,\n                                          num_workers = 4,\n                                          pin_memory = True,\n                                          shuffle = True,\n                                          persistent_workers = True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:25.382253Z","iopub.execute_input":"2026-03-31T16:13:25.382531Z","iopub.status.idle":"2026-03-31T16:13:25.390274Z","shell.execute_reply.started":"2026-03-31T16:13:25.382507Z","shell.execute_reply":"2026-03-31T16:13:25.389528Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.to(device)\ncriterion.to(device)\nevaluator.to(device)\n\nmodel.decoder_l0.apply(init_layer)\nmodel.decoder_l1.apply(init_layer)\nmodel.decoder_l2.apply(init_layer)\nmodel.decoder_l3.apply(init_layer)\nmodel.final_conv.apply(init_layer)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:25.391108Z","iopub.execute_input":"2026-03-31T16:13:25.391435Z","iopub.status.idle":"2026-03-31T16:13:25.454129Z","shell.execute_reply.started":"2026-03-31T16:13:25.391398Z","shell.execute_reply":"2026-03-31T16:13:25.453471Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm as ProgressBarView\n\nepoch_history: dict = {\n    \"BatchIdx\": [],\n    \"BWTrainIoU\": [],\n    \"BWTrainFocal\": [],\n    \"BWWeightedLossT\": [],\n    \"BWValIoU\": [],\n    \"BWValFocal\": [],\n    \"BWWeightedLossV\": []\n}\nbest_loss: float = float(\"inf\")\nckpoint_path = \"best_model.pth\"\n\nfor epoch_no in range(num_epochs):\n    print(\"\\n-------------------------------------------------\")\n    print(f\"Epoch no {epoch_no} started! \")\n    model.train(True)\n\n    pbar = ProgressBarView(TrainingLoader)\n    for batch_idx, (imgs, lbls, ns_s, nsp_s) in enumerate(pbar):\n        if (epoch_no + 1) % 4 == 0:\n            optimizer.zero_grad()\n            \n        imgs = imgs.to(device)\n        lbls = lbls.to(device)\n        ns_s = ns_s.to(device)\n        nsp_s = nsp_s.to(device)\n        \n        with torch.amp.autocast( 'cuda' ,dtype = torch.float16):\n            preds = model(imgs)\n            loss, iou, focal = criterion(preds, lbls, ns_s, nsp_s)\n            \n        grad_scaler.scale(loss).backward()\n        if (epoch_no + 1) % 4 == 0:\n            grad_scaler.step(optimizer)\n            grad_scaler.update()\n            \n        epoch_history[\"BatchIdx\"].append(batch_idx)\n        epoch_history[\"BWWeightedLossT\"].append(loss.item())\n        epoch_history[\"BWTrainIoU\"].append(iou)\n        epoch_history[\"BWTrainFocal\"].append(focal)\n        pbar.set_postfix({\n            \"iou\": f\"{iou :.3f}\",\n            \"focal\": f\"{focal :.3f}\",\n            \"w_loss\": f\"{loss.item() :.3f}\"\n        })\n    model.eval()\n    \n    avg_weighted_val_loss: float = 0.0\n\n    pbar = ProgressBarView(ValidationLoader)\n    with torch.no_grad():\n        for batch_idx, (imgs, lbls, ns_s, nsp_s) in enumerate(pbar):\n        \n            imgs = imgs.to(device)\n            lbls = lbls.to(device)\n            ns_s = ns_s.to(device)\n            nsp_s = nsp_s.to(device)\n            \n            with torch.amp.autocast('cuda',dtype = torch.float16):\n                preds = model(imgs)\n                loss, iou, focal = evaluator(preds, lbls, ns_s, nsp_s)\n                \n            epoch_history[\"BatchIdx\"].append(batch_idx)\n            epoch_history[\"BWWeightedLossV\"].append(loss.item())\n            epoch_history[\"BWValIoU\"].append(iou)\n            epoch_history[\"BWValFocal\"].append(focal)\n            avg_weighted_val_loss += loss.item()/len(ValidationLoader)\n            pbar.set_postfix({\n                \"iou\": f\"{iou :.3f}\",\n                \"focal\": f\"{focal :.3f}\",\n                \"w_loss\": f\"{loss.item() :.3f}\"\n            })\n    \n    if best_loss > avg_weighted_val_loss:\n        best_loss = avg_weighted_val_loss\n        ckpnt: dict = {\n            \"model_state_dict\": model.state_dict(),\n            \"optimizer_state_dict\": optimizer.state_dict(),\n            \"scheduler_state_dict\": scheduler.state_dict(),\n            \"epoch_no\": epoch_no,\n            \"loss\": best_loss\n        }\n        torch.save(ckpnt, ckpoint_path)\n        print(\"New best model found!\")\n    print(\"\\n-------------------------------------------------\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T16:13:25.455198Z","iopub.execute_input":"2026-03-31T16:13:25.455428Z","iopub.status.idle":"2026-03-31T17:57:07.933238Z","shell.execute_reply.started":"2026-03-31T16:13:25.455407Z","shell.execute_reply":"2026-03-31T17:57:07.931786Z"}},"outputs":[],"execution_count":null}]}