{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":23823,"databundleVersionId":1920183},{"sourceType":"datasetVersion","sourceId":2215761,"datasetId":1330639,"databundleVersionId":2257404},{"sourceType":"datasetVersion","sourceId":2153167,"datasetId":1229677,"databundleVersionId":2194269},{"sourceType":"datasetVersion","sourceId":2185494,"datasetId":1136885,"databundleVersionId":2226909},{"sourceType":"datasetVersion","sourceId":2182061,"datasetId":1136931,"databundleVersionId":2223452},{"sourceType":"datasetVersion","sourceId":2167436,"datasetId":1149113,"databundleVersionId":2208662},{"sourceType":"datasetVersion","sourceId":1983975,"datasetId":1182793,"databundleVersionId":2023128},{"sourceType":"datasetVersion","sourceId":2194412,"datasetId":1317588,"databundleVersionId":2235923},{"sourceType":"datasetVersion","sourceId":2188154,"datasetId":1313589,"databundleVersionId":2229604},{"sourceType":"datasetVersion","sourceId":2211937,"datasetId":1328360,"databundleVersionId":2253568},{"sourceType":"datasetVersion","sourceId":2211395,"datasetId":1325931,"databundleVersionId":2253025},{"sourceType":"datasetVersion","sourceId":2211693,"datasetId":1328208,"databundleVersionId":2253323},{"sourceType":"datasetVersion","sourceId":2165946,"datasetId":1300160,"databundleVersionId":2207150},{"sourceType":"datasetVersion","sourceId":2211791,"datasetId":1328270,"databundleVersionId":2253422},{"sourceType":"datasetVersion","sourceId":12285780,"datasetId":7742754,"databundleVersionId":12836666},{"sourceType":"datasetVersion","sourceId":12285881,"datasetId":7742827,"databundleVersionId":12836780},{"sourceType":"datasetVersion","sourceId":12285818,"datasetId":7742780,"databundleVersionId":12836710},{"sourceType":"datasetVersion","sourceId":3075714,"datasetId":849808,"databundleVersionId":3124377},{"sourceType":"datasetVersion","sourceId":2214724,"datasetId":1330046,"databundleVersionId":2256362},{"sourceType":"datasetVersion","sourceId":12285770,"datasetId":7742744,"databundleVersionId":12836656},{"sourceType":"modelInstanceVersion","sourceId":714579,"databundleVersionId":15267713,"modelInstanceId":543057,"modelId":556222},{"sourceType":"modelInstanceVersion","sourceId":720879,"databundleVersionId":15334696,"modelInstanceId":548329,"modelId":561104},{"sourceType":"modelInstanceVersion","sourceId":713699,"databundleVersionId":15257345,"modelInstanceId":542326,"modelId":555530},{"sourceType":"modelInstanceVersion","sourceId":721887,"databundleVersionId":15345506,"modelInstanceId":549179,"modelId":561885},{"sourceType":"modelInstanceVersion","sourceId":722362,"databundleVersionId":15350871,"modelInstanceId":549573,"modelId":562248},{"sourceType":"modelInstanceVersion","sourceId":720090,"databundleVersionId":15325029,"modelInstanceId":547682,"modelId":560482},{"sourceType":"modelInstanceVersion","sourceId":781441,"databundleVersionId":16012243,"modelInstanceId":596214,"modelId":608472},{"sourceType":"modelInstanceVersion","sourceId":787692,"databundleVersionId":16105868,"modelInstanceId":601089,"modelId":613297},{"sourceType":"modelInstanceVersion","sourceId":366994,"databundleVersionId":12097638,"modelInstanceId":304244,"modelId":324725},{"sourceType":"modelInstanceVersion","sourceId":422817,"databundleVersionId":12563584,"modelInstanceId":344552,"modelId":365846},{"sourceType":"modelInstanceVersion","sourceId":422821,"databundleVersionId":12563646,"modelInstanceId":344556,"modelId":365850},{"sourceType":"modelInstanceVersion","sourceId":422825,"databundleVersionId":12563685,"modelInstanceId":344559,"modelId":365853},{"sourceType":"modelInstanceVersion","sourceId":423743,"databundleVersionId":12573910,"modelInstanceId":345332,"modelId":366624}],"dockerImageVersionId":30097,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints\n!cp -r ../input/landmark-additional-packages/rwightman_gen-efficientnet-pytorch_master/rwightman_gen-efficientnet-pytorch_master /root/.cache/torch/hub\n!cp ../input/landmark-additional-packages/tf_efficientnet_b3_aa-84b4657e.pth /root/.cache/torch/hub/checkpoints/\n!cp ../input/landmark-additional-packages/tf_efficientnet_b5_ra-9a3e5369.pth /root/.cache/torch/hub/checkpoints/\n!cp ../input/landmark-additional-packages/se_resnext50_32x4d-a260b3a4.pth /root/.cache/torch/hub/checkpoints/\n!cp ../input/landmark-additional-packages/resnet50d_ra2-464e36ba.pth /root/.cache/torch/hub/checkpoints/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:43:21.522299Z","iopub.execute_input":"2026-03-15T14:43:21.522573Z","iopub.status.idle":"2026-03-15T14:43:34.148625Z","shell.execute_reply.started":"2026-03-15T14:43:21.522513Z","shell.execute_reply":"2026-03-15T14:43:34.147566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q ../input/landmark-additional-packages/timm-0.3.4-py3-none-any.whl\n!pip install -q ../input/landmark-additional-packages/geffnet-1.0.0-py3-none-any.whl\n!pip install -q ../input/landmark-additional-packages/EfficientNet-PyTorch/EfficientNet-PyTorch-master\n!pip install -q ../input/landmark-additional-packages/pycocotools-2.0.2/dist/pycocotools-2.0.2.tar\n!pip install -q ../input/landmark-additional-packages/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:43:34.150079Z","iopub.execute_input":"2026-03-15T14:43:34.150432Z","iopub.status.idle":"2026-03-15T14:46:51.101345Z","shell.execute_reply.started":"2026-03-15T14:43:34.150393Z","shell.execute_reply":"2026-03-15T14:46:51.100166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install \"/kaggle/input/hpapytorchzoozip/pytorch_zoo-master\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:46:51.103021Z","iopub.execute_input":"2026-03-15T14:46:51.103396Z","iopub.status.idle":"2026-03-15T14:47:29.592788Z","shell.execute_reply.started":"2026-03-15T14:46:51.103354Z","shell.execute_reply":"2026-03-15T14:47:29.591851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nimport sys\nsys.path.append('../input/hpa-singlecell-e050f56/hpa_singlecell-double_level_valid_all/')\n\nfrom torch import nn\nimport torch\nimport torch.nn.functional as F\nimport torchvision\nimport timm\nfrom torch.nn.parameter import Parameter\nimport albumentations as A\n\nfrom utils import parse_args, prepare_for_result\nfrom torch.utils.data import DataLoader, Dataset\nfrom losses import get_loss, get_class_balanced_weighted\nfrom dataloaders import get_dataloader\nfrom utils import load_matched_state\nfrom configs import Config\nfrom models import get_model\nfrom dataloaders.transform_loader import get_tfms\n\nimport base64\nimport zlib\nfrom pycocotools import _mask as coco_mask\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport cv2\nimport tqdm\nimport seaborn as sns","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:48:38.324971Z","iopub.execute_input":"2026-03-15T14:48:38.325298Z","iopub.status.idle":"2026-03-15T14:48:45.946695Z","shell.execute_reply.started":"2026-03-15T14:48:38.325268Z","shell.execute_reply":"2026-03-15T14:48:45.945847Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchvision\nprint(torchvision.__version__)\n\nfrom torchvision import transforms","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:48:45.948131Z","iopub.execute_input":"2026-03-15T14:48:45.948396Z","iopub.status.idle":"2026-03-15T14:48:45.9525Z","shell.execute_reply.started":"2026-03-15T14:48:45.948366Z","shell.execute_reply":"2026-03-15T14:48:45.951715Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import scipy.ndimage as ndi\nfrom scipy.sparse import csr_matrix\nfrom skimage import filters, measure, segmentation, transform, util\nfrom skimage.morphology import closing, disk, remove_small_holes, remove_small_objects","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:48:45.953737Z","iopub.execute_input":"2026-03-15T14:48:45.954063Z","iopub.status.idle":"2026-03-15T14:48:45.963515Z","shell.execute_reply.started":"2026-03-15T14:48:45.95403Z","shell.execute_reply":"2026-03-15T14:48:45.962717Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import psutil\nimport os\n\ndef print_memory_usage():\n    process = psutil.Process(os.getpid())\n    mem_info = process.memory_info()\n    print(f\"RAM Usada: {mem_info.rss / 1024 / 1024:.2f} MB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:34:00.516265Z","iopub.execute_input":"2026-03-15T15:34:00.516645Z","iopub.status.idle":"2026-03-15T15:34:00.520801Z","shell.execute_reply.started":"2026-03-15T15:34:00.516613Z","shell.execute_reply":"2026-03-15T15:34:00.519947Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Transforms","metadata":{}},{"cell_type":"code","source":"tensor_tfms = torchvision.transforms.Compose([\n            torchvision.transforms.ToTensor(),\n            torchvision.transforms.Normalize(mean=[0.485, 0.456, 0.406, 0.485], \n                                             std=[0.229, 0.224, 0.225, 0.229]),\n        ])\n\nimage_tfms = torchvision.transforms.Compose([\n            torchvision.transforms.ToTensor(),\n            torchvision.transforms.Normalize(mean=[0.485, 0.456, 0.406, 0.485], \n                                             std=[0.229, 0.224, 0.225, 0.229]),\n        ])\n\nres_tfms = A.Compose([\n    A.Resize(always_apply=False, p=1, height=512, width=512, interpolation=1)\n])\n\nimage_res_tfms = A.Compose([\n    A.Resize(always_apply=False, p=1, height=512, width=512, interpolation=1)\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:48:45.964618Z","iopub.execute_input":"2026-03-15T14:48:45.964966Z","iopub.status.idle":"2026-03-15T14:48:45.976245Z","shell.execute_reply.started":"2026-03-15T14:48:45.964934Z","shell.execute_reply":"2026-03-15T14:48:45.975655Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Seed","metadata":{}},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)  # if using GPU\n\n    # For convolutional determinism\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything(42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:48:45.977239Z","iopub.execute_input":"2026-03-15T14:48:45.977566Z","iopub.status.idle":"2026-03-15T14:48:45.992619Z","shell.execute_reply.started":"2026-03-15T14:48:45.977534Z","shell.execute_reply":"2026-03-15T14:48:45.992025Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Functions","metadata":{}},{"cell_type":"code","source":"def binary_mask_to_ascii(mask, mask_val=1):\n    \"\"\"Converts a binary mask into OID challenge encoding ascii text.\"\"\"\n    mask = np.where(mask==mask_val, 1, 0).astype(np.bool)\n    \n    # check input mask --\n    if mask.dtype != np.bool:\n        raise ValueError(f\"encode_binary_mask expects a binary mask, received dtype == {mask.dtype}\")\n\n    mask = np.squeeze(mask)\n    if len(mask.shape) != 2:\n        raise ValueError(f\"encode_binary_mask expects a 2d mask, received shape == {mask.shape}\")\n\n    # convert input mask to expected COCO API input --\n    mask_to_encode = mask.reshape(mask.shape[0], mask.shape[1], 1)\n    mask_to_encode = mask_to_encode.astype(np.uint8)\n    mask_to_encode = np.asfortranarray(mask_to_encode)\n\n    # RLE encode mask --\n    encoded_mask = coco_mask.encode(mask_to_encode)[0][\"counts\"]\n\n    # compress and base64 encoding --\n    binary_str = zlib.compress(encoded_mask, zlib.Z_BEST_COMPRESSION)\n    base64_str = base64.b64encode(binary_str)\n    return base64_str.decode()\n\ndef squarify(M,val):\n    (a,b,c)=M.shape\n    if a>b:\n        padding=((0,0),((a-b)//2,a-b-(a-b)//2),(0, 0))\n    else:\n        padding=(((b-a)//2,b-a-(b-a)//2),(0,0),(0, 0))\n    return np.pad(M,padding,mode='constant',constant_values=val)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:48:45.993651Z","iopub.execute_input":"2026-03-15T14:48:45.99397Z","iopub.status.idle":"2026-03-15T14:48:46.003553Z","shell.execute_reply.started":"2026-03-15T14:48:45.993934Z","shell.execute_reply":"2026-03-15T14:48:46.002768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CLASSES = np.asarray([\n'0. Nucleoplasm',\n'1. Nuclear-membrane',\n'2. Nucleoli',\n'3. Nucleoli-fibrillar-center',\n'4. Nuclear-speckles',\n'5. Nuclear-bodies',\n'6. Endoplasmic-reticulum',\n'7. Golgi apparatus',\n'8. Intermediate-filaments',\n'9. Actin-filaments',\n'10. Microtubules',\n'11. Mitotic-spindle',\n'12. Centrosome',\n'13. Plasma-membrane',\n'14. Mitochondria',\n'15. Aggresome',\n'16. Cytosol',\n'17. Vesicles',\n'18. Negative'\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:48:46.004563Z","iopub.execute_input":"2026-03-15T14:48:46.004765Z","iopub.status.idle":"2026-03-15T14:48:46.012928Z","shell.execute_reply.started":"2026-03-15T14:48:46.004746Z","shell.execute_reply":"2026-03-15T14:48:46.01215Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ResNest269 Model","metadata":{}},{"cell_type":"code","source":"#@title Model / Architecture\n\nimport math\nimport random\n\nimport cv2\nimport numpy as np\n\nimport torch\nimport torch.nn.functional as F\nfrom torch import nn\nfrom torch.nn import BatchNorm2d, Conv2d, Linear, Module, ReLU\nfrom torch.nn.modules.utils import _pair\nimport torch.utils.model_zoo as model_zoo\nfrom torch.optim.lr_scheduler import LambdaLR\n\n\ndef set_seed(seed):\n  random.seed(seed)\n  np.random.seed(seed)\n\n  torch.manual_seed(seed)\n  if torch.cuda.is_available():\n    torch.cuda.manual_seed_all(seed)\n\n\ndef rotation(x, k):\n  return torch.rot90(x, k, (1, 2))\n\n\ndef interleave(x, size):\n  s = list(x.shape)\n  return x.reshape([-1, size] + s[1:]).transpose(0, 1).reshape([-1] + s[1:])\n\n\ndef de_interleave(x, size):\n  s = list(x.shape)\n  return x.reshape([size, -1] + s[1:]).transpose(0, 1).reshape([-1] + s[1:])\n\n\ndef resize_tensor(tensors, size, mode='bilinear', align_corners=None):\n  return F.interpolate(tensors, size, mode=mode, align_corners=align_corners)\n\n\ndef gap2d(x, keepdims=False):\n  x = torch.mean(x.view(x.size(0), x.size(1), -1), -1)\n  if keepdims:\n    x = x.view(x.size(0), x.size(1), 1, 1)\n  return x\n\n\n# Losses\n\n\ndef L1_Loss(A_tensors, B_tensors):\n  return torch.abs(A_tensors - B_tensors)\n\n\ndef L2_Loss(A_tensors, B_tensors):\n  return torch.pow(A_tensors - B_tensors, 2)\n\n\n# ratio = 0.2, top=20%\ndef Online_Hard_Example_Mining(values, ratio=0.2):\n  b, c, h, w = values.size()\n  return torch.topk(values.reshape(b, -1), k=int(c * h * w * ratio), dim=-1)[0]\n\n\ndef shannon_entropy_loss(logits, activation=torch.sigmoid, epsilon=1e-5):\n  v = activation(logits)\n  return -torch.sum(v * torch.log(v + epsilon), dim=1).mean()\n\n\ndef make_cam(x, eps=1e-5, shift_min=False, global_norm=False, inplace=True):\n  x = F.relu(x)\n\n  if global_norm:\n    x_min = x.min() if shift_min else 0\n    x_max = x.max() - x_min\n  else:\n    b, c, h, w = x.size()\n    flat_x = x.view(b, c, -1)\n    x_min = flat_x.min(axis=-1)[0].view((b, c, 1, 1)) if shift_min else 0\n    x_max = flat_x.max(axis=-1)[0].view((b, c, 1, 1)) - x_min\n\n  if shift_min:\n    if inplace:\n      x -= x_min\n      x /= x_max + eps\n    else:\n      x = (x - x_min) / (x_max + eps)\n  else:\n    if inplace:\n      x /= x_max + eps\n    else:\n      x = x / (x_max + eps)\n\n  return x\n\n\ndef one_hot_embedding(label, classes):\n  \"\"\"Embedding labels to one-hot form.\n\n    Args:\n      labels: (int) class labels.\n      num_classes: (int) number of classes.\n\n    Returns:\n      (tensor) encoded labels, sized [N, #classes].\n    \"\"\"\n\n  vector = np.zeros((classes), dtype=np.float32)\n  if len(label) > 0:\n    vector[label] = 1.\n  return vector\n\n\ndef calculate_parameters(model):\n  return sum(param.numel() for param in model.parameters()) / 1000000.0\n\n\ndef get_learning_rate_from_optimizer(optimizer):\n  return optimizer.param_groups[0]['lr']\n\n\ndef set_trainable_layers(model, klass=None, trainable=False):\n  for m in model.modules():\n    if klass is None or isinstance(m, klass):\n      for prop in (\"weight\", \"bias\"):\n        w = getattr(m, prop, None)\n        if w is not None:\n          w.requires_grad = trainable\n\n\ndef to_numpy(tensor):\n  return tensor.cpu().detach().numpy()\n\n\ndef load_model(model, model_path, parallel=False, map_location=None, strict=True):\n  print(f'loading weights from `{model_path}`.')\n\n  state = torch.load(model_path, map_location=map_location)\n  if parallel:\n    model.module.load_state_dict(state, strict=strict)\n  else:\n    model.load_state_dict(state, strict=strict)\n\n\ndef save_model(model, model_path, parallel=False):\n  print(f'saving weights to `{model_path}`.')\n\n  if parallel:\n    torch.save(model.module.state_dict(), model_path)\n  else:\n    torch.save(model.state_dict(), model_path)\n\n\ndef transfer_model(pretrained_model, model):\n  pretrained_dict = pretrained_model.state_dict()\n  model_dict = model.state_dict()\n\n  pretrained_dict = {k: v for k, v in pretrained_dict.items() if k in model_dict}\n\n  model_dict.update(pretrained_dict)\n  model.load_state_dict(model_dict)\n\n\ndef get_learning_rate(optimizer):\n  lr = []\n  for param_group in optimizer.param_groups:\n    lr += [param_group['lr']]\n  return lr\n\n\ndef get_cosine_schedule_with_warmup(optimizer, warmup_iteration, max_iteration, cycles=7. / 16.):\n\n  def _lr_lambda(current_iteration):\n    if current_iteration < warmup_iteration:\n      return float(current_iteration) / float(max(1, warmup_iteration))\n\n    no_progress = float(current_iteration - warmup_iteration) / float(max(1, max_iteration - warmup_iteration))\n    return max(0., math.cos(math.pi * cycles * no_progress))\n\n  return LambdaLR(optimizer, _lr_lambda, -1)\n\n\ndef label_smoothing(labels, alpha):\n  if alpha:\n    return (1 - alpha) * labels + alpha * 0.5\n\n  return labels\n\n\nclass SplAtConv2d(Module):\n  \"\"\"Split-Attention Conv2d\n    \"\"\"\n\n  def __init__(\n    self,\n    in_channels,\n    channels,\n    kernel_size,\n    stride=(1, 1),\n    padding=(0, 0),\n    dilation=(1, 1),\n    groups=1,\n    bias=True,\n    radix=2,\n    reduction_factor=4,\n    rectify=False,\n    rectify_avg=False,\n    norm_layer=None,\n    dropblock_prob=0.0,\n    **kwargs\n  ):\n    super(SplAtConv2d, self).__init__()\n    padding = _pair(padding)\n    self.rectify = rectify and (padding[0] > 0 or padding[1] > 0)\n    self.rectify_avg = rectify_avg\n    inter_channels = max(in_channels * radix // reduction_factor, 32)\n    self.radix = radix\n    self.cardinality = groups\n    self.channels = channels\n    self.dropblock_prob = dropblock_prob\n    if self.rectify:\n      from rfconv import RFConv2d\n      self.conv = RFConv2d(\n        in_channels,\n        channels * radix,\n        kernel_size,\n        stride,\n        padding,\n        dilation,\n        groups=groups * radix,\n        bias=bias,\n        average_mode=rectify_avg,\n        **kwargs\n      )\n    else:\n      self.conv = Conv2d(\n        in_channels,\n        channels * radix,\n        kernel_size,\n        stride,\n        padding,\n        dilation,\n        groups=groups * radix,\n        bias=bias,\n        **kwargs\n      )\n    self.use_bn = norm_layer is not None\n    if self.use_bn:\n      self.bn0 = norm_layer(channels * radix)\n    self.relu = ReLU(inplace=True)\n    self.fc1 = Conv2d(channels, inter_channels, 1, groups=self.cardinality)\n    if self.use_bn:\n      self.bn1 = norm_layer(inter_channels)\n    self.fc2 = Conv2d(inter_channels, channels * radix, 1, groups=self.cardinality)\n    if dropblock_prob > 0.0:\n      self.dropblock = DropBlock2D(dropblock_prob, 3)\n    self.rsoftmax = rSoftMax(radix, groups)\n\n  def forward(self, x):\n    x = self.conv(x)\n    if self.use_bn:\n      x = self.bn0(x)\n    if self.dropblock_prob > 0.0:\n      x = self.dropblock(x)\n    x = self.relu(x)\n\n    batch, rchannel = x.shape[:2]\n    if self.radix > 1:\n      if torch.__version__ < '1.5':\n        splited = torch.split(x, int(rchannel // self.radix), dim=1)\n      else:\n        splited = torch.split(x, rchannel // self.radix, dim=1)\n      gap = sum(splited)\n    else:\n      gap = x\n    gap = F.adaptive_avg_pool2d(gap, 1)\n    gap = self.fc1(gap)\n\n    if self.use_bn:\n      gap = self.bn1(gap)\n    gap = self.relu(gap)\n\n    atten = self.fc2(gap)\n    atten = self.rsoftmax(atten).view(batch, -1, 1, 1)\n\n    if self.radix > 1:\n      if torch.__version__ < '1.5':\n        attens = torch.split(atten, int(rchannel // self.radix), dim=1)\n      else:\n        attens = torch.split(atten, rchannel // self.radix, dim=1)\n      out = sum([att * split for (att, split) in zip(attens, splited)])\n    else:\n      out = atten * x\n    return out.contiguous()\n\n\nclass rSoftMax(nn.Module):\n\n  def __init__(self, radix, cardinality):\n    super().__init__()\n    self.radix = radix\n    self.cardinality = cardinality\n\n  def forward(self, x):\n    batch = x.size(0)\n    if self.radix > 1:\n      x = x.view(batch, self.cardinality, self.radix, -1).transpose(1, 2)\n      x = F.softmax(x, dim=1)\n      x = x.reshape(batch, -1)\n    else:\n      x = torch.sigmoid(x)\n    return x\n\n\nclass DropBlock2D(object):\n\n  def __init__(self, *args, **kwargs):\n    raise NotImplementedError\n\n\nclass GlobalAvgPool2d(nn.Module):\n\n  def __init__(self):\n    \"\"\"Global average pooling over the input's spatial dimensions\"\"\"\n    super(GlobalAvgPool2d, self).__init__()\n\n  def forward(self, inputs):\n    return nn.functional.adaptive_avg_pool2d(inputs, 1).view(inputs.size(0), -1)\n\n\nclass Bottleneck(nn.Module):\n  \"\"\"ResNet Bottleneck\n    \"\"\"\n  # pylint: disable=unused-argument\n  expansion = 4\n\n  def __init__(\n    self,\n    inplanes,\n    planes,\n    stride=1,\n    downsample=None,\n    radix=1,\n    cardinality=1,\n    bottleneck_width=64,\n    avd=False,\n    avd_first=False,\n    dilation=1,\n    is_first=False,\n    rectified_conv=False,\n    rectify_avg=False,\n    norm_layer=None,\n    dropblock_prob=0.0,\n    last_gamma=False\n  ):\n    super(Bottleneck, self).__init__()\n    group_width = int(planes * (bottleneck_width / 64.)) * cardinality\n    self.conv1 = nn.Conv2d(inplanes, group_width, kernel_size=1, bias=False)\n    self.bn1 = norm_layer(group_width)\n    self.dropblock_prob = dropblock_prob\n    self.radix = radix\n    self.avd = avd and (stride > 1 or is_first)\n    self.avd_first = avd_first\n\n    if self.avd:\n      self.avd_layer = nn.AvgPool2d(3, stride, padding=1)\n      stride = 1\n\n    if dropblock_prob > 0.0:\n      self.dropblock1 = DropBlock2D(dropblock_prob, 3)\n      if radix == 1:\n        self.dropblock2 = DropBlock2D(dropblock_prob, 3)\n      self.dropblock3 = DropBlock2D(dropblock_prob, 3)\n\n    if radix >= 1:\n      self.conv2 = SplAtConv2d(\n        group_width,\n        group_width,\n        kernel_size=3,\n        stride=stride,\n        padding=dilation,\n        dilation=dilation,\n        groups=cardinality,\n        bias=False,\n        radix=radix,\n        rectify=rectified_conv,\n        rectify_avg=rectify_avg,\n        norm_layer=norm_layer,\n        dropblock_prob=dropblock_prob\n      )\n    elif rectified_conv:\n      from rfconv import RFConv2d\n      self.conv2 = RFConv2d(\n        group_width,\n        group_width,\n        kernel_size=3,\n        stride=stride,\n        padding=dilation,\n        dilation=dilation,\n        groups=cardinality,\n        bias=False,\n        average_mode=rectify_avg\n      )\n      self.bn2 = norm_layer(group_width)\n    else:\n      self.conv2 = nn.Conv2d(\n        group_width,\n        group_width,\n        kernel_size=3,\n        stride=stride,\n        padding=dilation,\n        dilation=dilation,\n        groups=cardinality,\n        bias=False\n      )\n      self.bn2 = norm_layer(group_width)\n\n    self.conv3 = nn.Conv2d(group_width, planes * 4, kernel_size=1, bias=False)\n    self.bn3 = norm_layer(planes * 4)\n\n    if last_gamma:\n      from torch.nn.init import zeros_\n      zeros_(self.bn3.weight)\n    self.relu = nn.ReLU(inplace=True)\n    self.downsample = downsample\n    self.dilation = dilation\n    self.stride = stride\n\n  def forward(self, x):\n    residual = x\n\n    out = self.conv1(x)\n    out = self.bn1(out)\n    if self.dropblock_prob > 0.0:\n      out = self.dropblock1(out)\n    out = self.relu(out)\n\n    if self.avd and self.avd_first:\n      out = self.avd_layer(out)\n\n    out = self.conv2(out)\n    if self.radix == 0:\n      out = self.bn2(out)\n      if self.dropblock_prob > 0.0:\n        out = self.dropblock2(out)\n      out = self.relu(out)\n\n    if self.avd and not self.avd_first:\n      out = self.avd_layer(out)\n\n    out = self.conv3(out)\n    out = self.bn3(out)\n    if self.dropblock_prob > 0.0:\n      out = self.dropblock3(out)\n\n    if self.downsample is not None:\n      residual = self.downsample(x)\n\n    out += residual\n    out = self.relu(out)\n\n    return out\n\n\nclass ResNet(nn.Module):\n  \"\"\"ResNet Variants\n\n    Parameters\n    ----------\n    block : Block\n        Class for the residual block. Options are BasicBlockV1, BottleneckV1.\n    layers : list of int\n        Numbers of layers in each block\n    classes : int, default 1000\n        Number of classification classes.\n    dilated : bool, default False\n        Applying dilation strategy to pretrained ResNet yielding a stride-8 model,\n        typically used in Semantic Segmentation.\n    norm_layer : object\n        Normalization layer used in backbone network (default: :class:`mxnet.gluon.nn.BatchNorm`;\n        for Synchronized Cross-GPU BachNormalization).\n\n    Reference:\n\n        - He, Kaiming, et al. \"Deep residual learning for image recognition.\" Proceedings of the IEEE conference on computer vision and pattern recognition. 2016.\n\n        - Yu, Fisher, and Vladlen Koltun. \"Multi-scale context aggregation by dilated convolutions.\"\n    \"\"\"\n\n  # pylint: disable=unused-variable\n  def __init__(\n    self,\n    block,\n    layers,\n    radix=1,\n    groups=1,\n    bottleneck_width=64,\n    num_classes=1000,\n    dilated=False,\n    dilation=1,\n    deep_stem=False,\n    stem_width=64,\n    avg_down=False,\n    rectified_conv=False,\n    rectify_avg=False,\n    avd=False,\n    avd_first=False,\n    final_drop=0.0,\n    dropblock_prob=0,\n    last_gamma=False,\n    norm_layer=nn.BatchNorm2d\n  ):\n    self.cardinality = groups\n    self.bottleneck_width = bottleneck_width\n    # ResNet-D params\n    self.stage_features = []\n    self.outplanes = self.inplanes = stem_width * 2 if deep_stem else 64\n    self.avg_down = avg_down\n    self.last_gamma = last_gamma\n    # ResNeSt params\n    self.radix = radix\n    self.avd = avd\n    self.avd_first = avd_first\n\n    super(ResNet, self).__init__()\n    self.rectified_conv = rectified_conv\n    self.rectify_avg = rectify_avg\n    if rectified_conv:\n      from rfconv import RFConv2d\n      conv_layer = RFConv2d\n    else:\n      conv_layer = nn.Conv2d\n    conv_kwargs = {'average_mode': rectify_avg} if rectified_conv else {}\n    if deep_stem:\n      self.conv1 = nn.Sequential(\n        conv_layer(3, stem_width, kernel_size=3, stride=2, padding=1, bias=False, **conv_kwargs),\n        norm_layer(stem_width),\n        nn.ReLU(inplace=True),\n        conv_layer(stem_width, stem_width, kernel_size=3, stride=1, padding=1, bias=False, **conv_kwargs),\n        norm_layer(stem_width),\n        nn.ReLU(inplace=True),\n        conv_layer(stem_width, stem_width * 2, kernel_size=3, stride=1, padding=1, bias=False, **conv_kwargs),\n      )\n    else:\n      self.conv1 = conv_layer(3, 64, kernel_size=7, stride=2, padding=3, bias=False, **conv_kwargs)\n    self.bn1 = norm_layer(self.inplanes)\n    self.relu = nn.ReLU(inplace=True)\n    self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n    self.layer1 = self._make_layer(block, 64, layers[0], norm_layer=norm_layer, is_first=False)\n    self.layer2 = self._make_layer(block, 128, layers[1], stride=2, norm_layer=norm_layer)\n    if dilated or dilation == 4:\n      self.layer3 = self._make_layer(\n        block, 256, layers[2], stride=1, dilation=2, norm_layer=norm_layer, dropblock_prob=dropblock_prob\n      )\n      self.layer4 = self._make_layer(\n        block, 512, layers[3], stride=1, dilation=4, norm_layer=norm_layer, dropblock_prob=dropblock_prob\n      )\n    elif dilation == 2:\n      self.layer3 = self._make_layer(\n        block, 256, layers[2], stride=2, dilation=1, norm_layer=norm_layer, dropblock_prob=dropblock_prob\n      )\n      self.layer4 = self._make_layer(\n        block, 512, layers[3], stride=1, dilation=2, norm_layer=norm_layer, dropblock_prob=dropblock_prob\n      )\n    else:\n      self.layer3 = self._make_layer(\n        block, 256, layers[2], stride=2, norm_layer=norm_layer, dropblock_prob=dropblock_prob\n      )\n      self.layer4 = self._make_layer(\n        block, 512, layers[3], stride=2, norm_layer=norm_layer, dropblock_prob=dropblock_prob\n      )\n\n    self.avgpool = GlobalAvgPool2d()\n    self.drop = nn.Dropout(final_drop) if final_drop > 0.0 else None\n    self.fc = nn.Linear(512 * block.expansion, num_classes)\n\n    for m in self.modules():\n      if isinstance(m, nn.Conv2d):\n        n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels\n        m.weight.data.normal_(0, math.sqrt(2. / n))\n      elif isinstance(m, norm_layer):\n        m.weight.data.fill_(1)\n        m.bias.data.zero_()\n\n  def _make_layer(\n    self, block, planes, blocks, stride=1, dilation=1, norm_layer=None, dropblock_prob=0.0, is_first=True\n  ):\n    downsample = None\n    if stride != 1 or self.inplanes != planes * block.expansion:\n      down_layers = []\n      if self.avg_down:\n        if dilation == 1:\n          down_layers.append(nn.AvgPool2d(kernel_size=stride, stride=stride, ceil_mode=True, count_include_pad=False))\n        else:\n          down_layers.append(nn.AvgPool2d(kernel_size=1, stride=1, ceil_mode=True, count_include_pad=False))\n        down_layers.append(nn.Conv2d(self.inplanes, planes * block.expansion, kernel_size=1, stride=1, bias=False))\n      else:\n        down_layers.append(nn.Conv2d(self.inplanes, planes * block.expansion, kernel_size=1, stride=stride, bias=False))\n      down_layers.append(norm_layer(planes * block.expansion))\n      downsample = nn.Sequential(*down_layers)\n\n    layers = []\n    if dilation == 1 or dilation == 2:\n      layers.append(\n        block(\n          self.inplanes,\n          planes,\n          stride,\n          downsample=downsample,\n          radix=self.radix,\n          cardinality=self.cardinality,\n          bottleneck_width=self.bottleneck_width,\n          avd=self.avd,\n          avd_first=self.avd_first,\n          dilation=1,\n          is_first=is_first,\n          rectified_conv=self.rectified_conv,\n          rectify_avg=self.rectify_avg,\n          norm_layer=norm_layer,\n          dropblock_prob=dropblock_prob,\n          last_gamma=self.last_gamma\n        )\n      )\n    elif dilation == 4:\n      layers.append(\n        block(\n          self.inplanes,\n          planes,\n          stride,\n          downsample=downsample,\n          radix=self.radix,\n          cardinality=self.cardinality,\n          bottleneck_width=self.bottleneck_width,\n          avd=self.avd,\n          avd_first=self.avd_first,\n          dilation=2,\n          is_first=is_first,\n          rectified_conv=self.rectified_conv,\n          rectify_avg=self.rectify_avg,\n          norm_layer=norm_layer,\n          dropblock_prob=dropblock_prob,\n          last_gamma=self.last_gamma\n        )\n      )\n    else:\n      raise RuntimeError(\"=> unknown dilation size: {}\".format(dilation))\n\n    self.outplanes = self.inplanes = planes * block.expansion\n    self.stage_features.append(self.outplanes)\n\n    for i in range(1, blocks):\n      layers.append(\n        block(\n          self.inplanes,\n          planes,\n          radix=self.radix,\n          cardinality=self.cardinality,\n          bottleneck_width=self.bottleneck_width,\n          avd=self.avd,\n          avd_first=self.avd_first,\n          dilation=dilation,\n          rectified_conv=self.rectified_conv,\n          rectify_avg=self.rectify_avg,\n          norm_layer=norm_layer,\n          dropblock_prob=dropblock_prob,\n          last_gamma=self.last_gamma\n        )\n      )\n\n    return nn.Sequential(*layers)\n\n  def forward(self, x):\n    x = self.conv1(x)\n    x = self.bn1(x)\n    x = self.relu(x)\n    x = self.maxpool(x)\n\n    x1 = self.layer1(x)\n    x2 = self.layer2(x1)\n    x3 = self.layer3(x2)\n    x4 = self.layer4(x3)\n\n    # print(x.size())\n\n    # x = self.avgpool(x)\n    #x = x.view(x.size(0), -1)\n    # x = torch.flatten(x, 1)\n    # if self.drop:\n    #   x = self.drop(x)\n    # x = self.fc(x)\n\n    return (x1, x2, x3, x4)\n\n_url_format = 'https://github.com/zhanghang1989/ResNeSt/releases/download/weights_step1/{}-{}.pth'\n\n_model_sha256 = {\n  name: checksum for checksum, name in [\n    ('528c19ca', 'resnest50'),\n    ('22405ba7', 'resnest101'),\n    ('75117900', 'resnest200'),\n    ('0cc87c48', 'resnest269'),\n  ]\n}\n\n\ndef short_hash(name):\n  if name not in _model_sha256:\n    raise ValueError('Pretrained model for {name} is not available.'.format(name=name))\n  return _model_sha256[name][:8]\n\n\nresnest_model_urls = {name: _url_format.format(name, short_hash(name)) for name in _model_sha256.keys()}\n\n\ndef resnest50(pretrained=False, root='~/.encoding/models', **kwargs):\n  model = ResNet(\n    Bottleneck, [3, 4, 6, 3],\n    radix=2,\n    groups=1,\n    bottleneck_width=64,\n    deep_stem=True,\n    stem_width=32,\n    avg_down=True,\n    avd=True,\n    avd_first=False,\n    **kwargs\n  )\n  if pretrained:\n    model.load_state_dict(\n      torch.hub.load_state_dict_from_url(resnest_model_urls['resnest50'], progress=True, check_hash=True)\n    )\n  return model\n\n\ndef resnest101(pretrained=False, root='~/.encoding/models', **kwargs):\n  model = ResNet(\n    Bottleneck, [3, 4, 23, 3],\n    radix=2,\n    groups=1,\n    bottleneck_width=64,\n    deep_stem=True,\n    stem_width=64,\n    avg_down=True,\n    avd=True,\n    avd_first=False,\n    **kwargs\n  )\n  if pretrained:\n    model.load_state_dict(\n      torch.hub.load_state_dict_from_url(resnest_model_urls['resnest101'], progress=True, check_hash=True)\n    )\n  return model\n\n\ndef resnest200(pretrained=False, root='~/.encoding/models', **kwargs):\n  model = ResNet(\n    Bottleneck, [3, 24, 36, 3],\n    radix=2,\n    groups=1,\n    bottleneck_width=64,\n    deep_stem=True,\n    stem_width=64,\n    avg_down=True,\n    avd=True,\n    avd_first=False,\n    **kwargs\n  )\n  if pretrained:\n    model.load_state_dict(\n      torch.hub.load_state_dict_from_url(resnest_model_urls['resnest200'], progress=True, check_hash=True)\n    )\n  return model\n\n\ndef resnest269(pretrained=False, root='~/.encoding/models', **kwargs):\n  model = ResNet(\n    Bottleneck, [3, 30, 48, 8],\n    radix=2,\n    groups=1,\n    bottleneck_width=64,\n    deep_stem=True,\n    stem_width=64,\n    avg_down=True,\n    avd=True,\n    avd_first=False,\n    **kwargs\n  )\n  if pretrained:\n    model.load_state_dict(\n      torch.hub.load_state_dict_from_url(resnest_model_urls['resnest269'], progress=True, check_hash=True)\n    )\n  return model\n\n\nclass FixedBatchNorm(nn.BatchNorm2d):\n\n  def forward(self, x):\n    return F.batch_norm(x, self.running_mean, self.running_var, self.weight, self.bias, training=False, eps=self.eps)\n\n\ndef group_norm(features):\n  return nn.GroupNorm(4, features)\n\n\n#######################################################################\n\ndef patch_conv_in_channels(model, layer_name, new_in_channels, copying_channel=0):\n  layer = getattr(model, layer_name)  # layer = model.conv1\n  # new_layer = layer.clone().detach()\n\n  if isinstance(layer, nn.Sequential):\n    cv_layers = list(layer.children())\n    cv = cv_layers[0]\n  else:\n    cv = layer\n\n  if not isinstance(cv, nn.Conv2d):\n    raise ValueError(f\"Cannot extract Conv2d from {cv}.\")\n\n  new_cv = nn.Conv2d(\n    in_channels=new_in_channels,\n    out_channels=cv.out_channels,\n    kernel_size=cv.kernel_size,\n    stride=cv.stride,\n    padding=cv.padding,\n    bias=cv.bias).requires_grad_()\n\n  with torch.no_grad():\n    new_cv.weight[:, :cv.in_channels, :, :] = cv.weight.data\n\n    for i in range(new_in_channels - cv.in_channels):\n        channel = cv.in_channels + i\n        new_cv.weight[:, channel:channel+1, :, :] = cv.weight[:, copying_channel:copying_channel+1, : :].data\n  new_cv.weight = nn.Parameter(new_cv.weight)\n\n  if isinstance(layer, nn.Sequential):\n    new_cv = nn.Sequential(new_cv, *cv_layers[1:])\n\n  setattr(model, layer_name, new_cv)  # model.conv1 = new_layer\n\n\ndef build_backbone(name, dilated, strides, norm_fn, weights='imagenet', channels=3, **kwargs):\n  dilation = 4 if dilated else 2\n\n  pretrained = weights == \"imagenet\"\n  model_fn = globals()[name]\n  model = model_fn(pretrained=pretrained, dilated=dilated, dilation=dilation, norm_layer=norm_fn)\n  if channels != 3:\n    patch_conv_in_channels(model, \"conv1\", channels)\n  if pretrained:\n    print(f'loading weights from {resnest_model_urls[name]}')\n\n  del model.avgpool\n  del model.fc\n\n  if weights and weights != 'imagenet':\n    print(f'loading weights from {weights}')\n    checkpoint = torch.load(weights, map_location=\"cpu\")\n    model.load_state_dict(checkpoint['state_dict'], strict=False)\n\n  stages = (\n    nn.Sequential(model.conv1, model.bn1, model.relu, model.maxpool),\n    model.layer1,\n    model.layer2,\n    model.layer3,\n    model.layer4,\n  )\n\n  return model, stages\n\n\nclass Backbone(nn.Module):\n\n  def __init__(\n    self,\n    model_name,\n    weights='imagenet',\n    channels=3,\n    mode='fix',\n    dilated=False,\n    strides=None,\n    trainable_stem=True,\n    trainable_stage4=True,\n    trainable_backbone=True,\n    backbone_kwargs={},\n  ):\n    super().__init__()\n\n    self.mode = mode\n    self.trainable_stem = trainable_stem\n    self.trainable_stage4 = trainable_stage4\n    self.trainable_backbone = trainable_backbone\n    self.not_training = []\n    self.from_scratch_layers = []\n\n    if mode == 'normal':\n      self.norm_fn = nn.BatchNorm2d\n    elif mode == 'fix':\n      self.norm_fn = FixedBatchNorm\n    else:\n      raise ValueError(f'Unknown mode {mode}. Must be `normal` or `fix`.')\n\n    backbone, stages = build_backbone(\n      model_name, dilated, strides, self.norm_fn, weights, channels, **backbone_kwargs,\n    )\n\n    self.backbone = backbone\n    self.stages = stages\n\n    if not self.trainable_backbone:\n      for s in stages:\n        set_trainable_layers(s, trainable=False)\n      self.not_training.extend(stages)\n    else:\n      if not self.trainable_stage4:\n        self.not_training.extend(stages[:-1])\n        for s in stages[:-1]:\n          set_trainable_layers(s, trainable=False)\n\n      elif not self.trainable_stem:\n        set_trainable_layers(stages[0], trainable=False)\n        self.not_training.append(stages[0])\n\n      if self.mode == \"fix\":\n        for s in stages:\n          set_trainable_layers(s, torch.nn.BatchNorm2d, trainable=False)\n          self.not_training.extend([m for m in s.modules() if isinstance(m, torch.nn.BatchNorm2d)])\n\n  def initialize(self, modules):\n    for m in modules:\n      if isinstance(m, nn.Conv2d):\n        torch.nn.init.kaiming_normal_(m.weight)\n      elif isinstance(m, nn.Linear):\n        nn.init.trunc_normal_(m.weight, std=.02)\n        if isinstance(m, nn.Linear) and m.bias is not None:\n            nn.init.constant_(m.bias, 0)\n      elif isinstance(m, (nn.BatchNorm2d, nn.SyncBatchNorm, nn.GroupNorm, nn.LayerNorm)):\n        nn.init.constant_(m.weight, 1.0)\n        nn.init.constant_(m.bias, 0)\n\n  def get_parameter_groups(self, exclude_partial_names=(), with_names=False):\n    names = ([], [], [], [])\n    groups = ([], [], [], [])\n\n    scratch_parameters = set()\n    all_parameters = set()\n\n    for layer in self.from_scratch_layers:\n      for name, param in layer.named_parameters():\n        if param in all_parameters:\n          continue\n        scratch_parameters.add(param)\n        all_parameters.add(param)\n\n        if not param.requires_grad:\n          continue\n        for p in exclude_partial_names:\n          if p in name:\n            continue\n\n        idx = 2 if \"weight\" in name else 3\n        names[idx].append(name)\n        groups[idx].append(param)\n\n    for name, param in self.named_parameters():\n      if param in all_parameters:\n        continue\n      all_parameters.add(param)\n\n      if not param.requires_grad or param in scratch_parameters:\n        continue\n      for p in exclude_partial_names:\n        if p in name:\n          continue\n\n      idx = 0 if \"weight\" in name else 1\n      names[idx].append(name)\n      groups[idx].append(param)\n\n    if with_names:\n      return groups, names\n\n    return groups\n\n  def train(self, mode=True):\n    super().train(mode)\n    for m in self.not_training:\n      m.eval()\n    return self\n\n\ndef gem(x, p=3, eps=1e-6):\n    return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1./p)\n\n\nclass GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM,self).__init__()\n        self.p = Parameter(torch.ones(1)*p)\n        self.eps = eps\n\n    def forward(self, x):\n        return gem(x, p=self.p, eps=self.eps)\n\n    def __repr__(self):\n        return self.__class__.__name__ + '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + ', ' + 'eps=' + str(self.eps) + ')'\n\n\nclass CellClassifier(Backbone):\n\n  def __init__(\n    self,\n    model_name,\n    num_classes=20,\n    backbone_weights=\"imagenet\",\n    channels=3,\n    mode='fix',\n    dilated=False,\n    strides=None,\n    trainable_stem=True,\n    trainable_stage4=True,\n    trainable_backbone=True,\n    **backbone_kwargs,\n  ):\n    super().__init__(\n      model_name,\n      channels=channels,\n      weights=backbone_weights,\n      mode=mode,\n      dilated=dilated,\n      strides=strides,\n      trainable_stem=trainable_stem,\n      trainable_stage4=trainable_stage4,\n      trainable_backbone=trainable_backbone,\n      backbone_kwargs=backbone_kwargs,\n    )\n\n    self.num_classes = num_classes\n\n    cin = self.backbone.outplanes\n    self.classifier = nn.Conv2d(cin, num_classes, 1, bias=False)\n\n    self.from_scratch_layers.extend([self.classifier])\n    self.initialize([self.classifier])\n\n    self.pool = GeM()\n    self.flatten = nn.Flatten()\n    self.dropout = nn.Dropout(p=0.5)\n\n    self.last_linear_cell = nn.Linear(\n      in_features=cin, \n      out_features=num_classes)\n    self.last_linear_image = nn.Linear(\n      in_features=cin, \n      out_features=num_classes)\n\n  def forward(self, x, cnt=16, with_cam=False):\n    if with_cam:\n      raise NotImplementedError(\n        \"CAM not currently supported in multi-view mode\")\n\n    outs = self.backbone(x)\n    features = outs[-1] if isinstance(outs, tuple) else outs\n\n    pooled = self.flatten(self.pool(features))\n    viewed_pooled = pooled.view(-1, cnt, pooled.shape[-1])\n    viewed_pooled = viewed_pooled.max(1)[0]\n\n    cell_logits = self.last_linear_cell(pooled)\n    image_logits = self.last_linear_image(viewed_pooled)\n\n    return cell_logits, image_logits\n\n\nclass ImageClassifier(Backbone):\n\n  def __init__(\n    self,\n    model_name,\n    num_classes=20,\n    backbone_weights=\"imagenet\",\n    channels=3,\n    mode='fix',\n    dilated=False,\n    strides=None,\n    trainable_stem=True,\n    trainable_stage4=True,\n    trainable_backbone=True,\n    **backbone_kwargs,\n  ):\n    super().__init__(\n      model_name,\n      channels=channels,\n      weights=backbone_weights,\n      mode=mode,\n      dilated=dilated,\n      strides=strides,\n      trainable_stem=trainable_stem,\n      trainable_stage4=trainable_stage4,\n      trainable_backbone=trainable_backbone,\n      backbone_kwargs=backbone_kwargs,\n    )\n\n    self.num_classes = num_classes\n\n    cin = self.backbone.outplanes\n    self.classifier = nn.Conv2d(cin, num_classes, 1, bias=False)\n\n    self.from_scratch_layers.extend([self.classifier])\n    self.initialize([self.classifier])\n\n  def forward(self, x, with_cam=False):\n    outs = self.backbone(x)\n    x = outs[-1] if isinstance(outs, tuple) else outs\n\n    if with_cam:\n      features = self.classifier(x)\n      logits = gap2d(features)\n      return logits, features\n    else:\n      x = gap2d(x, keepdims=True)\n      logits = self.classifier(x).view(-1, self.num_classes)\n      return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:48:46.014238Z","iopub.execute_input":"2026-03-15T14:48:46.014521Z","iopub.status.idle":"2026-03-15T14:48:46.229067Z","shell.execute_reply.started":"2026-03-15T14:48:46.014491Z","shell.execute_reply":"2026-03-15T14:48:46.228223Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## PNOC Model","metadata":{}},{"cell_type":"code","source":"#@title Model / Architecture\n\nclass PNOCClassifier(Backbone):\n\n  def __init__(\n    self,\n    model_name,\n    num_classes=20,\n    backbone_weights=\"imagenet\",\n    channels=3,\n    mode='fix',\n    dilated=False,\n    strides=None,\n    trainable_stem=True,\n    trainable_stage4=True,\n    trainable_backbone=True,\n    **backbone_kwargs,\n  ):\n    super().__init__(\n      model_name,\n      channels=channels,\n      weights=backbone_weights,\n      mode=mode,\n      dilated=dilated,\n      strides=strides,\n      trainable_stem=trainable_stem,\n      trainable_stage4=trainable_stage4,\n      trainable_backbone=trainable_backbone,\n      backbone_kwargs=backbone_kwargs,\n    )\n\n    self.num_classes = num_classes\n\n    cin = self.backbone.outplanes\n    self.classifier = nn.Conv2d(cin, num_classes, 1, bias=False)\n\n    self.from_scratch_layers.extend([self.classifier])\n    self.initialize([self.classifier])\n\n  def forward(self, x, with_cam=False):\n    outs = self.backbone(x)\n    x = outs[-1] if isinstance(outs, tuple) else outs\n\n    if with_cam:\n      features = self.classifier(x)\n      logits = gap2d(features)\n      return logits, features\n    else:\n      x = gap2d(x, keepdims=True)\n      logits = self.classifier(x).view(-1, self.num_classes)\n      return logits","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-03-15T14:48:46.230152Z","iopub.execute_input":"2026-03-15T14:48:46.230419Z","iopub.status.idle":"2026-03-15T14:48:46.242887Z","shell.execute_reply.started":"2026-03-15T14:48:46.230394Z","shell.execute_reply":"2026-03-15T14:48:46.242227Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## CSRM Model","metadata":{}},{"cell_type":"code","source":"#@title Model / Architecture\n\nfrom typing import Tuple, Union\n\n\nclass ASPPModule(nn.Module):\n\n  def __init__(self, inplanes, planes, kernel_size, padding, dilation, norm_fn=None):\n    super().__init__()\n    self.atrous_conv = nn.Conv2d(\n      inplanes, planes, kernel_size=kernel_size, stride=1, padding=padding, dilation=dilation, bias=False\n    )\n    self.bn = norm_fn(planes)\n    self.relu = nn.ReLU(inplace=True)\n\n    self.initialize([self.atrous_conv, self.bn])\n\n  def forward(self, x):\n    x = self.atrous_conv(x)\n    x = self.bn(x)\n    return self.relu(x)\n\n  def initialize(self, modules):\n    for m in modules:\n      if isinstance(m, nn.Conv2d):\n        torch.nn.init.kaiming_normal_(m.weight)\n      elif isinstance(m, nn.BatchNorm2d):\n        m.weight.data.fill_(1)\n        m.bias.data.zero_()\n\n\nclass ASPP(nn.Module):\n\n  def __init__(self, inplanes, output_stride, norm_fn):\n    super().__init__()\n\n    if output_stride == 16:\n      dilations = [1, 6, 12, 18]\n    elif output_stride == 8:\n      dilations = [1, 12, 24, 36]\n\n    self.aspp1 = ASPPModule(inplanes, 256, 1, padding=0, dilation=dilations[0], norm_fn=norm_fn)\n    self.aspp2 = ASPPModule(inplanes, 256, 3, padding=dilations[1], dilation=dilations[1], norm_fn=norm_fn)\n    self.aspp3 = ASPPModule(inplanes, 256, 3, padding=dilations[2], dilation=dilations[2], norm_fn=norm_fn)\n    self.aspp4 = ASPPModule(inplanes, 256, 3, padding=dilations[3], dilation=dilations[3], norm_fn=norm_fn)\n\n    self.global_avg_pool = nn.Sequential(\n      nn.AdaptiveAvgPool2d((1, 1)),\n      nn.Conv2d(inplanes, 256, 1, stride=1, bias=False),\n      norm_fn(256),\n      nn.ReLU(inplace=True),\n    )\n\n    self.conv1 = nn.Conv2d(1280, 256, 1, bias=False)\n    self.bn1 = norm_fn(256)\n    self.relu = nn.ReLU(inplace=True)\n    self.dropout = nn.Dropout(0.5)\n\n    self.initialize([self.conv1, self.bn1] + list(self.global_avg_pool.modules()))\n\n  def forward(self, x):\n    x1 = self.aspp1(x)\n    x2 = self.aspp2(x)\n    x3 = self.aspp3(x)\n    x4 = self.aspp4(x)\n\n    x5 = self.global_avg_pool(x)\n    x5 = F.interpolate(x5, size=x4.size()[2:], mode='bilinear', align_corners=True)\n\n    x = torch.cat((x1, x2, x3, x4, x5), dim=1)\n\n    x = self.conv1(x)\n    x = self.bn1(x)\n    x = self.relu(x)\n    x = self.dropout(x)\n\n    return x\n\n  def initialize(self, modules):\n    for m in modules:\n      if isinstance(m, nn.Conv2d):\n        torch.nn.init.kaiming_normal_(m.weight)\n      elif isinstance(m, nn.BatchNorm2d):\n        m.weight.data.fill_(1)\n        m.bias.data.zero_()\n\n\ndef resize_tensor(tensors, size, mode='bilinear', align_corners=None):\n  return F.interpolate(tensors, size, mode=mode, align_corners=align_corners)\n\n\ndef gap2d(x, keepdims=False):\n  x = torch.mean(x.view(x.size(0), x.size(1), -1), -1)\n  if keepdims:\n    x = x.view(x.size(0), x.size(1), 1, 1)\n  return x\n    \n\nclass CSRM(Backbone):\n\n  def __init__(\n    self,\n    model_name,\n    num_classes=20,\n    num_classes_segm=None,\n    channels=3,\n    backbone_weights=\"imagenet\",\n    mode='fix',\n    dilated=False,\n    strides=None,\n    trainable_stem=False,\n    trainable_stage4=True,\n    trainable_backbone=True,\n    use_group_norm=False,\n    use_sal_head=False,\n    use_rep_head=True,\n    rep_output_dim=256,\n    dropout: Union[float, Tuple[float, float]] = 0.1,\n  ):\n    super().__init__(\n      model_name,\n      channels=channels,\n      weights=backbone_weights,\n      mode=mode,\n      dilated=dilated,\n      strides=strides,\n      trainable_stem=trainable_stem,\n      trainable_stage4=trainable_stage4,\n      trainable_backbone=trainable_backbone,\n    )\n\n    low_level_cin = self.backbone.stage_features[0]\n    cin = self.backbone.outplanes\n    norm_fn = group_norm if use_group_norm else nn.BatchNorm2d\n\n    self.num_classes = num_classes\n    self.num_classes_segm = num_classes_segm or num_classes + 1\n    self.use_sal_head = use_sal_head\n    self.use_rep_head = use_rep_head\n    self.dropout = [dropout, dropout] if isinstance(dropout, float) else dropout\n\n    ## Pretrained parameters\n    # self.backbone = ...\n    self.classifier = nn.Conv2d(cin, self.num_classes, 1, bias=False)\n\n    ## Scratch parameters\n    self.aspp = ASPP(cin, output_stride=16, norm_fn=norm_fn)\n\n    self.project = nn.Sequential(\n      nn.Conv2d(low_level_cin, 48, 1, bias=False),\n      norm_fn(48),\n      nn.ReLU(inplace=True),\n    )\n\n    self.decoder = nn.Sequential(\n      nn.Conv2d(304, 256, 3, padding=1, bias=False),\n      norm_fn(256),\n      nn.ReLU(inplace=True),\n      nn.Dropout(self.dropout[0]),\n      nn.Conv2d(256, 256, 3, padding=1, bias=False),\n      norm_fn(256),\n      nn.ReLU(inplace=True),\n      nn.Dropout(self.dropout[1]),\n      nn.Conv2d(256, self.num_classes_segm, 1),\n    )\n\n    self.from_scratch_layers += [*self.aspp.modules(), *self.project.modules(), *self.decoder.modules()]\n\n    if use_sal_head:\n      self.saliency_head = nn.Conv2d(304, 1, 1)\n      self.from_scratch_layers += [self.saliency_head]\n\n    if use_rep_head:\n      self.representation = nn.Sequential(\n        nn.Conv2d(304, 256, 3, padding=1, bias=False),\n        norm_fn(256),\n        nn.ReLU(inplace=True),\n        nn.Conv2d(256, rep_output_dim, 1)\n      )\n      self.from_scratch_layers += [*self.representation.modules()]\n\n    self.initialize(self.from_scratch_layers)\n\n  def forward_features(self, x):\n    return self.backbone(x)\n\n  def forward(self, inputs, with_cam=False, with_saliency=False, with_mask=True, with_rep=False, resize_mask=True):\n    outs = self.forward_features(inputs)\n    features_s2 = outs[0]\n    features_s5 = outs[-1]\n\n    outputs = self.classification_branch(features_s5, with_cam=with_cam or with_saliency)\n\n    if not (with_saliency or with_mask or with_rep):\n      return outputs\n\n    features_s2 = self.project(features_s2)\n    features = self.aspp(features_s5)\n\n    features = resize_tensor(features, features_s2.size()[2:], align_corners=True)\n    features = torch.cat((features, features_s2), dim=1)\n\n    if with_saliency:\n      masks = self.saliency_branch(outputs[\"features_c\"], features)\n      outputs[\"masks_sal\"] = masks\n      if resize_mask:\n        outputs[\"masks_sal_large\"] = resize_tensor(masks, inputs.shape[2:], align_corners=True)\n\n    if with_mask:\n      masks = self.decoder(features)\n      outputs[\"masks_seg\"] = masks\n      if resize_mask:\n        outputs[\"masks_seg_large\"] = resize_tensor(masks, inputs.shape[2:], align_corners=True)\n\n    if with_rep and self.use_rep_head:\n      rep = self.representation(features)\n      outputs[\"rep\"] = rep\n\n    return outputs\n\n  def saliency_branch(self, features_c, features):\n    if not self.use_sal_head:\n      return features_c.sum(dim=1, keepdim=True)\n\n    x = self.saliency_head(features)\n    return x\n\n  def classification_branch(self, features, with_cam=False):\n    if with_cam:\n      features = self.classifier(features)\n      logits = gap2d(features)\n      return {\n        \"logits_c\": logits,\n        \"features_c\": features,\n      }\n\n    features = gap2d(features, keepdims=True)\n    logits = self.classifier(features).view(-1, self.num_classes)\n    return {\"logits\": logits}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:48:46.244143Z","iopub.execute_input":"2026-03-15T14:48:46.244476Z","iopub.status.idle":"2026-03-15T14:48:46.272462Z","shell.execute_reply.started":"2026-03-15T14:48:46.244443Z","shell.execute_reply":"2026-03-15T14:48:46.271639Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load ResNest269 Model","metadata":{}},{"cell_type":"code","source":"models = []\nimage_models = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:48:46.273538Z","iopub.execute_input":"2026-03-15T14:48:46.273863Z","iopub.status.idle":"2026-03-15T14:48:46.285271Z","shell.execute_reply.started":"2026-03-15T14:48:46.27383Z","shell.execute_reply":"2026-03-15T14:48:46.2845Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# WEIGHTS_PATHS = [\n#     '/kaggle/input/baseline-rs50/pytorch/default/1/hpa2nd-kaggle-rs50-lr0.0002-b16-aug2nd-adamw-eid-baseline-2/model-f0-best.pth']\n# DEVICE = 'cuda'\n\n# for weights_path in WEIGHTS_PATHS:\n#     model = CellClassifier(\n#         'resnest50',\n#         19,\n#         channels=4,\n#         mode='normal',\n#         dilated=False,\n#         backbone_weights=None,\n#     )\n#     load_model(model, weights_path, map_location=torch.device(DEVICE))\n#     model.eval()\n\n#     models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:48:46.28642Z","iopub.execute_input":"2026-03-15T14:48:46.286744Z","iopub.status.idle":"2026-03-15T14:48:46.295147Z","shell.execute_reply.started":"2026-03-15T14:48:46.286711Z","shell.execute_reply":"2026-03-15T14:48:46.294385Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"WEIGHTS_PATHS = [\n    '/kaggle/input/models/julianamidlej/lucas-vanilla/pytorch/default/1/hpa-512-rs269-lr0.1-b32-ls0.1-re-mix-sgd-ema1-cp-1.0-0.5-sum-2.pth']\nDEVICE = 'cuda'\n\nfor weights_path in WEIGHTS_PATHS:\n    model = ImageClassifier(\n        'resnest269',\n        19,\n        channels=4,\n        mode='normal',\n        dilated=False,\n        backbone_weights=None,\n    )\n    load_model(model, weights_path, map_location=torch.device(DEVICE))\n    model.eval()\n\n    image_models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:48:46.295996Z","iopub.execute_input":"2026-03-15T14:48:46.296209Z","iopub.status.idle":"2026-03-15T14:49:00.679068Z","shell.execute_reply.started":"2026-03-15T14:48:46.296177Z","shell.execute_reply":"2026-03-15T14:49:00.678402Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load CSRM Model","metadata":{}},{"cell_type":"code","source":"# WEIGHTS_PATHS = [\n#     '/kaggle/input/models/julianamidlej/csrm-pnoc/pytorch/default/1/hpa-rs269-csrm-lr0.007-m0.9-b32-ls0.1-classmix-r1.pth']\n# DEVICE = 'cuda'\n\n# for weights_path in WEIGHTS_PATHS:\n#     model = CSRM(\n#         'resnest269',\n#         19,\n#         channels=4,\n#         mode='normal',\n#         dilated=False,\n#         backbone_weights=None,\n#     )\n#     load_model(model, weights_path, map_location=torch.device(DEVICE), strict=False)\n#     model.eval()\n\n#     image_models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:00.680017Z","iopub.execute_input":"2026-03-15T14:49:00.680228Z","iopub.status.idle":"2026-03-15T14:49:00.683235Z","shell.execute_reply.started":"2026-03-15T14:49:00.680206Z","shell.execute_reply":"2026-03-15T14:49:00.682522Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loading models\n* b3\n* b5\n* r50d\n* r200d\n* se50","metadata":{}},{"cell_type":"code","source":"# ckpt = {\n#     0: 16\n# }\n\n# for i in range(1):\n#     cfg = Config.load_json('/kaggle/input/b3-f4/pytorch/default/1/b3_F4/config.json')\n#     model = get_model(cfg).cuda()\n#     load_matched_state(model, torch.load(\n#         f'/kaggle/input/b3-f4/pytorch/default/1/b3_F4/checkpoints/f4_epoch-18.pth'))\n#     _ = model.eval()\n#     models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:00.686416Z","iopub.execute_input":"2026-03-15T14:49:00.686641Z","iopub.status.idle":"2026-03-15T14:49:00.710348Z","shell.execute_reply.started":"2026-03-15T14:49:00.686606Z","shell.execute_reply":"2026-03-15T14:49:00.709734Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ckpt = {\n#     0: 8\n# }\n\n# for i in range(1):\n#     # if i in [2, 3, 4]: continue\n#     cfg = Config.load_json('/kaggle/input/b5/pytorch/default/1/b5/config.json')\n#     model = get_model(cfg).cuda()\n#     load_matched_state(model, torch.load(\n#         f'/kaggle/input/b5/pytorch/default/1/b5/checkpoints/f0_epoch-14.pth'))\n#     _ = model.eval()\n#     models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:00.712679Z","iopub.execute_input":"2026-03-15T14:49:00.712992Z","iopub.status.idle":"2026-03-15T14:49:00.719162Z","shell.execute_reply.started":"2026-03-15T14:49:00.71296Z","shell.execute_reply":"2026-03-15T14:49:00.718597Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ckpt = {\n#     0: 12\n# }\n\n# for i in range(1):\n#     # if i in [0, 3, 4]: continue\n#     cfg = Config.load_json('/kaggle/input/r50-f0/pytorch/default/1/r50_F0/config.json')\n#     model = get_model(cfg).cuda()\n#     load_matched_state(model, torch.load(\n#         f'/kaggle/input/r50-f0/pytorch/default/1/r50_F0/checkpoints/f0_epoch-14.pth'))\n#     _ = model.eval()\n#     models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:00.720033Z","iopub.execute_input":"2026-03-15T14:49:00.720249Z","iopub.status.idle":"2026-03-15T14:49:00.726968Z","shell.execute_reply.started":"2026-03-15T14:49:00.720229Z","shell.execute_reply":"2026-03-15T14:49:00.726168Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !cp ../input/landmark-additional-packages/resnet200d_ra2-bdba9bf9.pth /root/.cache/torch/hub/checkpoints/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:00.72788Z","iopub.execute_input":"2026-03-15T14:49:00.728103Z","iopub.status.idle":"2026-03-15T14:49:00.735823Z","shell.execute_reply.started":"2026-03-15T14:49:00.728057Z","shell.execute_reply":"2026-03-15T14:49:00.735115Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ckpt = {\n#     0: 13\n# }\n\n# for i in range(1):\n#     # if i in [0, 1, 4]: continue\n#     cfg = Config.load_json('/kaggle/input/r200/pytorch/default/1/r200d/config.json')\n#     model = get_model(cfg).cuda()\n#     load_matched_state(model, torch.load(\n#         f'/kaggle/input/r200/pytorch/default/1/r200d/checkpoints/f0_epoch-13.pth'))\n#     _ = model.eval()\n#     models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:00.736633Z","iopub.execute_input":"2026-03-15T14:49:00.736814Z","iopub.status.idle":"2026-03-15T14:49:00.746398Z","shell.execute_reply.started":"2026-03-15T14:49:00.736796Z","shell.execute_reply":"2026-03-15T14:49:00.745746Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ckpt = {\n#     0: 18\n# }\n\n# for i in range(1):\n#     # if i in [0, 1, 2]: continue\n#     cfg = Config.load_json('/kaggle/input/se50/pytorch/default/1/se50/config.json')\n#     model = get_model(cfg).cuda()\n#     load_matched_state(model, torch.load(\n#         f'/kaggle/input/se50/pytorch/default/1/se50/checkpoints/f0_epoch-18.pth'))\n#     _ = model.eval()\n#     models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:00.747195Z","iopub.execute_input":"2026-03-15T14:49:00.747429Z","iopub.status.idle":"2026-03-15T14:49:00.759914Z","shell.execute_reply.started":"2026-03-15T14:49:00.747407Z","shell.execute_reply":"2026-03-15T14:49:00.759368Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(models))\nprint(len(image_models))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:00.76076Z","iopub.execute_input":"2026-03-15T14:49:00.760959Z","iopub.status.idle":"2026-03-15T14:49:00.769694Z","shell.execute_reply.started":"2026-03-15T14:49:00.760939Z","shell.execute_reply":"2026-03-15T14:49:00.768811Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell Segmentator","metadata":{}},{"cell_type":"code","source":"def read_img(image_id_path):\n    img = cv2.imread(image_id_path, 0)\n    return img\n\ndef load_RGBY_images(image_id_path):\n    \n    red_image = read_img(image_id_path+\"_red.png\")\n    green_image = read_img(image_id_path+\"_green.png\")\n    blue_image = read_img(image_id_path+\"_blue.png\")\n    yellow_image = read_img(image_id_path+\"_yellow.png\")\n    \n    return red_image, green_image, blue_image, yellow_image\n\ndef encode_binary_mask(mask):\n\n    if mask.dtype != np.bool:\n        raise ValueError(\"encode_binary_mask expects a binary mask, received dtype == %s\" % mask.dtype)\n\n    mask = np.squeeze(mask)\n    if len(mask.shape) != 2:\n        raise ValueError(\"encode_binary_mask expects a 2d mask, received shape == %s\" % mask.shape)\n    \n    mask = mask.astype(np.uint8)\n    \n    mask_to_encode = mask.reshape(mask.shape[0], mask.shape[1], 1)\n    mask_to_encode = np.asfortranarray(mask_to_encode)\n    \n    encoded_mask = coco_mask.encode(mask_to_encode)[0][\"counts\"]\n    \n    binary_str = zlib.compress(encoded_mask, zlib.Z_BEST_COMPRESSION)\n    base64_str = base64.b64encode(binary_str)\n    \n    return base64_str.decode()\n\ndef compute_M(data):\n    cols = np.arange(data.size)\n    return csr_matrix((cols, (data.ravel(), cols)), shape=(data.max() + 1, data.size))\n\ndef get_indices_sparse(data):\n    M = compute_M(data)\n    return [np.unravel_index(row.data, data.shape) for row in M]\n\nclass CellSegmentator(object):\n    \"\"\"Uses pretrained DPN-Unet models to segment cells from images.\"\"\"\n\n    NORMALIZE = {\"mean\": [124 / 255, 117 / 255, 104 / 255], \"std\": [1 / (0.0167 * 255)] * 3}\n    \n    def __init__(\n            self,\n            nuclei_model=\"./nuclei_model.pth\",\n            cell_model=\"./cell_model.pth\",\n            model_width_height=None,\n            device=\"cuda\",\n            multi_channel_model=True,\n            return_without_scale_restore=False,\n            scale_factor=0.25,\n            padding=False\n    ):\n\n        if device != \"cuda\" and device != \"cpu\" and \"cuda\" not in device:\n            raise ValueError(f\"{device} is not a valid device (cuda/cpu)\")\n        if device != \"cpu\":\n            try:\n                assert torch.cuda.is_available()\n            except AssertionError:\n                print(\"No GPU found, using CPU.\", file=sys.stderr)\n                device = \"cpu\"\n                \n        self.device = device\n\n        if isinstance(nuclei_model, str):\n            if not os.path.exists(nuclei_model):\n                print(\n                    f\"Could not find {nuclei_model}. Downloading it now\",\n                    file=sys.stderr,\n                )\n                download_with_url(NUCLEI_MODEL_URL, nuclei_model)\n            nuclei_model = torch.load(\n                nuclei_model, map_location=torch.device(self.device)\n            )\n        if isinstance(nuclei_model, torch.nn.DataParallel) and device == \"cpu\":\n            nuclei_model = nuclei_model.module\n\n        self.nuclei_model = nuclei_model.to(self.device)\n\n        self.multi_channel_model = multi_channel_model\n        if isinstance(cell_model, str):\n            if not os.path.exists(cell_model):\n                print(\n                    f\"Could not find {cell_model}. Downloading it now\", file=sys.stderr\n                )\n                if self.multi_channel_model:\n                    download_with_url(MULTI_CHANNEL_CELL_MODEL_URL, cell_model)\n                else:\n                    download_with_url(TWO_CHANNEL_CELL_MODEL_URL, cell_model)\n            cell_model = torch.load(cell_model, map_location=torch.device(self.device))\n        self.cell_model = cell_model.to(self.device)\n        self.model_width_height = model_width_height\n        self.return_without_scale_restore = return_without_scale_restore\n        self.scale_factor = scale_factor\n        self.padding = padding\n\n    def _image_conversion(self, images):\n\n        microtubule_imgs, er_imgs, nuclei_imgs = images\n        if self.multi_channel_model:\n            if not isinstance(er_imgs, list):\n                raise ValueError(\"Please speicify the image path(s) for er channels!\")\n        else:\n            if not er_imgs is None:\n                raise ValueError(\n                    \"second channel should be None for two channel model predition!\"\n                )\n\n        if not isinstance(microtubule_imgs, list):\n            raise ValueError(\"The microtubule images should be a list\")\n        if not isinstance(nuclei_imgs, list):\n            raise ValueError(\"The microtubule images should be a list\")\n\n        if er_imgs:\n            if not len(microtubule_imgs) == len(er_imgs) == len(nuclei_imgs):\n                raise ValueError(\"The lists of images needs to be the same length\")\n        else:\n            if not len(microtubule_imgs) == len(nuclei_imgs):\n                raise ValueError(\"The lists of images needs to be the same length\")\n\n        if not all(isinstance(item, np.ndarray) for item in microtubule_imgs):\n            microtubule_imgs = [\n                os.path.expanduser(item) for _, item in enumerate(microtubule_imgs)\n            ]\n            nuclei_imgs = [\n                os.path.expanduser(item) for _, item in enumerate(nuclei_imgs)\n            ]\n\n            microtubule_imgs = list(\n                map(lambda item: imageio.imread(item), microtubule_imgs)\n            )\n            nuclei_imgs = list(map(lambda item: imageio.imread(item), nuclei_imgs))\n            if er_imgs:\n                er_imgs = [os.path.expanduser(item) for _, item in enumerate(er_imgs)]\n                er_imgs = list(map(lambda item: imageio.imread(item), er_imgs))\n\n        if not er_imgs:\n            er_imgs = [\n                np.zeros(item.shape, dtype=item.dtype)\n                for _, item in enumerate(microtubule_imgs)\n            ]\n        cell_imgs = list(\n            map(\n                lambda item: np.dstack((item[0], item[1], item[2])),\n                list(zip(microtubule_imgs, er_imgs, nuclei_imgs)),\n            )\n        )\n\n        return cell_imgs\n\n    def _pad(self, image):\n        \n        rows, cols = image.shape[:2]\n        self.scaled_shape = rows, cols\n        img_pad= cv2.copyMakeBorder(\n                    image,\n                    32,\n                    (32 - rows % 32),\n                    32,\n                    (32 - cols % 32),\n                    cv2.BORDER_REFLECT,\n                )\n        \n        return img_pad\n\n    def pred_nuclei(self, images):\n\n        def _preprocess(images):\n            if isinstance(images[0], str):\n                raise NotImplementedError('Currently the model requires images as numpy arrays, not paths.')\n                # images = [imageio.imread(image_path) for image_path in images]\n            self.target_shapes = [image.shape for image in images]\n            #print(images.shape)\n            #resize like in original implementation with https://scikit-image.org/docs/dev/api/skimage.transform.html#skimage.transform.resize\n            if self.model_width_height:\n                images = np.array([transform.resize(image, (self.model_width_height,self.model_width_height)) \n                                  for image in images])\n            else:\n                images = [transform.rescale(image, self.scale_factor) for image in images]\n\n            if self.padding:\n                images = [self._pad(image) for image in images]\n\n            nuc_images = np.array([np.dstack((image[..., 2], image[..., 2], image[..., 2])) if len(image.shape) >= 3\n                                   else np.dstack((image, image, image)) for image in images])\n            \n            nuc_images = nuc_images.transpose([0, 3, 1, 2])\n            #print(\"nuc\", nuc_images.shape)\n\n            return nuc_images\n\n        def _segment_helper(imgs):\n            with torch.no_grad():\n                mean = torch.as_tensor(self.NORMALIZE[\"mean\"], device=self.device)\n                std = torch.as_tensor(self.NORMALIZE[\"std\"], device=self.device)\n                imgs = torch.tensor(imgs).float()\n                imgs = imgs.to(self.device)\n                imgs = imgs.sub_(mean[:, None, None]).div_(std[:, None, None])\n\n                imgs = self.nuclei_model(imgs)\n                imgs = F.softmax(imgs, dim=1)\n                return imgs\n\n        preprocessed_imgs = _preprocess(images)\n        predictions = _segment_helper(preprocessed_imgs)\n        predictions = predictions.to(\"cpu\").numpy()\n        #dont restore scaling, just save and scale later ...\n        predictions = [self._restore_scaling(util.img_as_ubyte(pred), target_shape)\n                       for pred, target_shape in zip(predictions, self.target_shapes)]\n        return predictions\n\n    def _restore_scaling(self, n_prediction, target_shape):\n        \"\"\"Restore an image from scaling and padding.\n        This method is intended for internal use.\n        It takes the output from the nuclei model as input.\n        \"\"\"\n        n_prediction = n_prediction.transpose([1, 2, 0])\n        if self.padding:\n            n_prediction = n_prediction[\n                32 : 32 + self.scaled_shape[0], 32 : 32 + self.scaled_shape[1], ...\n            ]\n        n_prediction[..., 0] = 0\n        if not self.return_without_scale_restore:\n            n_prediction = cv2.resize(\n                n_prediction,\n                (target_shape[0], target_shape[1]),\n                #try INTER_NEAREST_EXACT\n                interpolation=cv2.INTER_AREA,\n            )\n        return n_prediction\n\n    def pred_cells(self, images, precombined=False):\n\n        def _preprocess(images):\n            self.target_shapes = [image.shape for image in images]\n            for image in images:\n                if not len(image.shape) == 3:\n                    raise ValueError(\"image should has 3 channels\")\n            #resize like in original implementation with https://scikit-image.org/docs/dev/api/skimage.transform.html#skimage.transform.resize\n            if self.model_width_height:\n                images = np.array([transform.resize(image, (self.model_width_height,self.model_width_height)) \n                                  for image in images])\n            else:\n                images = np.array([transform.rescale(image, self.scale_factor, multichannel=True) for image in images])\n\n            if self.padding:\n                images = np.array([self._pad(image) for image in images])\n\n            cell_images = images.transpose([0, 3, 1, 2])\n\n            return cell_images\n\n        def _segment_helper(imgs):\n            with torch.no_grad():\n                mean = torch.as_tensor(self.NORMALIZE[\"mean\"], device=self.device)\n                std = torch.as_tensor(self.NORMALIZE[\"std\"], device=self.device)\n                imgs = torch.tensor(imgs).float()\n                imgs = imgs.to(self.device)\n                imgs = imgs.sub_(mean[:, None, None]).div_(std[:, None, None])\n                imgs = self.cell_model(imgs)\n                imgs = F.softmax(imgs, dim=1)\n                return imgs\n\n        if not precombined:\n            images = self._image_conversion(images)\n        preprocessed_imgs = _preprocess(images)\n        predictions = _segment_helper(preprocessed_imgs)\n        predictions = predictions.to(\"cpu\").numpy()\n        predictions = [self._restore_scaling(util.img_as_ubyte(pred), target_shape)\n                       for pred, target_shape in zip(predictions, self.target_shapes)]\n        \n        return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:00.77087Z","iopub.execute_input":"2026-03-15T14:49:00.771105Z","iopub.status.idle":"2026-03-15T14:49:00.803303Z","shell.execute_reply.started":"2026-03-15T14:49:00.771082Z","shell.execute_reply":"2026-03-15T14:49:00.802538Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def __fill_holes(image):\n    \"\"\"Fill_holes for labelled image, with a unique number.\"\"\"\n    boundaries = segmentation.find_boundaries(image)\n    image = np.multiply(image, np.invert(boundaries))\n    image = ndi.binary_fill_holes(image > 0)\n    image = ndi.label(image)[0]\n    return image\n\ndef label_cell_vlad(nuclei_pred, cell_pred, img_size=512, return_nuclei_label=True):\n    \"\"\"Label the cells and the nuclei.\n\n    Keyword arguments:\n    nuclei_pred -- a 3D numpy array of a prediction from a nuclei image.\n    cell_pred -- a 3D numpy array of a prediction from a cell image.\n\n    Returns:\n    A tuple containing:\n    nuclei-label -- A nuclei mask data array.\n    cell-label  -- A cell mask data array.\n\n    0's in the data arrays indicate background while a continous\n    strech of a specific number indicates the area for a specific\n    cell.\n    The same value in cell mask and nuclei mask refers to the identical cell.\n\n    NOTE: The nuclei labeling from this function will be sligthly\n    different from the values in :func:`label_nuclei` as this version\n    will use information from the cell-predictions to make better\n    estimates.\n    \"\"\"\n    def __wsh(\n        mask_img,\n        threshold,\n        border_img,\n        seeds,\n        threshold_adjustment=0.35,\n        small_object_size_cutoff=10,\n    ):\n        img_copy = np.zeros_like(mask_img)\n        m = seeds * border_img  # * dt\n        img_copy[m > threshold + threshold_adjustment] = 1\n        img_copy = img_copy.astype(np.bool)\n        img_copy = remove_small_objects(img_copy, small_object_size_cutoff).astype(\n            np.uint8\n        )\n\n        mask_img = np.where(mask_img <= threshold, 0, 1)\n        mask_img = mask_img.astype(np.bool)\n        \n        ### New segmentation ###\n        mask_img = remove_small_holes(mask_img, int(63 * (img_size / 512)**2))\n        ########################\n        \n        mask_img = remove_small_objects(mask_img, 8).astype(np.uint8)\n        markers = ndi.label(img_copy, output=np.uint32)[0]\n        labeled_array = segmentation.watershed(\n            mask_img, markers, mask=mask_img, watershed_line=True\n        )\n        return labeled_array\n\n    nuclei_label = __wsh(\n        nuclei_pred[..., 2] / 255.0,\n        0.4,\n        1 - (nuclei_pred[..., 1] + cell_pred[..., 1]) / 255.0 > 0.05,\n        nuclei_pred[..., 2] / 255,\n        threshold_adjustment=-0.25,\n        small_object_size_cutoff=32,\n    )\n\n    # for hpa_image, to remove the small pseduo nuclei\n    \n    ### New segmentation ###\n    nuclei_label = remove_small_objects(nuclei_label, int(157 * (img_size / 512)**2))\n    ########################\n    \n    nuclei_label = measure.label(nuclei_label)\n    # this is to remove the cell borders' signal from cell mask.\n    # could use np.logical_and with some revision, to replace this func.\n    # Tuned for segmentation hpa images\n    threshold_value = max(0.22, filters.threshold_otsu(cell_pred[..., 2] / 255) * 0.5)\n    # exclude the green area first\n    cell_region = np.multiply(\n        cell_pred[..., 2] / 255 > threshold_value,\n        np.invert(np.asarray(cell_pred[..., 1] / 255 > 0.05, dtype=np.int8)),\n    )\n    sk = np.asarray(cell_region, dtype=np.int8)\n    \n    ################### CHANGES HERE ###################################\n \n    ####################### V4 ####################\n    distance_old = np.clip(cell_pred[..., 2], 255 * threshold_value, cell_pred[..., 2])\n    cell_label_old = segmentation.watershed(-distance_old, nuclei_label, mask=sk)\n    \n    ### New segmentation ###\n    cell_label_old = remove_small_objects(cell_label_old, int(344 * (img_size / 512)**2)).astype(np.uint8)\n    ########################\n    \n    distance = distance_old.copy()\n    distance[distance<225] = 0\n    \n    cell_label = segmentation.watershed(-distance, nuclei_label, mask=distance)\n    \n    ### New segmentation ###\n    cell_label = remove_small_objects(cell_label, int(344 * (img_size / 512)**2)).astype(np.uint8)\n    ########################\n    \n    unqs = np.unique(cell_label)\n    if 0 in unqs:\n        unqs = unqs[1:]\n        \n    lst = [cv2.dilate((cell_label==unq).astype(np.uint8), kernel=np.ones((10, 10), np.uint8), iterations=1) for unq in unqs]\n    cell_label = np.zeros_like(cell_label)\n    \n    for i, l in enumerate(lst):\n        cell_label[l==True] = unqs[i]\n                \n    for unq in unqs:\n        smth = cell_label_old==unq\n        if True not in smth[0,:] and True not in smth[-1,:] and True not in smth[:,0] and True not in smth[:,-1]:\n            cell_label[cell_label==unq] = 0\n            cell_label[cell_label_old==unq] = unq\n            \n    ################### CHANGES HERE ###################################\n    \n#     selem = disk(max(1, int(6 * 2048 / img_size)))\n    selem = disk(6)\n    cell_label = closing(cell_label, selem)\n    cell_label = __fill_holes(cell_label)\n    # this part is to use green channel, and extend cell label to green channel\n    # benefit is to exclude cells clear on border but without nucleus\n    sk = np.asarray(\n        np.add(\n            np.asarray(cell_label > 0, dtype=np.int8),\n            np.asarray(cell_pred[..., 1] / 255 > 0.05, dtype=np.int8),\n        )\n        > 0,\n        dtype=np.int8,\n    )\n    cell_label = segmentation.watershed(-distance, cell_label, mask=sk)\n    cell_label = __fill_holes(cell_label)\n    cell_label = np.asarray(cell_label > 0, dtype=np.uint8)\n    cell_label = measure.label(cell_label)\n    \n    ### New segmentation ###\n    cell_label = remove_small_objects(cell_label, int(344 * (img_size / 512)**2))\n    ########################\n    \n    cell_label = measure.label(cell_label)\n    cell_label = np.asarray(cell_label, dtype=np.uint16)\n    if not return_nuclei_label:\n        return cell_label\n    nuclei_label = np.multiply(cell_label > 0, nuclei_label) > 0\n    nuclei_label = measure.label(nuclei_label)\n    \n    ### New segmentation ###\n    nuclei_label = remove_small_objects(nuclei_label, int(157 * (img_size / 512)**2))\n    ########################\n    \n    nuclei_label = np.multiply(cell_label, nuclei_label > 0)\n\n    return nuclei_label, cell_label\n\ndef label_cell(nuclei_pred, cell_pred):\n    \"\"\"Label the cells and the nuclei.\n\n    Keyword arguments:\n    nuclei_pred -- a 3D numpy array of a prediction from a nuclei image.\n    cell_pred -- a 3D numpy array of a prediction from a cell image.\n\n    Returns:\n    A tuple containing:\n    nuclei-label -- A nuclei mask data array.\n    cell-label  -- A cell mask data array.\n\n    0's in the data arrays indicate background while a continous\n    strech of a specific number indicates the area for a specific\n    cell.\n    The same value in cell mask and nuclei mask refers to the identical cell.\n\n    NOTE: The nuclei labeling from this function will be sligthly\n    different from the values in :func:`label_nuclei` as this version\n    will use information from the cell-predictions to make better\n    estimates.\n    \"\"\"\n    def __wsh(\n        mask_img,\n        threshold,\n        border_img,\n        seeds,\n        threshold_adjustment=0.35,\n        small_object_size_cutoff=10,\n    ):\n        img_copy = np.copy(mask_img)\n        m = seeds * border_img  # * dt\n        img_copy[m <= threshold + threshold_adjustment] = 0\n        img_copy[m > threshold + threshold_adjustment] = 1\n        img_copy = img_copy.astype(np.bool)\n        img_copy = remove_small_objects(img_copy, small_object_size_cutoff).astype(\n            np.uint8\n        )\n\n        mask_img[mask_img <= threshold] = 0\n        mask_img[mask_img > threshold] = 1\n        mask_img = mask_img.astype(np.bool)\n        mask_img = remove_small_holes(mask_img, 1000)\n        mask_img = remove_small_objects(mask_img, 8).astype(np.uint8)\n        markers = ndi.label(img_copy, output=np.uint32)[0]\n        labeled_array = segmentation.watershed(\n            mask_img, markers, mask=mask_img, watershed_line=True\n        )\n        return labeled_array\n\n    nuclei_label = __wsh(\n        nuclei_pred[..., 2] / 255.0,\n        0.4,\n        1 - (nuclei_pred[..., 1] + cell_pred[..., 1]) / 255.0 > 0.05,\n        nuclei_pred[..., 2] / 255,\n        threshold_adjustment=-0.25,\n        small_object_size_cutoff=500,\n    )\n\n    # for hpa_image, to remove the small pseduo nuclei\n    nuclei_label = remove_small_objects(nuclei_label, 2500)\n    nuclei_label = measure.label(nuclei_label)\n    # this is to remove the cell borders' signal from cell mask.\n    # could use np.logical_and with some revision, to replace this func.\n    # Tuned for segmentation hpa images\n    threshold_value = max(0.22, filters.threshold_otsu(cell_pred[..., 2] / 255) * 0.5)\n    # exclude the green area first\n    cell_region = np.multiply(\n        cell_pred[..., 2] / 255 > threshold_value,\n        np.invert(np.asarray(cell_pred[..., 1] / 255 > 0.05, dtype=np.int8)),\n    )\n    sk = np.asarray(cell_region, dtype=np.int8)\n    distance = np.clip(cell_pred[..., 2], 255 * threshold_value, cell_pred[..., 2])\n    cell_label = segmentation.watershed(-distance, nuclei_label, mask=sk)\n    cell_label = remove_small_objects(cell_label, 5500).astype(np.uint8)\n    selem = disk(6)\n    cell_label = closing(cell_label, selem)\n    cell_label = __fill_holes(cell_label)\n    # this part is to use green channel, and extend cell label to green channel\n    # benefit is to exclude cells clear on border but without nucleus\n    sk = np.asarray(\n        np.add(\n            np.asarray(cell_label > 0, dtype=np.int8),\n            np.asarray(cell_pred[..., 1] / 255 > 0.05, dtype=np.int8),\n        )\n        > 0,\n        dtype=np.int8,\n    )\n    cell_label = segmentation.watershed(-distance, cell_label, mask=sk)\n    cell_label = __fill_holes(cell_label)\n    cell_label = np.asarray(cell_label > 0, dtype=np.uint8)\n    cell_label = measure.label(cell_label)\n    cell_label = remove_small_objects(cell_label, 5500)\n    cell_label = measure.label(cell_label)\n    cell_label = np.asarray(cell_label, dtype=np.uint16)\n    nuclei_label = np.multiply(cell_label > 0, nuclei_label) > 0\n    nuclei_label = measure.label(nuclei_label)\n    nuclei_label = remove_small_objects(nuclei_label, 2500)\n    nuclei_label = np.multiply(cell_label, nuclei_label > 0)\n\n    return nuclei_label, cell_label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:00.804443Z","iopub.execute_input":"2026-03-15T14:49:00.804732Z","iopub.status.idle":"2026-03-15T14:49:00.832034Z","shell.execute_reply.started":"2026-03-15T14:49:00.804693Z","shell.execute_reply":"2026-03-15T14:49:00.831248Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2nd\n# NUC_MODEL = '/kaggle/input/hpamisc/HPA-Cell-Segmentation-weights/dpn_unet_nuclei_v1.pth'\n# CELL_MODEL = '/kaggle/input/hpamisc/HPA-Cell-Segmentation-weights/dpn_unet_cell_3ch_v1.pth'\nimport os\n\n# Lucas\nNUC_MODEL = '/kaggle/input/hpacellsegmentatormodelweights/dpn_unet_nuclei_v1.pth'\nCELL_MODEL = '/kaggle/input/hpacellsegmentatormodelweights/dpn_unet_cell_3ch_v1.pth'\n\nSEGMENTATION_SCALE=0.25\nsegmentator = CellSegmentator(\n    NUC_MODEL,\n    CELL_MODEL,\n    device=\"cuda\",\n    multi_channel_model=True,\n    scale_factor=SEGMENTATION_SCALE,\n    padding=True,\n    return_without_scale_restore=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:00.833127Z","iopub.execute_input":"2026-03-15T14:49:00.833397Z","iopub.status.idle":"2026-03-15T14:49:09.916027Z","shell.execute_reply.started":"2026-03-15T14:49:00.833374Z","shell.execute_reply":"2026-03-15T14:49:09.915128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def remove_small_nuclei_on_border(nuclei_mask, nuclei_area_thresh=0.5):\n    \n    h,w = nuclei_mask.shape[:2]\n    \n    touching_border = []\n    \n    nuclei_uniques = np.unique(nuclei_mask)\n    if 0 in nuclei_uniques:\n        nuclei_uniques = nuclei_uniques[1:]\n        \n    for unq in nuclei_uniques:\n        idxs = np.where(nuclei_mask==unq)\n        y_min, y_max = min(idxs[0]), max(idxs[0])\n        x_min, x_max = min(idxs[1]), max(idxs[1])\n        if x_min == 0 or y_min == 0 or x_max == (w - 1) or y_max == (h - 1):\n            touching_border.append(unq)\n    \n    nuclei_areas = np.array([np.count_nonzero(nuclei_mask==unq) for unq in nuclei_uniques])\n    not_touching_border_idxs = [i for i, unq in enumerate(nuclei_uniques) if unq not in touching_border]\n\n    ignore_nuclei_idxs = nuclei_uniques[(nuclei_areas < np.median(nuclei_areas[not_touching_border_idxs])*nuclei_area_thresh) & np.isin(nuclei_uniques, touching_border)]\n    \n    return ignore_nuclei_idxs\n\ndef do_segmentation(ryb_image, blue_image):\n    \n    img_size = blue_image.shape[0] * SEGMENTATION_SCALE\n\n    nuc_segmentation = segmentator.pred_nuclei([blue_image])\n    cell_segmentation = segmentator.pred_cells([ryb_image], precombined=True)\n    \n    nuclei_mask, cell_mask = label_cell_vlad(nuc_segmentation[0], cell_segmentation[0], img_size, return_nuclei_label=True)\n    \n    cell_mask = cell_mask.astype(np.uint8)\n\n    # Remove border cells with nuclei_area < median/2\n    small_border_nuclei_idxs = remove_small_nuclei_on_border(nuclei_mask)\n\n    # Remove cells without nuclei\n    cell_uniques = np.unique(cell_mask)\n    if 0 in cell_uniques:\n        cell_uniques = cell_uniques[1:]\n        \n    nuclei_uniques = np.unique(nuclei_mask)\n    if 0 in nuclei_uniques:\n        nuclei_uniques = nuclei_uniques[1:]\n            \n    cells_without_nuclei_idxs = np.setdiff1d(np.union1d(cell_uniques, nuclei_uniques), np.intersect1d(cell_uniques, nuclei_uniques))\n    \n    if len(small_border_nuclei_idxs):\n        small_border_cell_mask = np.array([cell_mask != i for i in small_border_nuclei_idxs]).prod(axis=0).astype(np.uint8)\n    else:\n        small_border_cell_mask = np.ones(cell_mask.shape, dtype=np.uint8)\n        \n    for ig in np.union1d(small_border_nuclei_idxs, cells_without_nuclei_idxs):\n        cell_mask[cell_mask == ig] = 0\n        \n    return cell_mask, small_border_cell_mask\n\ndef do_basic_segmentation(ryb_image, blue_image):\n    \n    nuc_segmentation = segmentator.pred_nuclei([blue_image])\n    cell_segmentation = segmentator.pred_cells([ryb_image], precombined=True)\n\n    nuc_segmentation[0] = cv2.resize(\n        nuc_segmentation[0], (2048, 2048), interpolation=cv2.INTER_NEAREST)\n    cell_segmentation[0] = cv2.resize(\n        cell_segmentation[0], (2048, 2048), interpolation=cv2.INTER_NEAREST)\n    \n    nuclei_mask, cell_mask = label_cell(nuc_segmentation[0], cell_segmentation[0])\n    \n    cell_mask = cell_mask.astype(np.uint8)\n        \n    return cell_mask\n\ndef get_bboxes(cell_mask, scale_factor):\n    \n    cell_bboxes = {}\n\n    unqs = np.unique(cell_mask)\n    if 0 in unqs:\n        unqs = unqs[1:]\n    \n    bboxes = get_indices_sparse(cell_mask)\n    \n    for c in unqs:\n        w, h = bboxes[c]\n        x_0, x_1, y_0, y_1 = w.min(), w.max(), h.min(), h.max()\n        \n        bbox = [int(scale_factor*x_0), int((x_1+1)*scale_factor), int(y_0*scale_factor), int((y_1+1)*scale_factor)]\n        cell_bboxes[c] = bbox\n    \n    return cell_bboxes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:09.917384Z","iopub.execute_input":"2026-03-15T14:49:09.917727Z","iopub.status.idle":"2026-03-15T14:49:09.932343Z","shell.execute_reply.started":"2026-03-15T14:49:09.917681Z","shell.execute_reply":"2026-03-15T14:49:09.931546Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"directory = '/kaggle/input/hpa-single-cell-image-classification'\ntest_df = pd.read_csv(\"/kaggle/input/hpa-single-cell-image-classification/sample_submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:09.933499Z","iopub.execute_input":"2026-03-15T14:49:09.933776Z","iopub.status.idle":"2026-03-15T14:49:09.957158Z","shell.execute_reply.started":"2026-03-15T14:49:09.933751Z","shell.execute_reply":"2026-03-15T14:49:09.956587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"idx = 0\n\nimage_id, ImageWidth, ImageHeight, PredictionString = test_df.iloc[idx]\n\n# Load images\nred_image, green_image, blue_image, yellow_image = load_RGBY_images(f\"{directory}/test/{image_id}\")\n\nimage_size = red_image.shape[0]\n# print(\"Image size:\", image_size)\n\n# Segmentation\nryb_image = np.transpose(np.array([red_image, yellow_image, blue_image]), (1,2,0))\n\ncell_mask, small_border_cell_mask = do_segmentation(ryb_image, blue_image)\nscale_factor = 1/SEGMENTATION_SCALE\n\n# cell_mask = do_basic_segmentation(ryb_image, blue_image)\n# scale_factor = 1\n\n# print(\"Cell mask size:\", cell_mask.shape)\n# print(\"Small border size:\", small_border_cell_mask.shape[0])\n\n# Get bboxes\ncell_bboxes = get_bboxes(cell_mask, scale_factor)\n\n# RLE\ncell_mask = cv2.resize(cell_mask, (image_size, image_size), interpolation=cv2.INTER_NEAREST)\n# small_border_cell_mask = cv2.resize(small_border_cell_mask, (image_size, image_size), interpolation=cv2.INTER_NEAREST)\nrles = [encode_binary_mask(cell_mask==cell_id) for cell_id in cell_bboxes.keys()]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:09.958043Z","iopub.execute_input":"2026-03-15T14:49:09.958269Z","iopub.status.idle":"2026-03-15T14:49:13.890558Z","shell.execute_reply.started":"2026-03-15T14:49:09.958247Z","shell.execute_reply":"2026-03-15T14:49:13.889874Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nimport numpy as np\n\n\ndef visualize_masks_and_bboxes(\n    image,\n    cell_mask,\n    cell_bboxes,\n    max_cells=10\n):\n    \"\"\"\n    image: HxWx3\n    cell_mask: HxW (labels inteiros)\n    cell_bboxes: dict {cell_id: [y_min, y_max, x_min, x_max]}\n    \"\"\"\n\n    fig, axes = plt.subplots(2, 2, figsize=(14, 14))\n    ax_img, ax_mask, ax_bbox, ax_all = axes.flatten()\n\n    # -------------------------------------------------\n    # 1) Imagem original\n    # -------------------------------------------------\n    ax_img.imshow(image)\n    ax_img.set_title(\"Imagem original\")\n    ax_img.axis(\"off\")\n\n    # -------------------------------------------------\n    # 2) Máscaras\n    # -------------------------------------------------\n    ax_mask.imshow(cell_mask, cmap=\"tab20\")\n    ax_mask.set_title(\"Máscaras (labels)\")\n    ax_mask.axis(\"off\")\n\n    # -------------------------------------------------\n    # 3) Bounding boxes\n    # -------------------------------------------------\n    ax_bbox.imshow(image)\n    ax_bbox.set_title(\"Bounding boxes\")\n    ax_bbox.axis(\"off\")\n\n    # -------------------------------------------------\n    # 4) Tudo junto\n    # -------------------------------------------------\n    ax_all.imshow(image)\n    ax_all.set_title(\"Imagem + máscara + bbox\")\n    ax_all.axis(\"off\")\n\n    shown = 0\n\n    for cell_id, bbox in cell_bboxes.items():\n        if shown >= max_cells:\n            break\n\n        # bbox no formato correto\n        y_min, y_max, x_min, x_max = bbox\n        w = x_max - x_min\n        h = y_max - y_min\n\n        # máscara da célula\n        mask = (cell_mask == cell_id)\n\n        # ---------- overlay máscara (painel 4)\n        ax_all.imshow(\n            np.ma.masked_where(~mask, mask),\n            alpha=0.4\n        )\n\n        # ---------- bbox (painel 3 e 4)\n        for ax in [ax_bbox, ax_all]:\n            rect = patches.Rectangle(\n                (x_min, y_min),\n                w,\n                h,\n                linewidth=2,\n                edgecolor=\"red\",\n                facecolor=\"none\"\n            )\n            ax.add_patch(rect)\n\n            ax.text(\n                x_min,\n                y_min - 4,\n                f\"id={cell_id}\\n{w}x{h}\",\n                color=\"yellow\",\n                fontsize=8,\n                bbox=dict(facecolor=\"black\", alpha=0.6, pad=1)\n            )\n\n        shown += 1\n\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:13.891608Z","iopub.execute_input":"2026-03-15T14:49:13.891855Z","iopub.status.idle":"2026-03-15T14:49:13.900515Z","shell.execute_reply.started":"2026-03-15T14:49:13.891831Z","shell.execute_reply":"2026-03-15T14:49:13.899602Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_masks_and_bboxes(\n    image=ryb_image,\n    cell_mask=cell_mask,\n    cell_bboxes=cell_bboxes,\n    max_cells=30\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:13.901537Z","iopub.execute_input":"2026-03-15T14:49:13.901766Z","iopub.status.idle":"2026-03-15T14:49:21.959865Z","shell.execute_reply.started":"2026-03-15T14:49:13.901743Z","shell.execute_reply":"2026-03-15T14:49:21.958869Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# public_test = pd.read_csv('/kaggle/input/hpa-best-segmentation-take-use/best_segmentation_result.csv')\n# public_test['ID'].values","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:21.960899Z","iopub.execute_input":"2026-03-15T14:49:21.961122Z","iopub.status.idle":"2026-03-15T14:49:21.963971Z","shell.execute_reply.started":"2026-03-15T14:49:21.961099Z","shell.execute_reply":"2026-03-15T14:49:21.963189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# rows = []\n\n# for idx in tqdm.tqdm(range(len(test_df)), desc=\"Segmentando imagens\"):\n#     image_id, ImageWidth, ImageHeight, PredictionString = test_df.iloc[idx]\n    \n#     # if image_id not in public_test['ID'].values:\n#     #     continue\n\n#     # Load images\n#     red_image, green_image, blue_image, yellow_image = load_RGBY_images(\n#         f\"{directory}/test/{image_id}\"\n#     )\n\n#     image_size = red_image.shape[0]\n#     scale_factor = 1 / SEGMENTATION_SCALE\n\n#     # Segmentation\n#     ryb_image = np.transpose(\n#         np.array([red_image, yellow_image, blue_image]), (1, 2, 0)\n#     )\n#     # cell_mask, small_border_cell_mask = do_segmentation(\n#     #     ryb_image, blue_image\n#     # )\n#     cell_mask = do_basic_segmentation(\n#         ryb_image, blue_image\n#     )\n\n#     # Get bounding boxes\n#     scale_factor = 1\n#     cell_bboxes = get_bboxes(cell_mask, scale_factor)\n\n#     # Resize masks back to original size\n#     cell_mask = cv2.resize(\n#         cell_mask, (image_size, image_size), interpolation=cv2.INTER_NEAREST\n#     )\n#     # small_border_cell_mask = cv2.resize(\n#     #     small_border_cell_mask, (image_size, image_size), interpolation=cv2.INTER_NEAREST\n#     # )\n\n#     # RLE per cell\n#     for cell_id in cell_bboxes.keys():\n#         rle = encode_binary_mask(cell_mask == cell_id)\n\n#         rows.append({\n#             \"image_id\": image_id,\n#             \"cell_id\": cell_id,\n#             \"enc\": rle,\n#             \"fname\": f\"{image_id}_{cell_id}\",\n#         })","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:21.965107Z","iopub.execute_input":"2026-03-15T14:49:21.965374Z","iopub.status.idle":"2026-03-15T14:49:21.973585Z","shell.execute_reply.started":"2026-03-15T14:49:21.965319Z","shell.execute_reply":"2026-03-15T14:49:21.972867Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# cells_df = pd.DataFrame(rows)\n# tm = cells_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:21.974667Z","iopub.execute_input":"2026-03-15T14:49:21.974901Z","iopub.status.idle":"2026-03-15T14:49:21.985814Z","shell.execute_reply.started":"2026-03-15T14:49:21.974881Z","shell.execute_reply":"2026-03-15T14:49:21.985078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# tm\n# tm.to_csv('base_segmentation.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:21.986773Z","iopub.execute_input":"2026-03-15T14:49:21.986989Z","iopub.status.idle":"2026-03-15T14:49:21.994592Z","shell.execute_reply.started":"2026-03-15T14:49:21.986969Z","shell.execute_reply":"2026-03-15T14:49:21.993837Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_cell_mask_from_df(cell_df, image_id, H, W):\n    \"\"\"\n    cell_df: dataframe com colunas [image_id, cell_id, enc]\n    H, W: tamanho da imagem\n    \"\"\"\n    df_img = cell_df[cell_df[\"image_id\"] == image_id]\n\n    masks = []\n    cell_ids = []\n\n    for _, row in df_img.iterrows():\n        rle = {\n            \"size\": [H, W],\n            \"counts\": zlib.decompress(base64.b64decode(row[\"enc\"]))\n        }\n        masks.append(rle)\n        cell_ids.append(row[\"cell_id\"])\n\n    decoded = coco_mask.decode(masks)  # H x W x N\n\n    # máscara final labelada\n    cell_mask = np.zeros((H, W), dtype=np.int32)\n    for i, cid in enumerate(cell_ids):\n        cell_mask[decoded[..., i] == 1] = cid\n\n    return cell_mask\n\ndef load_rgb_image(directory, image_id):\n    r, g, b, y = load_RGBY_images(f\"{directory}/test/{image_id}\")\n    rgb = np.stack([r, g, b], axis=-1)\n    return rgb\n\ndef plot_image_and_masks(rgb_image, cell_mask):\n    fig, axes = plt.subplots(1, 2, figsize=(14, 6))\n\n    # Imagem\n    axes[0].imshow(rgb_image)\n    axes[0].set_title(\"Imagem\")\n    axes[0].axis(\"off\")\n\n    # Máscaras\n    axes[1].imshow(cell_mask, cmap=\"tab20\")\n    axes[1].set_title(\"Máscaras das células\")\n    axes[1].axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:21.995591Z","iopub.execute_input":"2026-03-15T14:49:21.995817Z","iopub.status.idle":"2026-03-15T14:49:22.004303Z","shell.execute_reply.started":"2026-03-15T14:49:21.995784Z","shell.execute_reply":"2026-03-15T14:49:22.003613Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# image_id = '0040581b-f1f2-4fbe-b043-b6bfea5404bb'\n\n# # carregar imagem\n# rgb_image = load_rgb_image(directory, image_id)\n# H, W = rgb_image.shape[:2]\n\n# # montar máscara\n# cell_mask = build_cell_mask_from_df(tm, image_id, H, W)\n\n# # plotar\n# plot_image_and_masks(rgb_image, cell_mask)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T14:49:22.005208Z","iopub.execute_input":"2026-03-15T14:49:22.005494Z","iopub.status.idle":"2026-03-15T14:49:22.015132Z","shell.execute_reply.started":"2026-03-15T14:49:22.00547Z","shell.execute_reply":"2026-03-15T14:49:22.014594Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## If we read from a csv","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/hpa-best-segmentation-take-use/best_segmentation_result.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:32:35.055282Z","iopub.execute_input":"2026-03-15T15:32:35.05563Z","iopub.status.idle":"2026-03-15T15:32:35.665301Z","shell.execute_reply.started":"2026-03-15T15:32:35.055596Z","shell.execute_reply":"2026-03-15T15:32:35.664647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Lista para armazenar os dicionários com informações de cada célula\nimgs = []\n\n# Itera sobre todas as linhas do DataFrame df\nfor i, x in df.iterrows():\n    \n    # Extrai os labels da string PredictionString (um a cada 3 itens)\n    label = x.PredictionString.split(' ')[0::3]\n    \n    # Extrai as probabilidades associadas aos labels (um a cada 3 itens)\n    prob = x.PredictionString.split(' ')[1::3]\n    \n    # Extrai as máscaras codificadas em RLE (um a cada 3 itens)\n    encodes = x.PredictionString.split(' ')[2::3]\n    \n    # Itera sobre os RLEs únicos (remove duplicatas com set)\n    for idx, enc in enumerate(list(set(encodes))):\n        imgs.append({\n            'image_id': x.ID,            # ID da imagem original\n            'cell_id': idx + 1,          # Índice da célula (começa do 1)\n            'enc': enc,                  # Máscara codificada em RLE\n            'fname': f'{x.ID}_{idx+1}',  # Nome do arquivo no formato <image_id>_<cell_id>\n        })\n\n# Converte a lista de dicionários em um novo DataFrame\ntm = pd.DataFrame(imgs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:32:35.666491Z","iopub.execute_input":"2026-03-15T15:32:35.666706Z","iopub.status.idle":"2026-03-15T15:32:46.087443Z","shell.execute_reply.started":"2026-03-15T15:32:35.666684Z","shell.execute_reply":"2026-03-15T15:32:46.086742Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:32:46.088875Z","iopub.execute_input":"2026-03-15T15:32:46.089095Z","iopub.status.idle":"2026-03-15T15:32:46.103667Z","shell.execute_reply.started":"2026-03-15T15:32:46.089073Z","shell.execute_reply":"2026-03-15T15:32:46.10276Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = df.drop(columns='PredictionString')\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:32:46.104923Z","iopub.execute_input":"2026-03-15T15:32:46.105189Z","iopub.status.idle":"2026-03-15T15:32:46.117524Z","shell.execute_reply.started":"2026-03-15T15:32:46.105163Z","shell.execute_reply":"2026-03-15T15:32:46.116816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_submission = pd.read_csv('../input/hpa-single-cell-image-classification/sample_submission.csv', index_col=0)\nsample_submission.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:32:46.118633Z","iopub.execute_input":"2026-03-15T15:32:46.118963Z","iopub.status.idle":"2026-03-15T15:32:46.136611Z","shell.execute_reply.started":"2026-03-15T15:32:46.118929Z","shell.execute_reply":"2026-03-15T15:32:46.135899Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = sample_submission.reset_index()[['ID', 'ImageWidth', 'ImageHeight']]\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:32:46.137559Z","iopub.execute_input":"2026-03-15T15:32:46.137775Z","iopub.status.idle":"2026-03-15T15:32:46.147662Z","shell.execute_reply.started":"2026-03-15T15:32:46.137753Z","shell.execute_reply":"2026-03-15T15:32:46.146906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def apply_mask(cell_img, cell_mask, use_mask=True):\n    if not use_mask:\n        return cell_img\n    mask_4ch = np.repeat(cell_mask[..., None], cell_img.shape[2], axis=2)\n    return cell_img * mask_4ch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:32:46.148697Z","iopub.execute_input":"2026-03-15T15:32:46.149022Z","iopub.status.idle":"2026-03-15T15:32:46.155224Z","shell.execute_reply.started":"2026-03-15T15:32:46.14899Z","shell.execute_reply":"2026-03-15T15:32:46.154656Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def place_center(canvas_size, cell_img):\n    h, w = cell_img.shape[:2]\n    canvas = np.zeros((canvas_size, canvas_size, cell_img.shape[2]), dtype=cell_img.dtype)\n\n    y0 = (canvas_size - h) // 2\n    x0 = (canvas_size - w) // 2\n\n    canvas[y0:y0+h, x0:x0+w] = cell_img\n    return canvas","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:32:46.156786Z","iopub.execute_input":"2026-03-15T15:32:46.157045Z","iopub.status.idle":"2026-03-15T15:32:46.1656Z","shell.execute_reply.started":"2026-03-15T15:32:46.157007Z","shell.execute_reply":"2026-03-15T15:32:46.16499Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_bag(cell_img, canvas_size=512):\n    \"\"\"\n    Repete a célula até preencher completamente um canvas canvas_size x canvas_size.\n    \"\"\"\n    h, w, c = cell_img.shape\n\n    # Quantas repetições são necessárias em cada eixo\n    rep_y = int(np.ceil(canvas_size / h))\n    rep_x = int(np.ceil(canvas_size / w))\n\n    # Repete a célula\n    tiled = np.tile(cell_img, (rep_y, rep_x, 1))\n\n    # Corta para o tamanho exato do canvas\n    canvas = tiled[:canvas_size, :canvas_size, :]\n\n    return canvas","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:32:46.166897Z","iopub.execute_input":"2026-03-15T15:32:46.167257Z","iopub.status.idle":"2026-03-15T15:32:46.174831Z","shell.execute_reply.started":"2026-03-15T15:32:46.167223Z","shell.execute_reply":"2026-03-15T15:32:46.17405Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SliceInferenceDataset(torch.utils.data.Dataset):\n    def __init__(\n        self,\n        df,\n        tta=1,\n        bag_of_cells=False,\n        centralize_cell=True,\n        mask=True,\n        cfg=None,\n        tfms=None\n    ):\n        self.df = df\n        self.iids = df.image_id.unique()\n        self.tta = tta\n        self.bag_of_cells = bag_of_cells\n        self.centralize_cell = centralize_cell\n        self.mask = mask                           \n        \n    def __len__(self):\n        return len(self.iids)\n\n    def __getitem__(self, idx):\n        iid = self.iids[idx]\n\n        mt = f'../input/hpa-single-cell-image-classification/test/{iid}_red.png'\n        er = f'../input/hpa-single-cell-image-classification/test/{iid}_yellow.png'\n        nu = f'../input/hpa-single-cell-image-classification/test/{iid}_blue.png'\n        pr = f'../input/hpa-single-cell-image-classification/test/{iid}_green.png'\n\n        r = cv2.imread(mt, 0).astype(np.float32) / 255.0\n        g = cv2.imread(pr, 0).astype(np.float32) / 255.0\n        b = cv2.imread(nu, 0).astype(np.float32) / 255.0\n        a = cv2.imread(er, 0).astype(np.float32) / 255.0\n        \n        sz = r.shape[0]\n        img = np.stack([r, g, b, a], -1)\n\n        img_resized = cv2.resize(img, (512, 512), interpolation=cv2.INTER_LINEAR)\n\n        cells = []\n\n        for i, x in self.df[self.df.image_id == iid].iterrows():\n\n            bd = base64.b64decode(x.enc)\n            zd = zlib.decompress(bd)\n            encoded = [{'counts': zd, 'size': (sz, sz)}]\n            mask_bin = coco_mask.decode(encoded)[:, :, 0]\n\n            if len(np.unique(mask_bin)) == 1:\n                continue\n\n            mask_bin_resized = cv2.resize(mask_bin, (512, 512), interpolation=cv2.INTER_NEAREST)\n\n            ys, xs = np.where(mask_bin_resized == 1)\n\n            cell = img_resized[ys.min():ys.max(), xs.min():xs.max()]\n            cell_mask = mask_bin_resized[ys.min():ys.max(), xs.min():xs.max()]\n\n            cell = apply_mask(cell, cell_mask, self.mask)\n\n            if self.centralize_cell:\n                cell = place_center(512, cell)\n\n            if self.bag_of_cells:\n                cell = make_bag(cell, 512)\n\n            cells.append((cell.astype(np.float32), x.fname))\n\n        if not cells:\n            return None\n        \n        BS, tta = len(cells) + 1, self.tta\n\n        raw_ipt = [e[0] for e in cells]\n\n        # ---------- TTA transforms ----------\n        if tta == 1:\n            transforms = [(False, 0)]\n        elif tta == 8:\n            transforms = [(flip, k) for flip in [False, True] for k in range(4)]\n        else:\n            raise ValueError(\"TTA configuration not supported\")\n\n        # ---------- Cell TTA ----------\n        ipts = []\n\n        for flip, k in transforms:\n\n            batch = []\n\n            for x in raw_ipt:\n\n                xr = x\n\n                if flip:\n                    xr = np.flip(xr, axis=1)\n\n                xr = np.rot90(xr, k).copy()\n                xr = res_tfms(image=xr)['image']\n                xr = tensor_tfms(xr)\n\n                batch.append(xr)\n\n            ipts.append(torch.stack(batch).float())\n\n        # ---------- Full image TTA ----------\n        image_tta = []\n\n        for flip, k in transforms:\n\n            xr = img\n\n            if flip:\n                xr = np.flip(xr, axis=1)\n\n            xr = np.rot90(xr, k).copy()\n            xr = image_tfms(image_res_tfms(image=xr)['image'])\n\n            image_tta.append(xr)\n\n        return image_tta, ipts, BS, len(cells), tta, iid, [x[1] for x in cells]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:32:46.175992Z","iopub.execute_input":"2026-03-15T15:32:46.176247Z","iopub.status.idle":"2026-03-15T15:32:46.193098Z","shell.execute_reply.started":"2026-03-15T15:32:46.176225Z","shell.execute_reply":"2026-03-15T15:32:46.19237Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_worker(worker_id):\n    worker_seed = torch.initial_seed() % 2**32\n    np.random.seed(worker_seed)\n    random.seed(worker_seed)\n\nsid = SliceInferenceDataset(tm, \n                            tta=1,\n                            bag_of_cells=True,\n                            centralize_cell=False,\n                            mask=True\n                           )\ndl = torch.utils.data.DataLoader(sid, batch_size=1, num_workers=2, worker_init_fn=seed_worker)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:32:46.194052Z","iopub.execute_input":"2026-03-15T15:32:46.194313Z","iopub.status.idle":"2026-03-15T15:32:46.206385Z","shell.execute_reply.started":"2026-03-15T15:32:46.194264Z","shell.execute_reply":"2026-03-15T15:32:46.205643Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\n\n# Move os modelos para GPU e define modo de avaliação (desativa dropout, batchnorm)\ncell_models = [model.cuda().eval() for model in models]\nimage_models = [model.cuda().eval() for model in image_models]\n\n# Listas para armazenar predições\ncell_predictions_dfs = []        # Para armazenar predições por célula\nimage_predictions_dfs = []       # Para armazenar predições por imagem (whole image prediction)\n\n# Loop sobre o DataLoader de inferência (dl)\ncontador = 0\nfor image_tta, cell_inputs_tta, batch_size_tensor, num_cells_tensor, tta_tensor, image_id, cell_filenames_raw in tqdm.tqdm(dl):\n    # Extração e formatação das variáveis\n    batch_size = batch_size_tensor.item()           # Tamanho do batch\n    num_tta = tta_tensor.item()                     # Número de Test Time Augmentations\n    image_id = image_id[0]                          # image_id (string)\n    cell_filenames = [name[0] for name in cell_filenames_raw]  # Nomes dos arquivos das células\n    num_cells = num_cells_tensor.item()             # Número de células na imagem\n    \n    cell_predictions_batches = []  # Lista para armazenar as predições de células\n    image_predictions_batches = [] # Lista para armazenar as predições da imagem inteira\n\n    # Processamento das células em blocos do tamanho do batch\n    for i in range(0, num_cells, batch_size):\n        with torch.no_grad():\n            res = []  # Predições por célula (para cada modelo com TTA)\n            exp = []  # Predições da imagem inteira (para cada modelo)\n            \n            # Inferência da imagem completa com modelos de imagem (ex: Puzzle-CAM)\n            imgs = torch.stack(\n                [image_tta[tt][0] for tt in range(num_tta)]\n            ).cuda().float()\n            \n            for model in image_models:\n                ifr = model(imgs)\n                exp.append(ifr.float())\n\n            # Inferência por TTA: aplica vários modelos às células com augmentations\n            for tta_idx in range(num_tta):\n                cell_batch = cell_inputs_tta[tta_idx][0].cuda()  # Recupera as células para essa TTA e move para GPU\n                for model in image_models: # Mantido conforme sua lógica original (usando image_models para células)\n                    cell_batch = cell_batch.cuda().float()\n                    ifr = model(cell_batch)\n                    res.append(ifr.float())  # Coleta saída da célula\n\n        # Média das predições de cada célula com sigmoid (multi-label)\n        cell_prob = [torch.sigmoid(r.cpu()) for r in res]\n        image_prob = [torch.sigmoid(r.cpu()) for r in exp]\n        \n        cell_probs_mean = np.stack(cell_prob).mean(0)  # Média ao longo de modelos/TTA\n        image_probs_mean = np.stack(image_prob).mean(0)  # Média ao longo de modelos/TTA\n\n        cell_predictions_batches.append(cell_probs_mean)  # Armazena predições das células\n        image_predictions_batches.append(image_probs_mean)  # Armazena predições da imagem\n\n    # Concatena predições por célula (final do loop da imagem)\n    cell_predictions = np.concatenate(cell_predictions_batches)  # Shape: (n_cells, n_classes)\n    cell_df = pd.DataFrame(cell_predictions, index=cell_filenames)  # DataFrame com scores para cada célula\n    \n    image_df = pd.DataFrame(\n        np.concatenate(image_predictions_batches).mean(0).reshape(1, 19), index=[image_id]\n    )  # DataFrame com score médio da imagem inteira\n\n    # Salva resultados\n    image_predictions_dfs.append(image_df)  # Uma linha por imagem\n    cell_predictions_dfs.append(cell_df)    # Uma linha por célula\n\n    del image_tta, cell_inputs_tta\n    gc.collect()\n    torch.cuda.empty_cache()\n\n    contador += 1\n    if contador % 10 == 0:\n        print_memory_usage()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:34:13.638887Z","iopub.execute_input":"2026-03-15T15:34:13.639211Z","iopub.status.idle":"2026-03-15T15:35:06.996577Z","shell.execute_reply.started":"2026-03-15T15:34:13.639178Z","shell.execute_reply":"2026-03-15T15:35:06.995382Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Concatena todos os DataFrames de predição por imagem (gerados durante a inferência)\n# Isso resulta em um DataFrame onde cada linha representa uma imagem (image_id) com predições para cada classe\nimage_level = pd.concat(image_predictions_dfs)\n\n# Junta o DataFrame `image_level` com o DataFrame `tm`, para associar os image_id com seus nomes de arquivo (fname)\n# - `reset_index()` move o índice (que contém image_id) para uma coluna chamada 'index'\n# - `merge(..., how='left')` garante que todos os dados de `image_level` sejam mantidos\n# - `on='index'` (de `image_level`) se conecta com `image_id` (de `tm`)\nimage_pred = image_level.reset_index().merge(\n    tm[['image_id', 'fname']],  # Subconjunto com apenas as colunas relevantes do `tm`\n    left_on='index',            # Usa o antigo índice como chave de junção\n    right_on='image_id',        # Faz a junção com a coluna image_id do `tm`\n    how='left'                  # Preserva todas as linhas de `image_level`\n)\n# Organiza o DataFrame final:\n# - Define 'fname' como índice, o que facilita o acesso por nome de célula depois\n# - Remove as colunas 'index' (image_id original do image_level) e 'image_id' (duplicado do tm)\nimage_pred = image_pred.set_index('fname').drop(['index', 'image_id'], axis=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:26:07.855193Z","iopub.execute_input":"2026-03-15T15:26:07.855546Z","iopub.status.idle":"2026-03-15T15:26:07.883124Z","shell.execute_reply.started":"2026-03-15T15:26:07.855512Z","shell.execute_reply":"2026-03-15T15:26:07.88247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Concatena todos os DataFrames de predições por célula (armazenados em `pdfs`)\n# Cada DataFrame em `pdfs` contém predições de classes para as células daquela imagem, indexadas por fname (ex: \"abc_1\")\npub_pred = pd.concat(cell_predictions_dfs)\n\n# Multiplica as predições por célula (`pub_pred`) pelas predições a nível de imagem (`image_pred`)\n# - Ambas têm o mesmo índice (fname, como \"abc_1\"), então a operação é feita elemento a elemento\n# - Isso pode ser interpretado como um reponderamento das predições da célula com base na imagem\n#   (por exemplo, se a imagem não tiver forte ativação para uma classe, a predição da célula será atenuada)\n\n# merge_pred = pub_pred + image_pred.loc[pub_pred.index]\nmerge_pred = pub_pred * image_pred.loc[pub_pred.index]\n# merge_pred = image_pred.loc[pub_pred.index]\n# merge_pred = pub_pred\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:26:08.738543Z","iopub.execute_input":"2026-03-15T15:26:08.73884Z","iopub.status.idle":"2026-03-15T15:26:08.778301Z","shell.execute_reply.started":"2026-03-15T15:26:08.738813Z","shell.execute_reply":"2026-03-15T15:26:08.777527Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## If any ensemble","metadata":{}},{"cell_type":"code","source":"merge_pred.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:26:10.307009Z","iopub.execute_input":"2026-03-15T15:26:10.307311Z","iopub.status.idle":"2026-03-15T15:26:10.326732Z","shell.execute_reply.started":"2026-03-15T15:26:10.307282Z","shell.execute_reply":"2026-03-15T15:26:10.325802Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Save prediction","metadata":{}},{"cell_type":"code","source":"df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:26:11.113742Z","iopub.execute_input":"2026-03-15T15:26:11.114054Z","iopub.status.idle":"2026-03-15T15:26:11.125006Z","shell.execute_reply.started":"2026-03-15T15:26:11.114014Z","shell.execute_reply":"2026-03-15T15:26:11.124286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = df.set_index('ID') # Define a coluna 'ID' como índice do DataFrame `df`\nmerge_pred.index.name = 'fname' # Define o nome do índice de `merge_pred` como 'fname'\nmerge_pred = merge_pred.reset_index() # Reseta o índice de `merge_pred`, transformando 'fname' em uma coluna normal\ntm = tm.set_index('fname') # Define a coluna 'fname' como índice do DataFrame `tm`\n\n# Cria uma nova coluna 'ID' em `merge_pred`, extraindo a parte antes do \"_\" da coluna 'fname'\n# Exemplo: se fname = \"abc_3\", então ID = \"abc\"\nmerge_pred['ID'] = merge_pred['fname'].str.split('_', expand=True)[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:26:12.729428Z","iopub.execute_input":"2026-03-15T15:26:12.729722Z","iopub.status.idle":"2026-03-15T15:26:12.741588Z","shell.execute_reply.started":"2026-03-15T15:26:12.729697Z","shell.execute_reply":"2026-03-15T15:26:12.740856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Lista para armazenar as predições finais formatadas por imagem\nj_pred = []\n\n# Itera sobre todos os IDs únicos (ou seja, imagens únicas) presentes no DataFrame `merge_pred`\nfor iid in merge_pred.ID.unique():\n    enc = ''  # Inicializa a string de predição para essa imagem\n\n    # Filtra `merge_pred` para obter apenas as linhas (células) associadas ao ID atual (imagem atual)\n    sub_df = merge_pred[merge_pred.ID == iid]\n\n    # Para cada célula da imagem\n    for idx, row in sub_df.iterrows():\n        # Para cada uma das 19 classes (0 a 18), cria um trecho do tipo: \"class prob mask\"\n        for i in range(19):\n            enc += f'{i} {row[i]} {tm.loc[row.fname].enc} '\n    \n    # Adiciona a predição formatada como um dicionário na lista `j_pred`\n    # Remove o último espaço de `enc` com `[:-1]`\n    j_pred.append({\n        'ID': iid,  # ID da imagem\n        'ImageWidth': df.loc[iid].ImageWidth,   # Largura da imagem (obtida do DataFrame original `df`)\n        'ImageHeight': df.loc[iid].ImageHeight, # Altura da imagem\n        'PredictionString': enc[:-1]  # String com todas as predições no formato esperado para submissão\n    })\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:26:12.973527Z","iopub.execute_input":"2026-03-15T15:26:12.973781Z","iopub.status.idle":"2026-03-15T15:26:13.671024Z","shell.execute_reply.started":"2026-03-15T15:26:12.973757Z","shell.execute_reply":"2026-03-15T15:26:13.670341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Converte a lista de dicionários `j_pred` em um DataFrame do pandas\nfast_sub = pd.DataFrame(j_pred)\n\n# Salva esse DataFrame em um arquivo CSV chamado 'pub.csv'\n# Esse arquivo será formatado como exigido para submissão na competição (ID, ImageWidth, ImageHeight, PredictionString)\nfast_sub.to_csv('pub.csv', index=False)\n\n# Define a coluna 'ID' como índice do DataFrame (útil para buscas e manipulações futuras)\nfast_sub = fast_sub.set_index('ID')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:26:15.003867Z","iopub.execute_input":"2026-03-15T15:26:15.004144Z","iopub.status.idle":"2026-03-15T15:26:15.197903Z","shell.execute_reply.started":"2026-03-15T15:26:15.004119Z","shell.execute_reply":"2026-03-15T15:26:15.197231Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fast_sub.head(2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:26:15.414684Z","iopub.execute_input":"2026-03-15T15:26:15.414936Z","iopub.status.idle":"2026-03-15T15:26:15.423682Z","shell.execute_reply.started":"2026-03-15T15:26:15.414911Z","shell.execute_reply":"2026-03-15T15:26:15.42292Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## save","metadata":{}},{"cell_type":"code","source":"# Substitui as previsões do DataFrame `sample_submission` pelas do `fast_sub`\n# drop(fast_sub.index): remove as linhas de 'sample_submission' que já estão em 'fast_sub' (com mesmo ID)\n# depois concatena com 'fast_sub' para manter as novas previsões\nsub2 = pd.concat([sample_submission.drop(fast_sub.index), fast_sub], axis=0)\n\n# Reordena as linhas para seguir exatamente a ordem original de 'sample_submission'\nsub2 = sub2.loc[sample_submission.index]\n\n# Salva o novo DataFrame como CSV no formato de submissão exigido\nsub2.to_csv('submission.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T15:26:17.515052Z","iopub.execute_input":"2026-03-15T15:26:17.515386Z","iopub.status.idle":"2026-03-15T15:26:17.558005Z","shell.execute_reply.started":"2026-03-15T15:26:17.515352Z","shell.execute_reply":"2026-03-15T15:26:17.557234Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}