{"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":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":"datasetVersion","sourceId":14478588,"datasetId":9247771,"databundleVersionId":15301743},{"sourceType":"modelInstanceVersion","sourceId":781441,"databundleVersionId":16012243,"modelInstanceId":596214,"modelId":608472},{"sourceType":"modelInstanceVersion","sourceId":364551,"databundleVersionId":12078278,"modelInstanceId":302470,"modelId":322964},{"sourceType":"modelInstanceVersion","sourceId":779540,"databundleVersionId":15992124,"modelInstanceId":594858,"modelId":607104},{"sourceType":"modelInstanceVersion","sourceId":778278,"databundleVersionId":15978839,"modelInstanceId":593914,"modelId":606182},{"sourceType":"modelInstanceVersion","sourceId":778577,"databundleVersionId":15982237,"modelInstanceId":594120,"modelId":606392},{"sourceType":"modelInstanceVersion","sourceId":781276,"databundleVersionId":16010476,"modelInstanceId":596086,"modelId":608339},{"sourceType":"modelInstanceVersion","sourceId":632149,"databundleVersionId":14373708,"modelInstanceId":476494,"modelId":492416},{"sourceType":"modelInstanceVersion","sourceId":736178,"databundleVersionId":15519459,"modelInstanceId":561270,"modelId":573890},{"sourceType":"modelInstanceVersion","sourceId":737101,"databundleVersionId":15532506,"modelInstanceId":562053,"modelId":574676},{"sourceType":"modelInstanceVersion","sourceId":737119,"databundleVersionId":15532869,"modelInstanceId":562069,"modelId":574690},{"sourceType":"modelInstanceVersion","sourceId":735302,"databundleVersionId":15506162,"modelInstanceId":560531,"modelId":573116},{"sourceType":"modelInstanceVersion","sourceId":442288,"databundleVersionId":12763968,"modelInstanceId":359282,"modelId":380590},{"sourceType":"modelInstanceVersion","sourceId":757661,"databundleVersionId":15777486,"modelInstanceId":578741,"modelId":591076},{"sourceType":"modelInstanceVersion","sourceId":449467,"databundleVersionId":12836461,"modelInstanceId":364886,"modelId":385770},{"sourceType":"modelInstanceVersion","sourceId":450264,"databundleVersionId":12844752,"modelInstanceId":365450,"modelId":386329},{"sourceType":"modelInstanceVersion","sourceId":446044,"databundleVersionId":12799163,"modelInstanceId":362202,"modelId":383169},{"sourceType":"modelInstanceVersion","sourceId":449386,"databundleVersionId":12835782,"modelInstanceId":364856,"modelId":385736},{"sourceType":"modelInstanceVersion","sourceId":443268,"databundleVersionId":12772589,"modelInstanceId":360032,"modelId":381190},{"sourceType":"modelInstanceVersion","sourceId":541009,"databundleVersionId":13508240,"modelInstanceId":417968,"modelId":435634},{"sourceType":"modelInstanceVersion","sourceId":735316,"databundleVersionId":15506319,"modelInstanceId":560545,"modelId":573130},{"sourceType":"modelInstanceVersion","sourceId":582826,"databundleVersionId":13768488,"modelInstanceId":435205,"modelId":452029},{"sourceType":"modelInstanceVersion","sourceId":581619,"databundleVersionId":13752384,"modelInstanceId":434168,"modelId":451031},{"sourceType":"modelInstanceVersion","sourceId":545433,"databundleVersionId":13522701,"modelInstanceId":418738,"modelId":436388},{"sourceType":"modelInstanceVersion","sourceId":582415,"databundleVersionId":13763090,"modelInstanceId":434854,"modelId":451700},{"sourceType":"modelInstanceVersion","sourceId":581617,"databundleVersionId":13752357,"modelInstanceId":434167,"modelId":451030},{"sourceType":"modelInstanceVersion","sourceId":431551,"databundleVersionId":12656772,"modelInstanceId":351788,"modelId":373054},{"sourceType":"modelInstanceVersion","sourceId":439909,"databundleVersionId":12751647,"modelInstanceId":358519,"modelId":379833},{"sourceType":"modelInstanceVersion","sourceId":447812,"databundleVersionId":12818345,"modelInstanceId":363556,"modelId":384422},{"sourceType":"modelInstanceVersion","sourceId":444859,"databundleVersionId":12787670,"modelInstanceId":361202,"modelId":382264},{"sourceType":"modelInstanceVersion","sourceId":444925,"databundleVersionId":12788066,"modelInstanceId":361243,"modelId":382303},{"sourceType":"modelInstanceVersion","sourceId":449282,"databundleVersionId":12834914,"modelInstanceId":364793,"modelId":385675},{"sourceType":"modelInstanceVersion","sourceId":642496,"databundleVersionId":14455147,"modelInstanceId":484513,"modelId":499991},{"sourceType":"modelInstanceVersion","sourceId":576312,"databundleVersionId":13706332,"modelInstanceId":431050,"modelId":447976},{"sourceType":"modelInstanceVersion","sourceId":576321,"databundleVersionId":13706386,"modelInstanceId":431053,"modelId":447979}],"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":"# !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-15T23:05:54.486133Z","iopub.execute_input":"2026-03-15T23:05:54.486395Z","iopub.status.idle":"2026-03-15T23:06:39.141321Z","shell.execute_reply.started":"2026-03-15T23:05:54.486372Z","shell.execute_reply":"2026-03-15T23:06:39.140144Z"}},"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-15T23:06:39.142635Z","iopub.execute_input":"2026-03-15T23:06:39.142926Z","iopub.status.idle":"2026-03-15T23:07:18.044813Z","shell.execute_reply.started":"2026-03-15T23:06:39.142898Z","shell.execute_reply":"2026-03-15T23:07:18.04385Z"}},"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\nfrom torchvision import transforms\n# import 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 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\n\nimport psutil\nimport os","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:07:18.062155Z","iopub.execute_input":"2026-03-15T23:07:18.062438Z","iopub.status.idle":"2026-03-15T23:07:21.472102Z","shell.execute_reply.started":"2026-03-15T23:07:18.062409Z","shell.execute_reply":"2026-03-15T23:07:21.471352Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PUBLIC_TEST = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:07:21.473061Z","iopub.execute_input":"2026-03-15T23:07:21.473309Z","iopub.status.idle":"2026-03-15T23:07:21.476652Z","shell.execute_reply.started":"2026-03-15T23:07:21.473284Z","shell.execute_reply":"2026-03-15T23:07:21.475796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def 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-15T23:07:21.478026Z","iopub.execute_input":"2026-03-15T23:07:21.478371Z","iopub.status.idle":"2026-03-15T23:07:21.487021Z","shell.execute_reply.started":"2026-03-15T23:07:21.478337Z","shell.execute_reply":"2026-03-15T23:07:21.486124Z"}},"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\ntta_tfms = A.Compose([\n    A.Resize(always_apply=False, p=1, height=256, width=256, interpolation=1),\n    A.HorizontalFlip(always_apply=False, p=0.5),\n    A.ShiftScaleRotate(always_apply=False, p=0.7, shift_limit_x=(-0.06, 0.06), shift_limit_y=(-0.06, 0.06), scale_limit=(-0.3, 0.3), rotate_limit=(-22.5, 22.5), interpolation=1, border_mode=2, value=None, mask_value=None),\n    A.RandomBrightnessContrast(always_apply=False, p=0.5, brightness_limit=(-0.2, 0.2), contrast_limit=(-0.2, 0.2), brightness_by_max=True),\n])\n\nres_tfms = A.Compose([\n    A.Resize(always_apply=False, p=1, height=256, width=256, 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-15T23:07:21.488177Z","iopub.execute_input":"2026-03-15T23:07:21.488497Z","iopub.status.idle":"2026-03-15T23:07:21.497407Z","shell.execute_reply.started":"2026-03-15T23:07:21.488464Z","shell.execute_reply":"2026-03-15T23:07:21.49654Z"}},"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-15T23:07:21.498517Z","iopub.execute_input":"2026-03-15T23:07:21.49878Z","iopub.status.idle":"2026-03-15T23:07:21.512319Z","shell.execute_reply.started":"2026-03-15T23:07:21.498752Z","shell.execute_reply":"2026-03-15T23:07:21.511585Z"}},"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-15T23:07:21.513352Z","iopub.execute_input":"2026-03-15T23:07:21.51362Z","iopub.status.idle":"2026-03-15T23:07:21.521444Z","shell.execute_reply.started":"2026-03-15T23:07:21.513595Z","shell.execute_reply":"2026-03-15T23:07:21.520734Z"}},"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-15T23:07:21.522351Z","iopub.execute_input":"2026-03-15T23:07:21.522564Z","iopub.status.idle":"2026-03-15T23:07:21.533298Z","shell.execute_reply.started":"2026-03-15T23:07:21.522542Z","shell.execute_reply":"2026-03-15T23:07:21.532489Z"}},"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.float)\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 CellClassifierV2(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    cin = self.backbone.outplanes\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    self.from_scratch_layers.extend([self.last_linear_cell, self.last_linear_image])\n    self.initialize([self.last_linear_cell, self.last_linear_image])\n\n  def forward(self, x, cnt=16, with_cam=False, cell_logits_to_image_logits=False):\n    cnt = torch.tensor(cnt)\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\n    if cell_logits_to_image_logits:\n      cell_logits = self.last_linear_cell(pooled)\n      cell_logits_split = torch.split(cell_logits, cnt.tolist())\n      image_logits = torch.stack([p.max(0).values for p in cell_logits_split])\n\n      return cell_logits, image_logits\n\n    pooled_split = torch.split(pooled, cnt.tolist())\n    pooled_per_img = torch.stack([p.max(0)[0] for p in pooled_split])\n\n    cell_logits = self.last_linear_cell(pooled)\n    image_logits = self.last_linear_image(pooled_per_img)\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-15T23:07:21.536433Z","iopub.execute_input":"2026-03-15T23:07:21.536741Z","iopub.status.idle":"2026-03-15T23:07:21.766407Z","shell.execute_reply.started":"2026-03-15T23:07:21.536676Z","shell.execute_reply":"2026-03-15T23:07:21.765521Z"},"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-15T23:07:21.768587Z","iopub.execute_input":"2026-03-15T23:07:21.768968Z","iopub.status.idle":"2026-03-15T23:07:21.780663Z","shell.execute_reply.started":"2026-03-15T23:07:21.768929Z","shell.execute_reply":"2026-03-15T23:07:21.779784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"WEIGHTS_PATHS = [\n    '/kaggle/input/models/julianamidlej/rs50-dino-20e-no-freeze/pytorch/default/1/hpa2nd-rs50-lr0.0002-b6-aug2nd-adamw-eid-down-task-dino/model-f5-e19.pth']\nDEVICE = 'cuda'\n\nfor 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-15T23:07:41.082422Z","iopub.execute_input":"2026-03-15T23:07:41.082778Z","iopub.status.idle":"2026-03-15T23:07:41.741558Z","shell.execute_reply.started":"2026-03-15T23:07:41.082714Z","shell.execute_reply":"2026-03-15T23:07:41.740672Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# WEIGHTS_PATHS = [\n#     '/kaggle/input/models/julianamidlej/base-ema-mixup-crop-res-m1/pytorch/default/1/hpa-512-rs269-lr0.1-b16-ls0.1-mixup-ema1-cp-m1.pth']\n# DEVICE = 'cuda'\n\n# for 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-15T23:07:49.015819Z","iopub.execute_input":"2026-03-15T23:07:49.016161Z","iopub.status.idle":"2026-03-15T23:07:49.019758Z","shell.execute_reply.started":"2026-03-15T23:07:49.01613Z","shell.execute_reply":"2026-03-15T23:07:49.018779Z"}},"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":"# model_paths = [\n#     '/kaggle/input/b3/pytorch/default/1/b3/checkpoints/f0_epoch-18.pth',\n#     '/kaggle/input/b3-f1/pytorch/default/1/b3_F1/checkpoints/f1_epoch-16.pth',\n#     '/kaggle/input/b3-f2/pytorch/default/1/b3_F2/checkpoints/f2_epoch-17.pth',\n#     '/kaggle/input/b3-f3/pytorch/default/1/b3_F3/checkpoints/f3_epoch-19.pth',\n#     '/kaggle/input/b3-f4/pytorch/default/1/b3_F4/checkpoints/f4_epoch-18.pth'\n# ]\n\n# for model_path in model_paths:\n#     cfg = Config.load_json('/kaggle/input/b3/pytorch/default/1/b3/config.json')\n#     model = get_model(cfg).cuda()\n#     load_matched_state(model, torch.load(model_path))\n#     _ = model.eval()\n#     models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:07:49.351522Z","iopub.execute_input":"2026-03-15T23:07:49.351827Z","iopub.status.idle":"2026-03-15T23:07:49.355128Z","shell.execute_reply.started":"2026-03-15T23:07:49.3518Z","shell.execute_reply":"2026-03-15T23:07:49.354338Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model_paths = [\n#     '/kaggle/input/b5/pytorch/default/1/b5/checkpoints/f0_epoch-14.pth',\n#     '/kaggle/input/b5-f1/pytorch/default/1/b5_F1/checkpoints/f1_epoch-14.pth',\n#     '/kaggle/input/b5-f2/pytorch/default/1/b5_F2/checkpoints/f2_epoch-15.pth',\n#     '/kaggle/input/b5-f3-v2/pytorch/default/1/b5_F3/checkpoints/f3_epoch-17.pth',\n#     '/kaggle/input/b5-f4/pytorch/default/1/b5_F4/checkpoints/f4_epoch-14.pth'   \n# ]\n\n# for model_path in model_paths:\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(model_path))\n#     _ = model.eval()\n#     models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:07:49.577619Z","iopub.execute_input":"2026-03-15T23:07:49.577983Z","iopub.status.idle":"2026-03-15T23:07:49.581385Z","shell.execute_reply.started":"2026-03-15T23:07:49.577951Z","shell.execute_reply":"2026-03-15T23:07:49.580564Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model_paths = [\n#     '/kaggle/input/r50-f0/pytorch/default/1/r50_F0/checkpoints/f0_epoch-14.pth',\n#     '/kaggle/input/r50-f1/pytorch/default/1/r50_F1/checkpoints/f1_epoch-14.pth',\n#     '/kaggle/input/r50-f2/pytorch/default/1/r50_F2/checkpoints/f2_epoch-13.pth',\n#     '/kaggle/input/r50-f3/pytorch/default/1/r50_F3/checkpoints/f3_epoch-15.pth',\n#     '/kaggle/input/r50-f4/pytorch/default/1/r50_F4/checkpoints/f4_epoch-14.pth'\n# ]\n\n# for model_path in model_paths:\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(model_path))\n#     _ = model.eval()\n#     models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:07:51.205159Z","iopub.execute_input":"2026-03-15T23:07:51.205506Z","iopub.status.idle":"2026-03-15T23:07:51.208927Z","shell.execute_reply.started":"2026-03-15T23:07:51.205471Z","shell.execute_reply":"2026-03-15T23:07:51.208088Z"}},"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-15T23:07:51.343905Z","iopub.execute_input":"2026-03-15T23:07:51.344273Z","iopub.status.idle":"2026-03-15T23:07:51.34792Z","shell.execute_reply.started":"2026-03-15T23:07:51.344238Z","shell.execute_reply":"2026-03-15T23:07:51.346926Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model_paths = [\n#     '/kaggle/input/r200/pytorch/default/1/r200d/checkpoints/f0_epoch-13.pth',\n#     '/kaggle/input/r200-f1/pytorch/default/1/r200d_F1/checkpoints/f1_epoch-15.pth',\n#     '/kaggle/input/r200-f2/pytorch/default/1/r200d_F2/checkpoints/f2_epoch-15.pth',\n#     '/kaggle/input/r200-f3/pytorch/default/1/r200d_F3/checkpoints/f3_epoch-14.pth',\n#     '/kaggle/input/r200-f4/pytorch/default/1/r200d_F4/checkpoints/f4_epoch-14.pth'\n# ]\n\n# for model_path in model_paths:\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(model_path))\n#     _ = model.eval()\n#     models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:07:51.47529Z","iopub.execute_input":"2026-03-15T23:07:51.47559Z","iopub.status.idle":"2026-03-15T23:07:51.479195Z","shell.execute_reply.started":"2026-03-15T23:07:51.475561Z","shell.execute_reply":"2026-03-15T23:07:51.47806Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model_paths = [\n#     '/kaggle/input/se50/pytorch/default/1/se50/checkpoints/f0_epoch-18.pth',\n#     '/kaggle/input/se50-f1/pytorch/default/1/se50_F1/checkpoints/f1_epoch-19.pth',\n#     '/kaggle/input/se50-f2/pytorch/default/1/se50_F2/checkpoints/f2_epoch-18.pth',\n#     '/kaggle/input/se50-f3/pytorch/default/1/se50_F3/checkpoints/f3_epoch-18.pth',\n#     '/kaggle/input/se50-f4/pytorch/default/1/se50_F4/checkpoints/f4_epoch-18.pth'\n# ]\n\n# for model_path in model_paths:\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(model_path))\n#     _ = model.eval()\n#     models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:07:51.633464Z","iopub.execute_input":"2026-03-15T23:07:51.633805Z","iopub.status.idle":"2026-03-15T23:07:51.637178Z","shell.execute_reply.started":"2026-03-15T23:07:51.633769Z","shell.execute_reply":"2026-03-15T23:07:51.636329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(models))\nprint(len(image_models))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:07:51.788269Z","iopub.execute_input":"2026-03-15T23:07:51.788593Z","iopub.status.idle":"2026-03-15T23:07:51.793003Z","shell.execute_reply.started":"2026-03-15T23:07:51.788561Z","shell.execute_reply":"2026-03-15T23:07:51.792142Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell Segmentator","metadata":{}},{"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,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-03-15T23:07:53.548819Z","iopub.execute_input":"2026-03-15T23:07:53.549154Z","iopub.status.idle":"2026-03-15T23:07:53.576617Z","shell.execute_reply.started":"2026-03-15T23:07:53.549125Z","shell.execute_reply":"2026-03-15T23:07:53.575744Z"}},"outputs":[],"execution_count":null},{"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-15T23:07:53.747979Z","iopub.execute_input":"2026-03-15T23:07:53.74827Z","iopub.status.idle":"2026-03-15T23:07:53.783269Z","shell.execute_reply.started":"2026-03-15T23:07:53.748244Z","shell.execute_reply":"2026-03-15T23:07:53.782388Z"},"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-15T23:07:53.882581Z","iopub.execute_input":"2026-03-15T23:07:53.88297Z","iopub.status.idle":"2026-03-15T23:08:03.874626Z","shell.execute_reply.started":"2026-03-15T23:07:53.88294Z","shell.execute_reply":"2026-03-15T23:08:03.873859Z"}},"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\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-15T23:08:03.876089Z","iopub.execute_input":"2026-03-15T23:08:03.876351Z","iopub.status.idle":"2026-03-15T23:08:03.888508Z","shell.execute_reply.started":"2026-03-15T23:08:03.876322Z","shell.execute_reply":"2026-03-15T23:08:03.887694Z"},"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-15T23:08:03.890353Z","iopub.execute_input":"2026-03-15T23:08:03.890718Z","iopub.status.idle":"2026-03-15T23:08:03.916063Z","shell.execute_reply.started":"2026-03-15T23:08:03.890671Z","shell.execute_reply":"2026-03-15T23:08:03.91537Z"}},"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\nimport gc\n\nif not PUBLIC_TEST:\n    rows = []\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        # 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    \n        # Get bounding boxes\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    \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            })\n    \n        del red_image, green_image, blue_image, yellow_image, ryb_image\n        del cell_mask, cell_bboxes \n        \n        gc.collect()\n        torch.cuda.empty_cache()\n    \n    tm = pd.DataFrame(rows)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:08:03.917291Z","iopub.execute_input":"2026-03-15T23:08:03.917525Z","iopub.status.idle":"2026-03-15T23:08:03.925439Z","shell.execute_reply.started":"2026-03-15T23:08:03.917503Z","shell.execute_reply":"2026-03-15T23:08:03.924495Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print_memory_usage()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:08:03.926821Z","iopub.execute_input":"2026-03-15T23:08:03.927157Z","iopub.status.idle":"2026-03-15T23:08:03.936455Z","shell.execute_reply.started":"2026-03-15T23:08:03.927123Z","shell.execute_reply":"2026-03-15T23:08:03.93562Z"}},"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-15T23:08:03.937489Z","iopub.execute_input":"2026-03-15T23:08:03.937784Z","iopub.status.idle":"2026-03-15T23:08:03.948038Z","shell.execute_reply.started":"2026-03-15T23:08:03.937758Z","shell.execute_reply":"2026-03-15T23:08:03.94731Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## If we read from a csv","metadata":{}},{"cell_type":"code","source":"precomputed_df = pd.read_csv('/kaggle/input/hpa-best-segmentation-take-use/best_segmentation_result.csv')\nprecomputed_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:08:03.949119Z","iopub.execute_input":"2026-03-15T23:08:03.949329Z","iopub.status.idle":"2026-03-15T23:08:04.902624Z","shell.execute_reply.started":"2026-03-15T23:08:03.949308Z","shell.execute_reply":"2026-03-15T23:08:04.901782Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if PUBLIC_TEST:\n    # Lista para armazenar os dicionários com informações de cada célula\n    imgs = []\n    \n    # Itera sobre todas as linhas do DataFrame precomputed_df\n    for i, x in precomputed_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\n    tm = pd.DataFrame(imgs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:08:04.904267Z","iopub.execute_input":"2026-03-15T23:08:04.904524Z","iopub.status.idle":"2026-03-15T23:08:05.295373Z","shell.execute_reply.started":"2026-03-15T23:08:04.904499Z","shell.execute_reply":"2026-03-15T23:08:05.294625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"precomputed_df = precomputed_df.drop(columns='PredictionString')\nprecomputed_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:08:05.296704Z","iopub.execute_input":"2026-03-15T23:08:05.297108Z","iopub.status.idle":"2026-03-15T23:08:05.307105Z","shell.execute_reply.started":"2026-03-15T23:08:05.29707Z","shell.execute_reply":"2026-03-15T23:08:05.306211Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tm.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:08:05.308456Z","iopub.execute_input":"2026-03-15T23:08:05.308824Z","iopub.status.idle":"2026-03-15T23:08:05.320935Z","shell.execute_reply.started":"2026-03-15T23:08:05.308788Z","shell.execute_reply":"2026-03-15T23:08:05.320152Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if PUBLIC_TEST:\n    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-15T23:08:05.322009Z","iopub.execute_input":"2026-03-15T23:08:05.322254Z","iopub.status.idle":"2026-03-15T23:08:07.169807Z","shell.execute_reply.started":"2026-03-15T23:08:05.32223Z","shell.execute_reply":"2026-03-15T23:08:07.168976Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_submission_df = pd.read_csv('../input/hpa-single-cell-image-classification/sample_submission.csv', index_col=0)\nsample_submission_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:08:07.171184Z","iopub.execute_input":"2026-03-15T23:08:07.171493Z","iopub.status.idle":"2026-03-15T23:08:07.1866Z","shell.execute_reply.started":"2026-03-15T23:08:07.171463Z","shell.execute_reply":"2026-03-15T23:08:07.18583Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SliceInferenceDataset(torch.utils.data.Dataset):\n    def __init__(self, df, tta=16, cfg=None, tfms=None, deterministic=False):\n        self.df = df                                  # DataFrame contendo informações de imagem_id e codificação da máscara\n        self.iids = self.df.image_id.unique()         # Lista única de IDs de imagem\n        self.tta = tta                                # Número de Test Time Augmentations a serem aplicados\n        self.deterministic = deterministic\n    \n    def __len__(self):\n        return len(self.iids)                         # Número total de imagens únicas\n\n    def __getitem__(self, idx):\n        iid = self.iids[idx]                          # Pega o image_id pelo índice\n        \n        # Lê os canais de imagem correspondentes a cada cor (em escala de cinza)\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        r = cv2.imread(mt, 0).astype(np.float) / 255.0\n        g = cv2.imread(pr, 0).astype(np.float) / 255.0\n        b = cv2.imread(nu, 0).astype(np.float) / 255.0\n        a = cv2.imread(er, 0).astype(np.float) / 255.0\n        \n        sz = r.shape[0]                               # Assume imagem quadrada (W = H)\n        img = np.stack([r, g, b, a], -1)              # Empilha os 4 canais para criar imagem 4-channel RGBY\n        \n        sli = []                                       # Lista para armazenar os crops das células (imagem + fname)\n        \n        # Itera sobre todas as células daquela imagem\n        for i, x in self.df[self.df.image_id == iid].iterrows():\n            # Decodifica a máscara RLE\n            bd = base64.b64decode(x.enc)\n            zd = zlib.decompress(bd)\n            encoded = [{'counts': zd, 'size': (sz, sz)}]\n            ded = coco_mask.decode(encoded)[:, :, 0]  # Máscara binária\n\n            # Se não tiver célula detectada\n            if len(np.unique(ded)) == 1:\n                continue\n\n            # Extrai bounding box da máscara\n            xr, yr = np.where(ded == 1)\n            sub = img[xr.min(): xr.max(), yr.min(): yr.max()]  # Recorta imagem\n            crop_sub_mask = ded[xr.min(): xr.max(), yr.min(): yr.max()]\n            crop_sub_mask = np.repeat(crop_sub_mask[:, :, np.newaxis], 4, axis=2)  # Expande máscara para 4 canais\n            \n            # Aplica máscara à imagem\n            r = sub * crop_sub_mask\n\n            # Ajusta para quadrado e redimensiona para 256×256\n            sli.append((cv2.resize(squarify(r, 0), (256, 256)).astype(np.float32), x.fname))\n\n        # Se nenhuma célula válida foi extraída da imagem, retorna None\n        if not sli:\n            return None\n        \n        BS, tta = len(sli) + 1, self.tta  # BS: batch size estimado (número de células + 1)\n        ipts = []                         # Lista de batches de células com TTA\n\n        raw_ipt = [e[0] for e in sli]     # Apenas as imagens (sem o fname)\n\n        # Aplica transformações de TTA\n        if tta == 1:\n            ipts.append(torch.stack([tensor_tfms(res_tfms(image=x)['image']) for x in raw_ipt]).float())\n        else:\n            if not self.deterministic:\n                for tt in range(tta):\n                    ipts.append(torch.stack([tensor_tfms(tta_tfms(image=x)['image']) for x in raw_ipt]).float())\n            else:\n                transforms = []\n                for flip in [False, True]:\n                    for k in range(4):\n                        transforms.append((flip, k))\n                for flip, k in transforms[:tta]:\n                    batch = []\n                    for x in raw_ipt:\n                        xr = x\n                        if flip:\n                            xr = np.flip(xr, axis=1)\n                        xr = np.rot90(xr, k).copy()\n                        xr = res_tfms(image=xr)['image']\n                        xr = tensor_tfms(xr)\n                        batch.append(xr)\n                    ipts.append(torch.stack(batch).float())\n\n        # Também transforma a imagem inteira para outro uso (ex: CAM)\n        image = image_tfms(image_res_tfms(image=img)['image'])\n            \n        # Retorna:\n        # - imagem da amostra inteira (RGBY)\n        # - células da imagem após TTA\n        # - batch size estimado\n        # - número de células detectadas\n        # - número de TTA\n        # - image_id\n        # - lista dos nomes de arquivos para cada célula\n        return image, ipts, BS, len(sli), tta, iid, [x[1] for x in sli]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:08:59.903026Z","iopub.execute_input":"2026-03-15T23:08:59.903323Z","iopub.status.idle":"2026-03-15T23:08:59.918392Z","shell.execute_reply.started":"2026-03-15T23:08:59.903298Z","shell.execute_reply":"2026-03-15T23:08:59.917471Z"}},"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, tta=8, deterministic=False)\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-15T23:09:00.099599Z","iopub.execute_input":"2026-03-15T23:09:00.09995Z","iopub.status.idle":"2026-03-15T23:09:00.105505Z","shell.execute_reply.started":"2026-03-15T23:09:00.099923Z","shell.execute_reply":"2026-03-15T23:09:00.104678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 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 full_image, 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 start_idx 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            for model in image_models:\n                full_image = full_image.float().cuda()\n                model = model.float() \n                exp.append(model(full_image))\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() \n                for model in cell_models:\n                    with torch.cuda.amp.autocast():\n                        cell_output, image_output = model(cell_batch, len(cell_batch))\n                    res.append(cell_output.float()) \n                    exp.append(image_output.float())\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)\n        image_probs_mean = np.stack(image_prob).mean(0)\n\n        cell_predictions_batches.append(cell_probs_mean)\n        image_predictions_batches.append(image_probs_mean)\n\n    # Concatena predições por célula (final do loop da imagem)\n    cell_predictions = np.concatenate(cell_predictions_batches)\n    cell_df = pd.DataFrame(cell_predictions, index=cell_filenames)\n    \n    image_df = pd.DataFrame(\n        np.concatenate(image_predictions_batches).mean(0).reshape(1, 19), \n        index=[image_id]\n    )\n\n    # Salva resultados\n    image_predictions_dfs.append(image_df)\n    cell_predictions_dfs.append(cell_df)\n\n    del full_image, 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-15T23:09:04.536527Z","iopub.execute_input":"2026-03-15T23:09:04.536862Z","iopub.status.idle":"2026-03-15T23:09:47.359897Z","shell.execute_reply.started":"2026-03-15T23:09:04.53683Z","shell.execute_reply":"2026-03-15T23:09:47.358606Z"}},"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-15T23:09:47.360835Z","iopub.status.idle":"2026-03-15T23:09:47.36115Z","shell.execute_reply":"2026-03-15T23:09:47.361006Z"}},"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)\nmerge_pred = pub_pred * image_pred.loc[pub_pred.index]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:07:27.892266Z","iopub.status.idle":"2026-03-15T23:07:27.892601Z","shell.execute_reply":"2026-03-15T23:07:27.892448Z"}},"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-15T23:07:27.893548Z","iopub.status.idle":"2026-03-15T23:07:27.893892Z","shell.execute_reply":"2026-03-15T23:07:27.893733Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Save prediction","metadata":{}},{"cell_type":"code","source":"merge_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-15T23:07:27.894953Z","iopub.status.idle":"2026-03-15T23:07:27.895256Z","shell.execute_reply":"2026-03-15T23:07:27.895112Z"}},"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': sample_submission_df.loc[iid].ImageWidth,   # Largura da imagem (obtida do DataFrame original `df`)\n        'ImageHeight': sample_submission_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-15T23:07:27.895932Z","iopub.status.idle":"2026-03-15T23:07:27.896257Z","shell.execute_reply":"2026-03-15T23:07:27.896109Z"}},"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-15T23:07:27.897037Z","iopub.status.idle":"2026-03-15T23:07:27.897392Z","shell.execute_reply":"2026-03-15T23:07:27.897213Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fast_sub.head(2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:07:27.898181Z","iopub.status.idle":"2026-03-15T23:07:27.898512Z","shell.execute_reply":"2026-03-15T23:07:27.898352Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## save","metadata":{}},{"cell_type":"code","source":"sample_submission_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:07:27.899219Z","iopub.status.idle":"2026-03-15T23:07:27.899549Z","shell.execute_reply":"2026-03-15T23:07:27.899387Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Substitui as previsões do DataFrame `sample_submission_df` pelas do `fast_sub`\n# drop(fast_sub.index): remove as linhas de 'sample_submission_df' 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_df.drop(fast_sub.index), fast_sub], axis=0)\n\n# Reordena as linhas para seguir exatamente a ordem original de 'sample_submission_df'\nsub2 = sub2.loc[sample_submission_df.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-15T23:07:27.900321Z","iopub.status.idle":"2026-03-15T23:07:27.90066Z","shell.execute_reply":"2026-03-15T23:07:27.900501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T23:07:27.901458Z","iopub.status.idle":"2026-03-15T23:07:27.901788Z","shell.execute_reply":"2026-03-15T23:07:27.901614Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}