torch_em.model.unetr

   1from functools import partial
   2from collections import OrderedDict
   3from typing import Optional, Tuple, Union, Literal
   4
   5import torch
   6import torch.nn as nn
   7import torch.nn.functional as F
   8
   9from .vit import get_vision_transformer
  10from .unet import Decoder, ConvBlock2d, ConvBlock3d, Upsampler2d, Upsampler3d, _update_conv_kwargs
  11
  12try:
  13    from micro_sam.util import get_sam_model
  14except ImportError:
  15    get_sam_model = None
  16
  17try:
  18    from micro_sam.v2.util import get_sam2_model
  19except ImportError:
  20    get_sam2_model = None
  21
  22try:
  23    from micro_sam3.util import get_sam3_model
  24except ImportError:
  25    get_sam3_model = None
  26
  27
  28#
  29# UNETR IMPLEMENTATION [Vision Transformer (ViT from SAM / CellposeSAM / SAM2 / SAM3 / DINOv2 / DINOv3 / MAE / ScaleMAE) + UNet Decoder from `torch_em`]  # noqa
  30#
  31
  32
  33def _check_input_normalization_range(
  34    x: torch.Tensor,
  35    expected_range: Optional[Tuple[float, float]],
  36    unit_scale_max: Optional[float] = None,
  37) -> None:
  38    """Check whether raw inputs match the value range expected by the model normalizer.
  39
  40    Args:
  41        x: The input tensor to validate.
  42        expected_range: The (min, max) value range the input must lie within. If None, all checks are skipped.
  43        unit_scale_max: If set, raises an error when the input's maximum value is at or below this threshold,
  44            catching inputs that are likely in the wrong scale (e.g. [0, 1] instead of [0, 255] for SAM1).
  45    """
  46    if expected_range is None:
  47        return
  48
  49    if not torch.all(torch.isfinite(x)):
  50        raise ValueError("The input contains NaN or infinite values before normalization.")
  51
  52    min_value, max_value = expected_range
  53    if torch.any((x < min_value) | (x > max_value)):
  54        actual_min, actual_max = torch.aminmax(x.detach())
  55        raise ValueError(
  56            "The input is outside the expected scale before normalization: "
  57            f"expected values in [{min_value}, {max_value}], got [{actual_min.item()}, {actual_max.item()}]. "
  58            "Please check whether the raw inputs should be scaled to [0, 1] or kept in [0, 255] "
  59            "before applying the pretrained normalization statistics."
  60        )
  61
  62    if unit_scale_max is not None:
  63        actual_max = x.detach().max().item()
  64        if actual_max <= unit_scale_max:
  65            raise ValueError(
  66                f"The input maximum value ({actual_max:.4f}) suggests the input is in the wrong scale: "
  67                f"expected inputs with values in [{min_value}, {max_value}], "
  68                f"but the maximum is only {actual_max:.4f}. "
  69                "Please check whether the raw inputs should be scaled to [0, 255] instead of [0, 1]."
  70            )
  71
  72
  73def _as_stats(mean, std, device, dtype, is_3d: bool):
  74    view_shape = (1, -1, 1, 1, 1) if is_3d else (1, -1, 1, 1)
  75    pixel_mean = torch.tensor(mean, device=device, dtype=dtype).view(*view_shape)
  76    pixel_std = torch.tensor(std, device=device, dtype=dtype).view(*view_shape)
  77    return pixel_mean, pixel_std
  78
  79
  80class UNETRBase(nn.Module):
  81    """Base class for implementing a UNETR.
  82
  83    Args:
  84        img_size: The size of the input for the image encoder. Input images will be resized to match this size.
  85        backbone: The name of the vision transformer implementation.
  86            One of "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3"
  87            (see all combinations below)
  88        encoder: The vision transformer. Can either be a name, such as "vit_b"
  89            (see all combinations for this below) or a torch module.
  90        decoder: The convolutional decoder.
  91        out_channels: The number of output channels of the UNETR.
  92        use_sam_stats: Whether to normalize the input data with the statistics of the
  93            pretrained SAM / SAM2 / SAM3 model.
  94        use_dino_stats: Whether to normalize the input data with the statistics of the
  95            pretrained DINOv2 / DINOv3 model.
  96        use_imagenet_stats: Whether to normalize with standard ImageNet statistics, i.e.
  97            mean - (0.485, 0.456, 0.406) and std - (0.229, 0.224, 0.225), raw inputs between range [0, 1].
  98            Use this with the 'torchvision' backbone when loading pretrained weights.
  99        use_mae_stats: Whether to normalize the input data with the statistics of the pretrained MAE model.
 100        resize_input: Whether to resize the input images to match `img_size`.
 101            By default, it resizes the inputs to match the `img_size`.
 102        encoder_checkpoint: Checkpoint for initializing the vision transformer.
 103            Can either be a filepath or an already loaded checkpoint.
 104        final_activation: The activation to apply to the UNETR output.
 105        use_skip_connection: Whether to use skip connections. By default, it uses skip connections.
 106        embed_dim: The embedding dimensionality, corresponding to the output dimension of the vision transformer.
 107        use_conv_transpose: Whether to use transposed convolutions instead of resampling for upsampling.
 108            By default, it uses resampling for upsampling.
 109        perform_range_checks: Whether to validate the input value range before normalization on each forward pass.
 110            You can disable the checks to avoid GPU sync overhead during training when inputs are known to be correct.
 111        initial_features: The number of features of the finest decoder level. The features per level are
 112            'initial_features * gain ** i', so this scales the decoder parameters quadratically.
 113
 114        NOTE: The currently supported combinations of 'backbone' x 'encoder' are the following:
 115
 116        SAM_family_models:
 117            - 'sam' x 'vit_b'
 118            - 'sam' x 'vit_l'
 119            - 'sam' x 'vit_h'
 120            - 'sam2' x 'hvit_t'
 121            - 'sam2' x 'hvit_s'
 122            - 'sam2' x 'hvit_b'
 123            - 'sam2' x 'hvit_l'
 124            - 'sam3' x 'vit_pe'
 125            - 'cellpose_sam' x 'vit_l'
 126
 127        DINO_family_models:
 128            - 'dinov2' x 'vit_s'
 129            - 'dinov2' x 'vit_b'
 130            - 'dinov2' x 'vit_l'
 131            - 'dinov2' x 'vit_g'
 132            - 'dinov2' x 'vit_s_reg4'
 133            - 'dinov2' x 'vit_b_reg4'
 134            - 'dinov2' x 'vit_l_reg4'
 135            - 'dinov2' x 'vit_g_reg4'
 136            - 'dinov3' x 'vit_s'
 137            - 'dinov3' x 'vit_s+'
 138            - 'dinov3' x 'vit_b'
 139            - 'dinov3' x 'vit_l'
 140            - 'dinov3' x 'vit_l+'
 141            - 'dinov3' x 'vit_h+'
 142            - 'dinov3' x 'vit_7b'
 143
 144        MAE_family_models:
 145            - 'mae' x 'vit_b'
 146            - 'mae' x 'vit_l'
 147            - 'mae' x 'vit_h'
 148            - 'scalemae' x 'vit_b'
 149            - 'scalemae' x 'vit_l'
 150            - 'scalemae' x 'vit_h'
 151
 152        torchvision_models:
 153            - 'torchvision' x 'vit_b_16'
 154            - 'torchvision' x 'vit_b_32'
 155            - 'torchvision' x 'vit_l_16'
 156            - 'torchvision' x 'vit_l_32'
 157            - 'torchvision' x 'vit_h_14'
 158    """
 159    def __init__(
 160        self,
 161        img_size: int = 1024,
 162        backbone: Literal[
 163            "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3", "torchvision"
 164        ] = "sam",
 165        encoder: Optional[Union[nn.Module, str]] = "vit_b",
 166        decoder: Optional[nn.Module] = None,
 167        out_channels: int = 1,
 168        use_sam_stats: bool = False,
 169        use_mae_stats: bool = False,
 170        use_dino_stats: bool = False,
 171        use_imagenet_stats: bool = False,
 172        resize_input: bool = True,
 173        encoder_checkpoint: Optional[Union[str, OrderedDict]] = None,
 174        final_activation: Optional[Union[str, nn.Module]] = None,
 175        use_skip_connection: bool = True,
 176        embed_dim: Optional[int] = None,
 177        use_conv_transpose: bool = False,
 178        perform_range_checks: bool = True,
 179        initial_features: int = 64,
 180        **kwargs
 181    ) -> None:
 182        super().__init__()
 183
 184        self.img_size = img_size
 185        self.use_sam_stats = use_sam_stats
 186        self.use_mae_stats = use_mae_stats
 187        self.use_dino_stats = use_dino_stats
 188        self.use_imagenet_stats = use_imagenet_stats
 189        self.use_skip_connection = use_skip_connection
 190        self.resize_input = resize_input
 191        self.perform_range_checks = perform_range_checks
 192        self.use_conv_transpose = use_conv_transpose
 193        self.initial_features = initial_features
 194        self.backbone = backbone
 195
 196        if isinstance(encoder, str):  # e.g. "vit_b" / "hvit_b" / "vit_pe"
 197            print(f"Using {encoder} from {backbone.upper()}")
 198            self.encoder = get_vision_transformer(img_size=img_size, backbone=backbone, model=encoder, **kwargs)
 199
 200            if encoder_checkpoint is not None:
 201                self._load_encoder_from_checkpoint(backbone=backbone, encoder=encoder, checkpoint=encoder_checkpoint)
 202
 203            if embed_dim is None:
 204                embed_dim = self.encoder.embed_dim
 205
 206            # For SAM1 encoder, if 'apply_neck' is applied, the embedding dimension must change.
 207            if hasattr(self.encoder, "apply_neck") and self.encoder.apply_neck:
 208                embed_dim = self.encoder.neck[2].out_channels  # the value is 256
 209
 210        else:  # `nn.Module` ViT backbone
 211            self.encoder = encoder
 212
 213            have_neck = False
 214            for name, _ in self.encoder.named_parameters():
 215                if name.startswith("neck"):
 216                    have_neck = True
 217
 218            if embed_dim is None:
 219                if have_neck:
 220                    embed_dim = self.encoder.neck[2].out_channels  # the value is 256
 221                else:
 222                    embed_dim = self.encoder.patch_embed.proj.out_channels
 223
 224        self.embed_dim = embed_dim
 225        self.final_activation = self._get_activation(final_activation)
 226
 227    def _load_encoder_from_checkpoint(self, backbone, encoder, checkpoint):
 228        """Function to load pretrained weights to the image encoder.
 229        """
 230        if isinstance(checkpoint, str):
 231            if backbone == "sam" and isinstance(encoder, str):
 232                # If we have a SAM encoder, then we first try to load the full SAM Model
 233                # (using micro_sam) and otherwise fall back on directly loading the encoder state
 234                # from the checkpoint
 235                try:
 236                    _, model = get_sam_model(model_type=encoder, checkpoint_path=checkpoint, return_sam=True)
 237                    encoder_state = model.image_encoder.state_dict()
 238                except Exception:
 239                    # Try loading the encoder state directly from a checkpoint.
 240                    encoder_state = torch.load(checkpoint, weights_only=False)
 241
 242            elif backbone == "cellpose_sam" and isinstance(encoder, str):
 243                # The architecture matches CellposeSAM exactly (same rel_pos sizes),
 244                # so weights load directly without any interpolation.
 245                encoder_state = torch.load(checkpoint, map_location="cpu", weights_only=False)
 246                # Handle DataParallel/DistributedDataParallel prefix.
 247                if any(k.startswith("module.") for k in encoder_state.keys()):
 248                    encoder_state = OrderedDict(
 249                        {k[len("module."):]: v for k, v in encoder_state.items()}
 250                    )
 251                # Extract encoder weights from CellposeSAM checkpoint format (strip 'encoder.' prefix).
 252                if any(k.startswith("encoder.") for k in encoder_state.keys()):
 253                    encoder_state = OrderedDict(
 254                        {k[len("encoder."):]: v for k, v in encoder_state.items() if k.startswith("encoder.")}
 255                    )
 256
 257            elif backbone == "sam2" and isinstance(encoder, str):
 258                # If we have a SAM2 encoder, then we first try to load the full SAM2 Model.
 259                # (using micro_sam2) and otherwise fall back on directly loading the encoder state
 260                # from the checkpoint
 261                try:
 262                    model = get_sam2_model(model_type=encoder, checkpoint_path=checkpoint)
 263                    encoder_state = model.image_encoder.state_dict()
 264                except Exception:
 265                    # Try loading the encoder state directly from a checkpoint.
 266                    encoder_state = torch.load(checkpoint, weights_only=False)
 267
 268            elif backbone == "sam3" and isinstance(encoder, str):
 269                # If we have a SAM3 encoder, then we first try to load the full SAM3 Model.
 270                # (using micro_sam3) and otherwise fall back on directly loading the encoder state
 271                # from the checkpoint
 272                try:
 273                    model = get_sam3_model(checkpoint_path=checkpoint)
 274                    encoder_state = model.backbone.vision_backbone.state_dict()
 275                    # Let's align loading the encoder weights with expected parameter names
 276                    encoder_state = {
 277                        k[len("trunk."):] if k.startswith("trunk.") else k: v for k, v in encoder_state.items()
 278                    }
 279                    # And drop the 'convs' and 'sam2_convs' - these seem like some upsampling blocks.
 280                    encoder_state = {
 281                        k: v for k, v in encoder_state.items()
 282                        if not (k.startswith("convs.") or k.startswith("sam2_convs."))
 283                    }
 284                except Exception:
 285                    # Try loading the encoder state directly from a checkpoint.
 286                    encoder_state = torch.load(checkpoint, weights_only=False)
 287
 288            elif backbone == "mae":
 289                # vit initialization hints from:
 290                #     - https://github.com/facebookresearch/mae/blob/main/main_finetune.py#L233-L242
 291                encoder_state = torch.load(checkpoint, weights_only=False)["model"]
 292                encoder_state = OrderedDict({
 293                    k: v for k, v in encoder_state.items() if (k != "mask_token" and not k.startswith("decoder"))
 294                })
 295                # Let's remove the `head` from our current encoder (as the MAE pretrained don't expect it)
 296                current_encoder_state = self.encoder.state_dict()
 297                if ("head.weight" in current_encoder_state) and ("head.bias" in current_encoder_state):
 298                    del self.encoder.head
 299
 300            elif backbone == "scalemae":
 301                # Load the encoder state directly from a checkpoint.
 302                encoder_state = torch.load(checkpoint)["model"]
 303                encoder_state = OrderedDict({
 304                    k: v for k, v in encoder_state.items()
 305                    if not k.startswith(("mask_token", "decoder", "fcn", "fpn", "pos_embed"))
 306                })
 307
 308                # Let's remove the `head` from our current encoder (as the MAE pretrained don't expect it)
 309                current_encoder_state = self.encoder.state_dict()
 310                if ("head.weight" in current_encoder_state) and ("head.bias" in current_encoder_state):
 311                    del self.encoder.head
 312
 313                if "pos_embed" in current_encoder_state:  # NOTE: ScaleMAE uses 'pos. embeddings' in a diff. format.
 314                    del self.encoder.pos_embed
 315
 316            elif backbone in ["dinov2", "dinov3"]:  # Load the encoder state directly from a checkpoint.
 317                encoder_state = torch.load(checkpoint)
 318
 319            elif backbone == "torchvision":
 320                encoder_state = torch.load(checkpoint, weights_only=False)
 321
 322            else:
 323                raise ValueError(
 324                    f"We don't support either the '{backbone}' backbone or the '{encoder}' model combination (or both)."
 325                )
 326
 327        else:
 328            encoder_state = checkpoint
 329
 330        if backbone == "torchvision":
 331            if "state_dict" in encoder_state:
 332                encoder_state = encoder_state["state_dict"]
 333            encoder_state = {k: v for k, v in encoder_state.items() if not k.startswith("heads.")}
 334
 335        self.encoder.load_state_dict(encoder_state)
 336
 337    def _get_activation(self, activation):
 338        return_activation = None
 339        if activation is None:
 340            return None
 341        if isinstance(activation, nn.Module):
 342            return activation
 343        if isinstance(activation, str):
 344            return_activation = getattr(nn, activation, None)
 345        if return_activation is None:
 346            raise ValueError(f"Invalid activation: {activation}")
 347
 348        return return_activation()
 349
 350    @staticmethod
 351    def get_preprocess_shape(oldh: int, oldw: int, long_side_length: int) -> Tuple[int, int]:
 352        """Compute the output size given input size and target long side length.
 353
 354        Args:
 355            oldh: The input image height.
 356            oldw: The input image width.
 357            long_side_length: The longest side length for resizing.
 358
 359        Returns:
 360            The new image height.
 361            The new image width.
 362        """
 363        scale = long_side_length * 1.0 / max(oldh, oldw)
 364        newh, neww = oldh * scale, oldw * scale
 365        neww = int(neww + 0.5)
 366        newh = int(newh + 0.5)
 367        return (newh, neww)
 368
 369    def resize_longest_side(self, image: torch.Tensor) -> torch.Tensor:
 370        """Resize the image so that the longest side has the correct length.
 371
 372        Expects batched images with shape BxCxHxW OR BxCxDxHxW and float format.
 373
 374        Args:
 375            image: The input image.
 376
 377        Returns:
 378            The resized image.
 379        """
 380        if image.ndim == 4:  # i.e. 2d image
 381            target_size = self.get_preprocess_shape(image.shape[2], image.shape[3], self.encoder.img_size)
 382            return F.interpolate(image, target_size, mode="bilinear", align_corners=False, antialias=True)
 383        elif image.ndim == 5:  # i.e. 3d volume
 384            B, C, Z, H, W = image.shape
 385            target_size = self.get_preprocess_shape(H, W, self.img_size)
 386            return F.interpolate(image, (Z, *target_size), mode="trilinear", align_corners=False)
 387        else:
 388            raise ValueError("Expected 4d or 5d inputs, got", image.shape)
 389
 390    def _as_stats(self, mean, std, device, dtype, is_3d: bool):
 391        """@private
 392        """
 393        return _as_stats(mean, std, device, dtype, is_3d)
 394
 395    def _check_input_normalization_range(self, x: torch.Tensor, expected_range: Optional[Tuple[float, float]]) -> None:
 396        """@private
 397        """
 398        _check_input_normalization_range(x, expected_range)
 399
 400    def encode(self, x: torch.Tensor):
 401        """Preprocess the input and run the image encoder.
 402
 403        Args:
 404            x: The input tensor.
 405
 406        Returns:
 407            The encoder features to pass to `decode` and the spatial shape after preprocessing.
 408        """
 409        raise NotImplementedError
 410
 411    def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor:
 412        """Run the convolutional decoder on the encoder features.
 413
 414        Args:
 415            features: The encoder features returned by `encode`.
 416            input_shape: The spatial shape after preprocessing, returned by `encode`.
 417            original_shape: The spatial shape of the original input.
 418
 419        Returns:
 420            The UNETR output, resized to `original_shape`.
 421        """
 422        raise NotImplementedError
 423
 424    def forward(self, x: torch.Tensor) -> torch.Tensor:
 425        """Apply the UNETR to the input data.
 426
 427        Args:
 428            x: The input tensor.
 429
 430        Returns:
 431            The UNETR output.
 432        """
 433        features, input_shape = self.encode(x)
 434        return self.decode(features, input_shape, tuple(x.shape[2:]))
 435
 436    def preprocess(self, x: torch.Tensor) -> torch.Tensor:
 437        """@private
 438        """
 439        return preprocess_vit_inputs(
 440            x,
 441            use_sam_stats=self.use_sam_stats,
 442            backbone=self.backbone,
 443            use_mae_stats=self.use_mae_stats,
 444            use_dino_stats=self.use_dino_stats,
 445            use_imagenet_stats=self.use_imagenet_stats,
 446            resize_input=self.resize_input,
 447            img_size=self.img_size,
 448            encoder_img_size=self.encoder.img_size,
 449            perform_range_checks=self.perform_range_checks,
 450        )
 451
 452    def postprocess_masks(
 453        self, masks: torch.Tensor, input_size: Tuple[int, ...], original_size: Tuple[int, ...],
 454    ) -> torch.Tensor:
 455        """@private
 456        """
 457        if masks.ndim == 4:  # i.e. 2d labels
 458            masks = F.interpolate(
 459                masks,
 460                (self.encoder.img_size, self.encoder.img_size),
 461                mode="bilinear",
 462                align_corners=False,
 463            )
 464            masks = masks[..., : input_size[0], : input_size[1]]
 465            masks = F.interpolate(masks, original_size, mode="bilinear", align_corners=False)
 466
 467        elif masks.ndim == 5:  # i.e. 3d volumetric labels
 468            masks = F.interpolate(
 469                masks,
 470                (input_size[0], self.img_size, self.img_size),
 471                mode="trilinear",
 472                align_corners=False,
 473            )
 474            masks = masks[..., :input_size[0], :input_size[1], :input_size[2]]
 475            masks = F.interpolate(masks, original_size, mode="trilinear", align_corners=False)
 476
 477        else:
 478            raise ValueError("Expected 4d or 5d labels, got", masks.shape)
 479
 480        return masks
 481
 482
 483def preprocess_vit_inputs(
 484    x: torch.Tensor,
 485    use_sam_stats: bool = False,
 486    backbone: str = "sam",
 487    use_mae_stats: bool = False,
 488    use_dino_stats: bool = False,
 489    use_imagenet_stats: bool = False,
 490    resize_input: bool = True,
 491    img_size: int = 1024,
 492    encoder_img_size: int = 1024,
 493    perform_range_checks: bool = True,
 494) -> Tuple[torch.Tensor, Tuple]:
 495    """Preprocess inputs for ViT-backbones in UNETR models.
 496
 497    Handles normalization stat selection, input range validation, optional resizing to the longest side,
 498    and padding to `encoder_img_size`. Can be used as a standalone function without a model instance.
 499
 500    Args:
 501        x: Input tensor of shape (B, C, H, W) for 2D or (B, C, Z, H, W) for 3D.
 502        use_sam_stats: Whether to normalize with SAM/SAM2/SAM3 backbone statistics.
 503        backbone: The backbone name - controls which SAM stats are used when `use_sam_stats=True`.
 504        use_mae_stats: Whether to normalize with MAE statistics.
 505        use_dino_stats: Whether to normalize with DINOv2/DINOv3 statistics.
 506        use_imagenet_stats: Whether to normalize with standard ImageNet statistics
 507            (mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), range [0, 1]).
 508            Use this for torchvision pretrained backbones.
 509        resize_input: Whether to resize the input to the longest side before padding.
 510        img_size: The model image size, used for 3D resize.
 511        encoder_img_size: The encoder image size, used for 2D resize and padding.
 512        perform_range_checks: Whether to validate the expected input value range before normalization.
 513            You can disable the checks to avoid GPU sync overhead during training when inputs are known to be correct.
 514
 515    Returns:
 516        The preprocessed tensor and the spatial shape after resizing (before padding).
 517    """
 518    is_3d = (x.ndim == 5)
 519    device, dtype = x.device, x.dtype
 520    mean, std = (0.0, 0.0, 0.0), (1.0, 1.0, 1.0)
 521    expected_range = None
 522    unit_scale_max = None
 523
 524    if use_sam_stats:
 525        if backbone == "sam2":
 526            mean, std = (0.485, 0.456, 0.406), (0.229, 0.224, 0.225)
 527            expected_range = (0.0, 1.0)
 528        elif backbone == "sam3":
 529            mean, std = (0.5, 0.5, 0.5), (0.5, 0.5, 0.5)
 530            expected_range = (0.0, 1.0)
 531        else:  # sam1 / default
 532            mean, std = (123.675, 116.28, 103.53), (58.395, 57.12, 57.375)
 533            expected_range = (0.0, 255.0)
 534            unit_scale_max = 1.0
 535    elif use_mae_stats:  # TODO: add mean std from mae / scalemae experiments (or open up arguments for this)
 536        raise NotImplementedError
 537    elif use_dino_stats or use_imagenet_stats:
 538        mean, std = (0.485, 0.456, 0.406), (0.229, 0.224, 0.225)
 539        expected_range = (0.0, 1.0)
 540    else:
 541        mean, std = (0.0, 0.0, 0.0), (1.0, 1.0, 1.0)
 542        expected_range = None
 543
 544    if perform_range_checks:
 545        _check_input_normalization_range(x, expected_range, unit_scale_max)
 546    pixel_mean, pixel_std = _as_stats(mean, std, device=device, dtype=dtype, is_3d=is_3d)
 547
 548    if resize_input:
 549        if x.ndim == 4:
 550            target_size = UNETRBase.get_preprocess_shape(x.shape[2], x.shape[3], encoder_img_size)
 551            x = F.interpolate(x, target_size, mode="bilinear", align_corners=False, antialias=True)
 552        elif x.ndim == 5:
 553            B, C, Z, H, W = x.shape
 554            target_size = UNETRBase.get_preprocess_shape(H, W, img_size)
 555            x = F.interpolate(x, (Z, *target_size), mode="trilinear", align_corners=False)
 556
 557    input_shape = x.shape[-3:] if is_3d else x.shape[-2:]
 558
 559    x = (x - pixel_mean) / pixel_std
 560    h, w = x.shape[-2:]
 561    padh = encoder_img_size - h
 562    padw = encoder_img_size - w
 563
 564    if is_3d:
 565        x = F.pad(x, (0, padw, 0, padh, 0, 0))
 566    else:
 567        x = F.pad(x, (0, padw, 0, padh))
 568
 569    return x, input_shape
 570
 571
 572class UNETR(UNETRBase):
 573    """A (2d-only) UNet Transformer using a vision transformer as encoder and a convolutional decoder.
 574    """
 575    def __init__(
 576        self,
 577        img_size: int = 1024,
 578        backbone: Literal[
 579            "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3", "torchvision"
 580        ] = "sam",
 581        encoder: Optional[Union[nn.Module, str]] = "vit_b",
 582        decoder: Optional[nn.Module] = None,
 583        out_channels: int = 1,
 584        use_sam_stats: bool = False,
 585        use_mae_stats: bool = False,
 586        use_dino_stats: bool = False,
 587        use_imagenet_stats: bool = False,
 588        resize_input: bool = True,
 589        encoder_checkpoint: Optional[Union[str, OrderedDict]] = None,
 590        final_activation: Optional[Union[str, nn.Module]] = None,
 591        use_skip_connection: bool = True,
 592        embed_dim: Optional[int] = None,
 593        use_conv_transpose: bool = False,
 594        perform_range_checks: bool = True,
 595        **kwargs
 596    ) -> None:
 597
 598        super().__init__(
 599            img_size=img_size,
 600            backbone=backbone,
 601            encoder=encoder,
 602            decoder=decoder,
 603            out_channels=out_channels,
 604            use_sam_stats=use_sam_stats,
 605            use_mae_stats=use_mae_stats,
 606            use_dino_stats=use_dino_stats,
 607            use_imagenet_stats=use_imagenet_stats,
 608            resize_input=resize_input,
 609            encoder_checkpoint=encoder_checkpoint,
 610            final_activation=final_activation,
 611            use_skip_connection=use_skip_connection,
 612            embed_dim=embed_dim,
 613            use_conv_transpose=use_conv_transpose,
 614            perform_range_checks=perform_range_checks,
 615            **kwargs,
 616        )
 617
 618        encoder = self.encoder
 619
 620        if backbone == "sam2" and hasattr(encoder, "trunk"):
 621            in_chans = encoder.trunk.patch_embed.proj.in_channels
 622        elif hasattr(encoder, "in_chans"):
 623            in_chans = encoder.in_chans
 624        else:  # `nn.Module` ViT backbone.
 625            try:
 626                in_chans = encoder.patch_embed.proj.in_channels
 627            except AttributeError:  # for getting the input channels while using 'vit_t' from MobileSam
 628                in_chans = encoder.patch_embed.seq[0].c.in_channels
 629
 630        # parameters for the decoder network
 631        depth = 3
 632        gain = 2
 633        features_decoder = [self.initial_features * gain ** i for i in range(depth + 1)][::-1]
 634        scale_factors = depth * [2]
 635        self.out_channels = out_channels
 636
 637        # choice of upsampler - to use (bilinear interpolation + conv) or conv transpose
 638        _upsampler = SingleDeconv2DBlock if use_conv_transpose else Upsampler2d
 639
 640        self.decoder = decoder or Decoder(
 641            features=features_decoder,
 642            scale_factors=scale_factors[::-1],
 643            conv_block_impl=ConvBlock2d,
 644            sampler_impl=_upsampler,
 645        )
 646
 647        if use_skip_connection:
 648            self.deconv1 = Deconv2DBlock(
 649                in_channels=self.embed_dim,
 650                out_channels=features_decoder[0],
 651                use_conv_transpose=use_conv_transpose,
 652            )
 653            self.deconv2 = nn.Sequential(
 654                Deconv2DBlock(
 655                    in_channels=self.embed_dim,
 656                    out_channels=features_decoder[0],
 657                    use_conv_transpose=use_conv_transpose,
 658                ),
 659                Deconv2DBlock(
 660                    in_channels=features_decoder[0],
 661                    out_channels=features_decoder[1],
 662                    use_conv_transpose=use_conv_transpose,
 663                )
 664            )
 665            self.deconv3 = nn.Sequential(
 666                Deconv2DBlock(
 667                    in_channels=self.embed_dim,
 668                    out_channels=features_decoder[0],
 669                    use_conv_transpose=use_conv_transpose,
 670                ),
 671                Deconv2DBlock(
 672                    in_channels=features_decoder[0],
 673                    out_channels=features_decoder[1],
 674                    use_conv_transpose=use_conv_transpose,
 675                ),
 676                Deconv2DBlock(
 677                    in_channels=features_decoder[1],
 678                    out_channels=features_decoder[2],
 679                    use_conv_transpose=use_conv_transpose,
 680                )
 681            )
 682            self.deconv4 = ConvBlock2d(in_chans, features_decoder[-1])
 683        else:
 684            self.deconv1 = Deconv2DBlock(
 685                in_channels=self.embed_dim,
 686                out_channels=features_decoder[0],
 687                use_conv_transpose=use_conv_transpose,
 688            )
 689            self.deconv2 = Deconv2DBlock(
 690                in_channels=features_decoder[0],
 691                out_channels=features_decoder[1],
 692                use_conv_transpose=use_conv_transpose,
 693            )
 694            self.deconv3 = Deconv2DBlock(
 695                in_channels=features_decoder[1],
 696                out_channels=features_decoder[2],
 697                use_conv_transpose=use_conv_transpose,
 698            )
 699            self.deconv4 = Deconv2DBlock(
 700                in_channels=features_decoder[2],
 701                out_channels=features_decoder[3],
 702                use_conv_transpose=use_conv_transpose,
 703            )
 704
 705        self.base = ConvBlock2d(self.embed_dim, features_decoder[0])
 706        self.out_conv = nn.Conv2d(features_decoder[-1], out_channels, 1)
 707        self.deconv_out = _upsampler(
 708            scale_factor=2, in_channels=features_decoder[-1], out_channels=features_decoder[-1]
 709        )
 710        self.decoder_head = ConvBlock2d(2 * features_decoder[-1], features_decoder[-1])
 711
 712    def encode(self, x: torch.Tensor):
 713        """Preprocess the input and run the image encoder.
 714
 715        Args:
 716            x: The input tensor of shape (B, C, Y, X).
 717
 718        Returns:
 719            The features as a tuple of the image embeddings, the list of intermediate encoder outputs
 720            (None if the encoder returns only the embeddings) and the preprocessed input, which the
 721            skip connections consume, and the spatial shape after preprocessing.
 722        """
 723        # Reshape the inputs to the shape expected by the encoder
 724        # and normalize the inputs if normalization is part of the model.
 725        x, input_shape = self.preprocess(x)
 726
 727        encoder_outputs = self.encoder(x)
 728
 729        if isinstance(encoder_outputs[-1], list):
 730            # `encoder_outputs` can be arranged in only two forms:
 731            #   - either we only return the image embeddings
 732            #   - or, we return the image embeddings and the "list" of global attention layers
 733            z12, from_encoder = encoder_outputs
 734        else:
 735            z12, from_encoder = encoder_outputs, None
 736
 737        return (z12, from_encoder, x), input_shape
 738
 739    def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor:
 740        """Run the convolutional decoder on the encoder features.
 741
 742        Args:
 743            features: The tuple returned by `encode`.
 744            input_shape: The spatial shape (Y, X) after preprocessing.
 745            original_shape: The spatial shape (Y, X) of the original input.
 746
 747        Returns:
 748            The UNETR output, resized to `original_shape`.
 749        """
 750        z12, from_encoder, x = features
 751
 752        if self.use_skip_connection:
 753            from_encoder = from_encoder[::-1]
 754            z9 = self.deconv1(from_encoder[0])
 755            z6 = self.deconv2(from_encoder[1])
 756            z3 = self.deconv3(from_encoder[2])
 757            z0 = self.deconv4(x)
 758
 759        else:
 760            z9 = self.deconv1(z12)
 761            z6 = self.deconv2(z9)
 762            z3 = self.deconv3(z6)
 763            z0 = self.deconv4(z3)
 764
 765        updated_from_encoder = [z9, z6, z3]
 766
 767        x = self.base(z12)
 768        x = self.decoder(x, encoder_inputs=updated_from_encoder)
 769        x = self.deconv_out(x)
 770
 771        x = torch.cat([x, z0], dim=1)
 772        x = self.decoder_head(x)
 773
 774        x = self.out_conv(x)
 775        if self.final_activation is not None:
 776            x = self.final_activation(x)
 777
 778        return self.postprocess_masks(x, input_shape, original_shape)
 779
 780
 781class UNETR2D(UNETR):
 782    """A two-dimensional UNet Transformer using a vision transformer as encoder and a convolutional decoder.
 783    """
 784    pass
 785
 786
 787class UNETR3D(UNETRBase):
 788    """A three dimensional UNet Transformer using a vision transformer as encoder and a convolutional decoder.
 789    """
 790    def __init__(
 791        self,
 792        img_size: int = 1024,
 793        backbone: Literal[
 794            "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3", "torchvision"
 795        ] = "sam",
 796        encoder: Optional[Union[nn.Module, str]] = "hvit_b",
 797        decoder: Optional[nn.Module] = None,
 798        out_channels: int = 1,
 799        use_sam_stats: bool = False,
 800        use_mae_stats: bool = False,
 801        use_dino_stats: bool = False,
 802        use_imagenet_stats: bool = False,
 803        resize_input: bool = True,
 804        encoder_checkpoint: Optional[Union[str, OrderedDict]] = None,
 805        final_activation: Optional[Union[str, nn.Module]] = None,
 806        use_skip_connection: bool = False,
 807        embed_dim: Optional[int] = None,
 808        use_conv_transpose: bool = False,
 809        use_strip_pooling: bool = True,
 810        perform_range_checks: bool = True,
 811        **kwargs
 812    ):
 813        if use_skip_connection:
 814            raise NotImplementedError("The framework cannot handle skip connections atm.")
 815        if use_conv_transpose:
 816            raise NotImplementedError("It's not enabled to switch between interpolation and transposed convolutions.")
 817
 818        # Sort the `embed_dim` out
 819        embed_dim = 256 if embed_dim is None else embed_dim
 820
 821        super().__init__(
 822            img_size=img_size,
 823            backbone=backbone,
 824            encoder=encoder,
 825            decoder=decoder,
 826            out_channels=out_channels,
 827            use_sam_stats=use_sam_stats,
 828            use_mae_stats=use_mae_stats,
 829            use_dino_stats=use_dino_stats,
 830            use_imagenet_stats=use_imagenet_stats,
 831            resize_input=resize_input,
 832            encoder_checkpoint=encoder_checkpoint,
 833            final_activation=final_activation,
 834            use_skip_connection=use_skip_connection,
 835            embed_dim=embed_dim,
 836            use_conv_transpose=use_conv_transpose,
 837            perform_range_checks=perform_range_checks,
 838            **kwargs,
 839        )
 840
 841        # The 3d convolutional decoder.
 842        # First, get the important parameters for the decoder.
 843        depth = 3
 844        gain = 2
 845        features_decoder = [self.initial_features * gain ** i for i in range(depth + 1)][::-1]
 846        scale_factors = [1, 2, 2]
 847        self.out_channels = out_channels
 848
 849        # The mapping blocks.
 850        self.deconv1 = Deconv3DBlock(
 851            in_channels=embed_dim,
 852            out_channels=features_decoder[0],
 853            scale_factor=scale_factors,
 854            use_strip_pooling=use_strip_pooling,
 855        )
 856        self.deconv2 = Deconv3DBlock(
 857            in_channels=features_decoder[0],
 858            out_channels=features_decoder[1],
 859            scale_factor=scale_factors,
 860            use_strip_pooling=use_strip_pooling,
 861        )
 862        self.deconv3 = Deconv3DBlock(
 863            in_channels=features_decoder[1],
 864            out_channels=features_decoder[2],
 865            scale_factor=scale_factors,
 866            use_strip_pooling=use_strip_pooling,
 867        )
 868        self.deconv4 = Deconv3DBlock(
 869            in_channels=features_decoder[2],
 870            out_channels=features_decoder[3],
 871            scale_factor=scale_factors,
 872            use_strip_pooling=use_strip_pooling,
 873        )
 874
 875        # The core decoder block.
 876        self.decoder = decoder or Decoder(
 877            features=features_decoder,
 878            scale_factors=[scale_factors] * depth,
 879            conv_block_impl=partial(ConvBlock3dWithStrip, use_strip_pooling=use_strip_pooling),
 880            sampler_impl=Upsampler3d,
 881        )
 882
 883        # And the final upsampler to match the expected dimensions.
 884        self.deconv_out = Deconv3DBlock(  # NOTE: changed `end_up` to `deconv_out`
 885            in_channels=features_decoder[-1],
 886            out_channels=features_decoder[-1],
 887            scale_factor=scale_factors,
 888            use_strip_pooling=use_strip_pooling,
 889        )
 890
 891        # Additional conjunction blocks.
 892        self.base = ConvBlock3dWithStrip(
 893            in_channels=embed_dim,
 894            out_channels=features_decoder[0],
 895            use_strip_pooling=use_strip_pooling,
 896        )
 897
 898        # And the output layers.
 899        self.decoder_head = ConvBlock3dWithStrip(
 900            in_channels=2 * features_decoder[-1],
 901            out_channels=features_decoder[-1],
 902            use_strip_pooling=use_strip_pooling,
 903        )
 904        self.out_conv = nn.Conv3d(features_decoder[-1], out_channels, 1)
 905
 906    def encode(self, x: torch.Tensor):
 907        """Preprocess the input and run the image encoder on every z-slice.
 908
 909        Args:
 910            x: Inputs of expected shape (B, C, Z, Y, X), where Z considers flexible inputs.
 911
 912        Returns:
 913            The encoder features of shape (B, D, Z, Y', X') and the spatial shape after preprocessing.
 914        """
 915        Z = x.shape[2]
 916        x, input_shape = self.preprocess(x)
 917        features = torch.stack([self.encoder(x[:, :, i])[0] for i in range(Z)], dim=2)
 918        return features, input_shape
 919
 920    def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor:
 921        """Run the convolutional decoder on the encoder features.
 922
 923        Args:
 924            features: Encoder features of shape (B, D, Z, Y', X'), see `encode`.
 925            input_shape: The spatial shape (Z, Y, X) after preprocessing.
 926            original_shape: The spatial shape (Z, Y, X) of the original input.
 927
 928        Returns:
 929            The UNETR output, resized to `original_shape`.
 930        """
 931        # Prepare the counterparts for the decoder.
 932        # NOTE: The section below is sequential, there's no skip connections atm.
 933        z9 = self.deconv1(features)
 934        z6 = self.deconv2(z9)
 935        z3 = self.deconv3(z6)
 936        z0 = self.deconv4(z3)
 937
 938        updated_from_encoder = [z9, z6, z3]
 939
 940        # Align the features through the base block.
 941        x = self.base(features)
 942        # Run the decoder
 943        x = self.decoder(x, encoder_inputs=updated_from_encoder)
 944        x = self.deconv_out(x)  # NOTE before `end_up`
 945
 946        # And the final output head.
 947        x = torch.cat([x, z0], dim=1)
 948        x = self.decoder_head(x)
 949        x = self.out_conv(x)
 950        if self.final_activation is not None:
 951            x = self.final_activation(x)
 952
 953        # Postprocess the output back to original size.
 954        return self.postprocess_masks(x, input_shape, original_shape)
 955
 956#
 957#  ADDITIONAL FUNCTIONALITIES
 958#
 959
 960
 961def _strip_pooling_layers(enabled, channels) -> nn.Module:
 962    return DepthStripPooling(channels) if enabled else nn.Identity()
 963
 964
 965class DepthStripPooling(nn.Module):
 966    """@private
 967    """
 968    def __init__(self, channels: int, reduction: int = 4):
 969        """Block for strip pooling along the depth dimension (only).
 970
 971        eg. for 3D (Z > 1) - it aggregates global context across depth by adaptive avg pooling
 972        to Z=1, and then passes through a small 1x1x1 MLP, then broadcasts it back to Z to
 973        modulate the original features (using a gated residual).
 974
 975        For 2D (Z == 1): returns input unchanged (no-op).
 976
 977        Args:
 978            channels: The output channels.
 979            reduction: The reduction of the hidden layers.
 980        """
 981        super().__init__()
 982        hidden = max(1, channels // reduction)
 983        self.conv1 = nn.Conv3d(channels, hidden, kernel_size=1)
 984        self.bn1 = nn.BatchNorm3d(hidden)
 985        self.relu = nn.ReLU(inplace=True)
 986        self.conv2 = nn.Conv3d(hidden, channels, kernel_size=1)
 987
 988    def forward(self, x: torch.Tensor) -> torch.Tensor:
 989        if x.dim() != 5:
 990            raise ValueError(f"DepthStripPooling expects 5D tensors as input, got '{x.shape}'.")
 991
 992        B, C, Z, H, W = x.shape
 993        if Z == 1:  # i.e. always the case of all 2d.
 994            return x  # We simply do nothing there.
 995
 996        # We pool only along the depth dimension: i.e. target shape (B, C, 1, H, W).
 997        # A plain mean over Z is the same operation as adaptive_avg_pool3d to (1, H, W), but its
 998        # reduction kernel is several times faster at full resolution.
 999        feat = x.mean(dim=2, keepdim=True)
1000        feat = self.conv1(feat)
1001        feat = self.bn1(feat)
1002        feat = self.relu(feat)
1003        feat = self.conv2(feat)
1004        gate = torch.sigmoid(feat).expand(B, C, Z, H, W)  # Broadcast the collapsed depth context back to all slices
1005
1006        # Gated residual fusion
1007        return x * gate + x
1008
1009
1010class Deconv3DBlock(nn.Module):
1011    """@private
1012    """
1013    def __init__(
1014        self,
1015        scale_factor,
1016        in_channels,
1017        out_channels,
1018        kernel_size=3,
1019        anisotropic_kernel=True,
1020        use_strip_pooling=True,
1021    ):
1022        super().__init__()
1023        conv_block_kwargs = {
1024            "in_channels": out_channels,
1025            "out_channels": out_channels,
1026            "kernel_size": kernel_size,
1027            "padding": ((kernel_size - 1) // 2),
1028        }
1029        if anisotropic_kernel:
1030            conv_block_kwargs = _update_conv_kwargs(conv_block_kwargs, scale_factor)
1031
1032        self.block = nn.Sequential(
1033            Upsampler3d(scale_factor, in_channels, out_channels),
1034            nn.Conv3d(**conv_block_kwargs),
1035            nn.BatchNorm3d(out_channels),
1036            nn.ReLU(True),
1037            _strip_pooling_layers(enabled=use_strip_pooling, channels=out_channels),
1038        )
1039
1040    def forward(self, x):
1041        return self.block(x)
1042
1043
1044class ConvBlock3dWithStrip(nn.Module):
1045    """@private
1046    """
1047    def __init__(
1048        self, in_channels: int, out_channels: int, use_strip_pooling: bool = True, **kwargs
1049    ):
1050        super().__init__()
1051        self.block = nn.Sequential(
1052            ConvBlock3d(in_channels, out_channels, **kwargs),
1053            _strip_pooling_layers(enabled=use_strip_pooling, channels=out_channels),
1054        )
1055
1056    def forward(self, x):
1057        return self.block(x)
1058
1059
1060class SingleDeconv2DBlock(nn.Module):
1061    """@private
1062    """
1063    def __init__(self, scale_factor, in_channels, out_channels):
1064        super().__init__()
1065        self.block = nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2, stride=2, padding=0, output_padding=0)
1066
1067    def forward(self, x):
1068        return self.block(x)
1069
1070
1071class SingleConv2DBlock(nn.Module):
1072    """@private
1073    """
1074    def __init__(self, in_channels, out_channels, kernel_size):
1075        super().__init__()
1076        self.block = nn.Conv2d(
1077            in_channels, out_channels, kernel_size=kernel_size, stride=1, padding=((kernel_size - 1) // 2)
1078        )
1079
1080    def forward(self, x):
1081        return self.block(x)
1082
1083
1084class Conv2DBlock(nn.Module):
1085    """@private
1086    """
1087    def __init__(self, in_channels, out_channels, kernel_size=3):
1088        super().__init__()
1089        self.block = nn.Sequential(
1090            SingleConv2DBlock(in_channels, out_channels, kernel_size),
1091            nn.BatchNorm2d(out_channels),
1092            nn.ReLU(True)
1093        )
1094
1095    def forward(self, x):
1096        return self.block(x)
1097
1098
1099class Deconv2DBlock(nn.Module):
1100    """@private
1101    """
1102    def __init__(self, in_channels, out_channels, kernel_size=3, use_conv_transpose=True):
1103        super().__init__()
1104        _upsampler = SingleDeconv2DBlock if use_conv_transpose else Upsampler2d
1105        self.block = nn.Sequential(
1106            _upsampler(scale_factor=2, in_channels=in_channels, out_channels=out_channels),
1107            SingleConv2DBlock(out_channels, out_channels, kernel_size),
1108            nn.BatchNorm2d(out_channels),
1109            nn.ReLU(True)
1110        )
1111
1112    def forward(self, x):
1113        return self.block(x)
class UNETRBase(torch.nn.modules.module.Module):
 81class UNETRBase(nn.Module):
 82    """Base class for implementing a UNETR.
 83
 84    Args:
 85        img_size: The size of the input for the image encoder. Input images will be resized to match this size.
 86        backbone: The name of the vision transformer implementation.
 87            One of "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3"
 88            (see all combinations below)
 89        encoder: The vision transformer. Can either be a name, such as "vit_b"
 90            (see all combinations for this below) or a torch module.
 91        decoder: The convolutional decoder.
 92        out_channels: The number of output channels of the UNETR.
 93        use_sam_stats: Whether to normalize the input data with the statistics of the
 94            pretrained SAM / SAM2 / SAM3 model.
 95        use_dino_stats: Whether to normalize the input data with the statistics of the
 96            pretrained DINOv2 / DINOv3 model.
 97        use_imagenet_stats: Whether to normalize with standard ImageNet statistics, i.e.
 98            mean - (0.485, 0.456, 0.406) and std - (0.229, 0.224, 0.225), raw inputs between range [0, 1].
 99            Use this with the 'torchvision' backbone when loading pretrained weights.
100        use_mae_stats: Whether to normalize the input data with the statistics of the pretrained MAE model.
101        resize_input: Whether to resize the input images to match `img_size`.
102            By default, it resizes the inputs to match the `img_size`.
103        encoder_checkpoint: Checkpoint for initializing the vision transformer.
104            Can either be a filepath or an already loaded checkpoint.
105        final_activation: The activation to apply to the UNETR output.
106        use_skip_connection: Whether to use skip connections. By default, it uses skip connections.
107        embed_dim: The embedding dimensionality, corresponding to the output dimension of the vision transformer.
108        use_conv_transpose: Whether to use transposed convolutions instead of resampling for upsampling.
109            By default, it uses resampling for upsampling.
110        perform_range_checks: Whether to validate the input value range before normalization on each forward pass.
111            You can disable the checks to avoid GPU sync overhead during training when inputs are known to be correct.
112        initial_features: The number of features of the finest decoder level. The features per level are
113            'initial_features * gain ** i', so this scales the decoder parameters quadratically.
114
115        NOTE: The currently supported combinations of 'backbone' x 'encoder' are the following:
116
117        SAM_family_models:
118            - 'sam' x 'vit_b'
119            - 'sam' x 'vit_l'
120            - 'sam' x 'vit_h'
121            - 'sam2' x 'hvit_t'
122            - 'sam2' x 'hvit_s'
123            - 'sam2' x 'hvit_b'
124            - 'sam2' x 'hvit_l'
125            - 'sam3' x 'vit_pe'
126            - 'cellpose_sam' x 'vit_l'
127
128        DINO_family_models:
129            - 'dinov2' x 'vit_s'
130            - 'dinov2' x 'vit_b'
131            - 'dinov2' x 'vit_l'
132            - 'dinov2' x 'vit_g'
133            - 'dinov2' x 'vit_s_reg4'
134            - 'dinov2' x 'vit_b_reg4'
135            - 'dinov2' x 'vit_l_reg4'
136            - 'dinov2' x 'vit_g_reg4'
137            - 'dinov3' x 'vit_s'
138            - 'dinov3' x 'vit_s+'
139            - 'dinov3' x 'vit_b'
140            - 'dinov3' x 'vit_l'
141            - 'dinov3' x 'vit_l+'
142            - 'dinov3' x 'vit_h+'
143            - 'dinov3' x 'vit_7b'
144
145        MAE_family_models:
146            - 'mae' x 'vit_b'
147            - 'mae' x 'vit_l'
148            - 'mae' x 'vit_h'
149            - 'scalemae' x 'vit_b'
150            - 'scalemae' x 'vit_l'
151            - 'scalemae' x 'vit_h'
152
153        torchvision_models:
154            - 'torchvision' x 'vit_b_16'
155            - 'torchvision' x 'vit_b_32'
156            - 'torchvision' x 'vit_l_16'
157            - 'torchvision' x 'vit_l_32'
158            - 'torchvision' x 'vit_h_14'
159    """
160    def __init__(
161        self,
162        img_size: int = 1024,
163        backbone: Literal[
164            "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3", "torchvision"
165        ] = "sam",
166        encoder: Optional[Union[nn.Module, str]] = "vit_b",
167        decoder: Optional[nn.Module] = None,
168        out_channels: int = 1,
169        use_sam_stats: bool = False,
170        use_mae_stats: bool = False,
171        use_dino_stats: bool = False,
172        use_imagenet_stats: bool = False,
173        resize_input: bool = True,
174        encoder_checkpoint: Optional[Union[str, OrderedDict]] = None,
175        final_activation: Optional[Union[str, nn.Module]] = None,
176        use_skip_connection: bool = True,
177        embed_dim: Optional[int] = None,
178        use_conv_transpose: bool = False,
179        perform_range_checks: bool = True,
180        initial_features: int = 64,
181        **kwargs
182    ) -> None:
183        super().__init__()
184
185        self.img_size = img_size
186        self.use_sam_stats = use_sam_stats
187        self.use_mae_stats = use_mae_stats
188        self.use_dino_stats = use_dino_stats
189        self.use_imagenet_stats = use_imagenet_stats
190        self.use_skip_connection = use_skip_connection
191        self.resize_input = resize_input
192        self.perform_range_checks = perform_range_checks
193        self.use_conv_transpose = use_conv_transpose
194        self.initial_features = initial_features
195        self.backbone = backbone
196
197        if isinstance(encoder, str):  # e.g. "vit_b" / "hvit_b" / "vit_pe"
198            print(f"Using {encoder} from {backbone.upper()}")
199            self.encoder = get_vision_transformer(img_size=img_size, backbone=backbone, model=encoder, **kwargs)
200
201            if encoder_checkpoint is not None:
202                self._load_encoder_from_checkpoint(backbone=backbone, encoder=encoder, checkpoint=encoder_checkpoint)
203
204            if embed_dim is None:
205                embed_dim = self.encoder.embed_dim
206
207            # For SAM1 encoder, if 'apply_neck' is applied, the embedding dimension must change.
208            if hasattr(self.encoder, "apply_neck") and self.encoder.apply_neck:
209                embed_dim = self.encoder.neck[2].out_channels  # the value is 256
210
211        else:  # `nn.Module` ViT backbone
212            self.encoder = encoder
213
214            have_neck = False
215            for name, _ in self.encoder.named_parameters():
216                if name.startswith("neck"):
217                    have_neck = True
218
219            if embed_dim is None:
220                if have_neck:
221                    embed_dim = self.encoder.neck[2].out_channels  # the value is 256
222                else:
223                    embed_dim = self.encoder.patch_embed.proj.out_channels
224
225        self.embed_dim = embed_dim
226        self.final_activation = self._get_activation(final_activation)
227
228    def _load_encoder_from_checkpoint(self, backbone, encoder, checkpoint):
229        """Function to load pretrained weights to the image encoder.
230        """
231        if isinstance(checkpoint, str):
232            if backbone == "sam" and isinstance(encoder, str):
233                # If we have a SAM encoder, then we first try to load the full SAM Model
234                # (using micro_sam) and otherwise fall back on directly loading the encoder state
235                # from the checkpoint
236                try:
237                    _, model = get_sam_model(model_type=encoder, checkpoint_path=checkpoint, return_sam=True)
238                    encoder_state = model.image_encoder.state_dict()
239                except Exception:
240                    # Try loading the encoder state directly from a checkpoint.
241                    encoder_state = torch.load(checkpoint, weights_only=False)
242
243            elif backbone == "cellpose_sam" and isinstance(encoder, str):
244                # The architecture matches CellposeSAM exactly (same rel_pos sizes),
245                # so weights load directly without any interpolation.
246                encoder_state = torch.load(checkpoint, map_location="cpu", weights_only=False)
247                # Handle DataParallel/DistributedDataParallel prefix.
248                if any(k.startswith("module.") for k in encoder_state.keys()):
249                    encoder_state = OrderedDict(
250                        {k[len("module."):]: v for k, v in encoder_state.items()}
251                    )
252                # Extract encoder weights from CellposeSAM checkpoint format (strip 'encoder.' prefix).
253                if any(k.startswith("encoder.") for k in encoder_state.keys()):
254                    encoder_state = OrderedDict(
255                        {k[len("encoder."):]: v for k, v in encoder_state.items() if k.startswith("encoder.")}
256                    )
257
258            elif backbone == "sam2" and isinstance(encoder, str):
259                # If we have a SAM2 encoder, then we first try to load the full SAM2 Model.
260                # (using micro_sam2) and otherwise fall back on directly loading the encoder state
261                # from the checkpoint
262                try:
263                    model = get_sam2_model(model_type=encoder, checkpoint_path=checkpoint)
264                    encoder_state = model.image_encoder.state_dict()
265                except Exception:
266                    # Try loading the encoder state directly from a checkpoint.
267                    encoder_state = torch.load(checkpoint, weights_only=False)
268
269            elif backbone == "sam3" and isinstance(encoder, str):
270                # If we have a SAM3 encoder, then we first try to load the full SAM3 Model.
271                # (using micro_sam3) and otherwise fall back on directly loading the encoder state
272                # from the checkpoint
273                try:
274                    model = get_sam3_model(checkpoint_path=checkpoint)
275                    encoder_state = model.backbone.vision_backbone.state_dict()
276                    # Let's align loading the encoder weights with expected parameter names
277                    encoder_state = {
278                        k[len("trunk."):] if k.startswith("trunk.") else k: v for k, v in encoder_state.items()
279                    }
280                    # And drop the 'convs' and 'sam2_convs' - these seem like some upsampling blocks.
281                    encoder_state = {
282                        k: v for k, v in encoder_state.items()
283                        if not (k.startswith("convs.") or k.startswith("sam2_convs."))
284                    }
285                except Exception:
286                    # Try loading the encoder state directly from a checkpoint.
287                    encoder_state = torch.load(checkpoint, weights_only=False)
288
289            elif backbone == "mae":
290                # vit initialization hints from:
291                #     - https://github.com/facebookresearch/mae/blob/main/main_finetune.py#L233-L242
292                encoder_state = torch.load(checkpoint, weights_only=False)["model"]
293                encoder_state = OrderedDict({
294                    k: v for k, v in encoder_state.items() if (k != "mask_token" and not k.startswith("decoder"))
295                })
296                # Let's remove the `head` from our current encoder (as the MAE pretrained don't expect it)
297                current_encoder_state = self.encoder.state_dict()
298                if ("head.weight" in current_encoder_state) and ("head.bias" in current_encoder_state):
299                    del self.encoder.head
300
301            elif backbone == "scalemae":
302                # Load the encoder state directly from a checkpoint.
303                encoder_state = torch.load(checkpoint)["model"]
304                encoder_state = OrderedDict({
305                    k: v for k, v in encoder_state.items()
306                    if not k.startswith(("mask_token", "decoder", "fcn", "fpn", "pos_embed"))
307                })
308
309                # Let's remove the `head` from our current encoder (as the MAE pretrained don't expect it)
310                current_encoder_state = self.encoder.state_dict()
311                if ("head.weight" in current_encoder_state) and ("head.bias" in current_encoder_state):
312                    del self.encoder.head
313
314                if "pos_embed" in current_encoder_state:  # NOTE: ScaleMAE uses 'pos. embeddings' in a diff. format.
315                    del self.encoder.pos_embed
316
317            elif backbone in ["dinov2", "dinov3"]:  # Load the encoder state directly from a checkpoint.
318                encoder_state = torch.load(checkpoint)
319
320            elif backbone == "torchvision":
321                encoder_state = torch.load(checkpoint, weights_only=False)
322
323            else:
324                raise ValueError(
325                    f"We don't support either the '{backbone}' backbone or the '{encoder}' model combination (or both)."
326                )
327
328        else:
329            encoder_state = checkpoint
330
331        if backbone == "torchvision":
332            if "state_dict" in encoder_state:
333                encoder_state = encoder_state["state_dict"]
334            encoder_state = {k: v for k, v in encoder_state.items() if not k.startswith("heads.")}
335
336        self.encoder.load_state_dict(encoder_state)
337
338    def _get_activation(self, activation):
339        return_activation = None
340        if activation is None:
341            return None
342        if isinstance(activation, nn.Module):
343            return activation
344        if isinstance(activation, str):
345            return_activation = getattr(nn, activation, None)
346        if return_activation is None:
347            raise ValueError(f"Invalid activation: {activation}")
348
349        return return_activation()
350
351    @staticmethod
352    def get_preprocess_shape(oldh: int, oldw: int, long_side_length: int) -> Tuple[int, int]:
353        """Compute the output size given input size and target long side length.
354
355        Args:
356            oldh: The input image height.
357            oldw: The input image width.
358            long_side_length: The longest side length for resizing.
359
360        Returns:
361            The new image height.
362            The new image width.
363        """
364        scale = long_side_length * 1.0 / max(oldh, oldw)
365        newh, neww = oldh * scale, oldw * scale
366        neww = int(neww + 0.5)
367        newh = int(newh + 0.5)
368        return (newh, neww)
369
370    def resize_longest_side(self, image: torch.Tensor) -> torch.Tensor:
371        """Resize the image so that the longest side has the correct length.
372
373        Expects batched images with shape BxCxHxW OR BxCxDxHxW and float format.
374
375        Args:
376            image: The input image.
377
378        Returns:
379            The resized image.
380        """
381        if image.ndim == 4:  # i.e. 2d image
382            target_size = self.get_preprocess_shape(image.shape[2], image.shape[3], self.encoder.img_size)
383            return F.interpolate(image, target_size, mode="bilinear", align_corners=False, antialias=True)
384        elif image.ndim == 5:  # i.e. 3d volume
385            B, C, Z, H, W = image.shape
386            target_size = self.get_preprocess_shape(H, W, self.img_size)
387            return F.interpolate(image, (Z, *target_size), mode="trilinear", align_corners=False)
388        else:
389            raise ValueError("Expected 4d or 5d inputs, got", image.shape)
390
391    def _as_stats(self, mean, std, device, dtype, is_3d: bool):
392        """@private
393        """
394        return _as_stats(mean, std, device, dtype, is_3d)
395
396    def _check_input_normalization_range(self, x: torch.Tensor, expected_range: Optional[Tuple[float, float]]) -> None:
397        """@private
398        """
399        _check_input_normalization_range(x, expected_range)
400
401    def encode(self, x: torch.Tensor):
402        """Preprocess the input and run the image encoder.
403
404        Args:
405            x: The input tensor.
406
407        Returns:
408            The encoder features to pass to `decode` and the spatial shape after preprocessing.
409        """
410        raise NotImplementedError
411
412    def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor:
413        """Run the convolutional decoder on the encoder features.
414
415        Args:
416            features: The encoder features returned by `encode`.
417            input_shape: The spatial shape after preprocessing, returned by `encode`.
418            original_shape: The spatial shape of the original input.
419
420        Returns:
421            The UNETR output, resized to `original_shape`.
422        """
423        raise NotImplementedError
424
425    def forward(self, x: torch.Tensor) -> torch.Tensor:
426        """Apply the UNETR to the input data.
427
428        Args:
429            x: The input tensor.
430
431        Returns:
432            The UNETR output.
433        """
434        features, input_shape = self.encode(x)
435        return self.decode(features, input_shape, tuple(x.shape[2:]))
436
437    def preprocess(self, x: torch.Tensor) -> torch.Tensor:
438        """@private
439        """
440        return preprocess_vit_inputs(
441            x,
442            use_sam_stats=self.use_sam_stats,
443            backbone=self.backbone,
444            use_mae_stats=self.use_mae_stats,
445            use_dino_stats=self.use_dino_stats,
446            use_imagenet_stats=self.use_imagenet_stats,
447            resize_input=self.resize_input,
448            img_size=self.img_size,
449            encoder_img_size=self.encoder.img_size,
450            perform_range_checks=self.perform_range_checks,
451        )
452
453    def postprocess_masks(
454        self, masks: torch.Tensor, input_size: Tuple[int, ...], original_size: Tuple[int, ...],
455    ) -> torch.Tensor:
456        """@private
457        """
458        if masks.ndim == 4:  # i.e. 2d labels
459            masks = F.interpolate(
460                masks,
461                (self.encoder.img_size, self.encoder.img_size),
462                mode="bilinear",
463                align_corners=False,
464            )
465            masks = masks[..., : input_size[0], : input_size[1]]
466            masks = F.interpolate(masks, original_size, mode="bilinear", align_corners=False)
467
468        elif masks.ndim == 5:  # i.e. 3d volumetric labels
469            masks = F.interpolate(
470                masks,
471                (input_size[0], self.img_size, self.img_size),
472                mode="trilinear",
473                align_corners=False,
474            )
475            masks = masks[..., :input_size[0], :input_size[1], :input_size[2]]
476            masks = F.interpolate(masks, original_size, mode="trilinear", align_corners=False)
477
478        else:
479            raise ValueError("Expected 4d or 5d labels, got", masks.shape)
480
481        return masks

Base class for implementing a UNETR.

Arguments:
  • img_size: The size of the input for the image encoder. Input images will be resized to match this size.
  • backbone: The name of the vision transformer implementation. One of "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3" (see all combinations below)
  • encoder: The vision transformer. Can either be a name, such as "vit_b" (see all combinations for this below) or a torch module.
  • decoder: The convolutional decoder.
  • out_channels: The number of output channels of the UNETR.
  • use_sam_stats: Whether to normalize the input data with the statistics of the pretrained SAM / SAM2 / SAM3 model.
  • use_dino_stats: Whether to normalize the input data with the statistics of the pretrained DINOv2 / DINOv3 model.
  • use_imagenet_stats: Whether to normalize with standard ImageNet statistics, i.e. mean - (0.485, 0.456, 0.406) and std - (0.229, 0.224, 0.225), raw inputs between range [0, 1]. Use this with the 'torchvision' backbone when loading pretrained weights.
  • use_mae_stats: Whether to normalize the input data with the statistics of the pretrained MAE model.
  • resize_input: Whether to resize the input images to match img_size. By default, it resizes the inputs to match the img_size.
  • encoder_checkpoint: Checkpoint for initializing the vision transformer. Can either be a filepath or an already loaded checkpoint.
  • final_activation: The activation to apply to the UNETR output.
  • use_skip_connection: Whether to use skip connections. By default, it uses skip connections.
  • embed_dim: The embedding dimensionality, corresponding to the output dimension of the vision transformer.
  • use_conv_transpose: Whether to use transposed convolutions instead of resampling for upsampling. By default, it uses resampling for upsampling.
  • perform_range_checks: Whether to validate the input value range before normalization on each forward pass. You can disable the checks to avoid GPU sync overhead during training when inputs are known to be correct.
  • initial_features: The number of features of the finest decoder level. The features per level are 'initial_features * gain ** i', so this scales the decoder parameters quadratically.
  • NOTE: The currently supported combinations of 'backbone' x 'encoder' are the following:
  • SAM_family_models: - 'sam' x 'vit_b'
    • 'sam' x 'vit_l'
    • 'sam' x 'vit_h'
    • 'sam2' x 'hvit_t'
    • 'sam2' x 'hvit_s'
    • 'sam2' x 'hvit_b'
    • 'sam2' x 'hvit_l'
    • 'sam3' x 'vit_pe'
    • 'cellpose_sam' x 'vit_l'
  • DINO_family_models: - 'dinov2' x 'vit_s'
    • 'dinov2' x 'vit_b'
    • 'dinov2' x 'vit_l'
    • 'dinov2' x 'vit_g'
    • 'dinov2' x 'vit_s_reg4'
    • 'dinov2' x 'vit_b_reg4'
    • 'dinov2' x 'vit_l_reg4'
    • 'dinov2' x 'vit_g_reg4'
    • 'dinov3' x 'vit_s'
    • 'dinov3' x 'vit_s+'
    • 'dinov3' x 'vit_b'
    • 'dinov3' x 'vit_l'
    • 'dinov3' x 'vit_l+'
    • 'dinov3' x 'vit_h+'
    • 'dinov3' x 'vit_7b'
  • MAE_family_models: - 'mae' x 'vit_b'
    • 'mae' x 'vit_l'
    • 'mae' x 'vit_h'
    • 'scalemae' x 'vit_b'
    • 'scalemae' x 'vit_l'
    • 'scalemae' x 'vit_h'
  • torchvision_models: - 'torchvision' x 'vit_b_16'
    • 'torchvision' x 'vit_b_32'
    • 'torchvision' x 'vit_l_16'
    • 'torchvision' x 'vit_l_32'
    • 'torchvision' x 'vit_h_14'
UNETRBase( img_size: int = 1024, backbone: Literal['sam', 'sam2', 'sam3', 'cellpose_sam', 'mae', 'scalemae', 'dinov2', 'dinov3', 'torchvision'] = 'sam', encoder: Union[torch.nn.modules.module.Module, str, NoneType] = 'vit_b', decoder: Optional[torch.nn.modules.module.Module] = None, out_channels: int = 1, use_sam_stats: bool = False, use_mae_stats: bool = False, use_dino_stats: bool = False, use_imagenet_stats: bool = False, resize_input: bool = True, encoder_checkpoint: Union[str, collections.OrderedDict, NoneType] = None, final_activation: Union[torch.nn.modules.module.Module, str, NoneType] = None, use_skip_connection: bool = True, embed_dim: Optional[int] = None, use_conv_transpose: bool = False, perform_range_checks: bool = True, initial_features: int = 64, **kwargs)
160    def __init__(
161        self,
162        img_size: int = 1024,
163        backbone: Literal[
164            "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3", "torchvision"
165        ] = "sam",
166        encoder: Optional[Union[nn.Module, str]] = "vit_b",
167        decoder: Optional[nn.Module] = None,
168        out_channels: int = 1,
169        use_sam_stats: bool = False,
170        use_mae_stats: bool = False,
171        use_dino_stats: bool = False,
172        use_imagenet_stats: bool = False,
173        resize_input: bool = True,
174        encoder_checkpoint: Optional[Union[str, OrderedDict]] = None,
175        final_activation: Optional[Union[str, nn.Module]] = None,
176        use_skip_connection: bool = True,
177        embed_dim: Optional[int] = None,
178        use_conv_transpose: bool = False,
179        perform_range_checks: bool = True,
180        initial_features: int = 64,
181        **kwargs
182    ) -> None:
183        super().__init__()
184
185        self.img_size = img_size
186        self.use_sam_stats = use_sam_stats
187        self.use_mae_stats = use_mae_stats
188        self.use_dino_stats = use_dino_stats
189        self.use_imagenet_stats = use_imagenet_stats
190        self.use_skip_connection = use_skip_connection
191        self.resize_input = resize_input
192        self.perform_range_checks = perform_range_checks
193        self.use_conv_transpose = use_conv_transpose
194        self.initial_features = initial_features
195        self.backbone = backbone
196
197        if isinstance(encoder, str):  # e.g. "vit_b" / "hvit_b" / "vit_pe"
198            print(f"Using {encoder} from {backbone.upper()}")
199            self.encoder = get_vision_transformer(img_size=img_size, backbone=backbone, model=encoder, **kwargs)
200
201            if encoder_checkpoint is not None:
202                self._load_encoder_from_checkpoint(backbone=backbone, encoder=encoder, checkpoint=encoder_checkpoint)
203
204            if embed_dim is None:
205                embed_dim = self.encoder.embed_dim
206
207            # For SAM1 encoder, if 'apply_neck' is applied, the embedding dimension must change.
208            if hasattr(self.encoder, "apply_neck") and self.encoder.apply_neck:
209                embed_dim = self.encoder.neck[2].out_channels  # the value is 256
210
211        else:  # `nn.Module` ViT backbone
212            self.encoder = encoder
213
214            have_neck = False
215            for name, _ in self.encoder.named_parameters():
216                if name.startswith("neck"):
217                    have_neck = True
218
219            if embed_dim is None:
220                if have_neck:
221                    embed_dim = self.encoder.neck[2].out_channels  # the value is 256
222                else:
223                    embed_dim = self.encoder.patch_embed.proj.out_channels
224
225        self.embed_dim = embed_dim
226        self.final_activation = self._get_activation(final_activation)

Initialize internal Module state, shared by both nn.Module and ScriptModule.

img_size
use_sam_stats
use_mae_stats
use_dino_stats
use_imagenet_stats
use_skip_connection
resize_input
perform_range_checks
use_conv_transpose
initial_features
backbone
embed_dim
final_activation
@staticmethod
def get_preprocess_shape(oldh: int, oldw: int, long_side_length: int) -> Tuple[int, int]:
351    @staticmethod
352    def get_preprocess_shape(oldh: int, oldw: int, long_side_length: int) -> Tuple[int, int]:
353        """Compute the output size given input size and target long side length.
354
355        Args:
356            oldh: The input image height.
357            oldw: The input image width.
358            long_side_length: The longest side length for resizing.
359
360        Returns:
361            The new image height.
362            The new image width.
363        """
364        scale = long_side_length * 1.0 / max(oldh, oldw)
365        newh, neww = oldh * scale, oldw * scale
366        neww = int(neww + 0.5)
367        newh = int(newh + 0.5)
368        return (newh, neww)

Compute the output size given input size and target long side length.

Arguments:
  • oldh: The input image height.
  • oldw: The input image width.
  • long_side_length: The longest side length for resizing.
Returns:

The new image height. The new image width.

def resize_longest_side(self, image: torch.Tensor) -> torch.Tensor:
370    def resize_longest_side(self, image: torch.Tensor) -> torch.Tensor:
371        """Resize the image so that the longest side has the correct length.
372
373        Expects batched images with shape BxCxHxW OR BxCxDxHxW and float format.
374
375        Args:
376            image: The input image.
377
378        Returns:
379            The resized image.
380        """
381        if image.ndim == 4:  # i.e. 2d image
382            target_size = self.get_preprocess_shape(image.shape[2], image.shape[3], self.encoder.img_size)
383            return F.interpolate(image, target_size, mode="bilinear", align_corners=False, antialias=True)
384        elif image.ndim == 5:  # i.e. 3d volume
385            B, C, Z, H, W = image.shape
386            target_size = self.get_preprocess_shape(H, W, self.img_size)
387            return F.interpolate(image, (Z, *target_size), mode="trilinear", align_corners=False)
388        else:
389            raise ValueError("Expected 4d or 5d inputs, got", image.shape)

Resize the image so that the longest side has the correct length.

Expects batched images with shape BxCxHxW OR BxCxDxHxW and float format.

Arguments:
  • image: The input image.
Returns:

The resized image.

def encode(self, x: torch.Tensor):
401    def encode(self, x: torch.Tensor):
402        """Preprocess the input and run the image encoder.
403
404        Args:
405            x: The input tensor.
406
407        Returns:
408            The encoder features to pass to `decode` and the spatial shape after preprocessing.
409        """
410        raise NotImplementedError

Preprocess the input and run the image encoder.

Arguments:
  • x: The input tensor.
Returns:

The encoder features to pass to decode and the spatial shape after preprocessing.

def decode( self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor:
412    def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor:
413        """Run the convolutional decoder on the encoder features.
414
415        Args:
416            features: The encoder features returned by `encode`.
417            input_shape: The spatial shape after preprocessing, returned by `encode`.
418            original_shape: The spatial shape of the original input.
419
420        Returns:
421            The UNETR output, resized to `original_shape`.
422        """
423        raise NotImplementedError

Run the convolutional decoder on the encoder features.

Arguments:
  • features: The encoder features returned by encode.
  • input_shape: The spatial shape after preprocessing, returned by encode.
  • original_shape: The spatial shape of the original input.
Returns:

The UNETR output, resized to original_shape.

def forward(self, x: torch.Tensor) -> torch.Tensor:
425    def forward(self, x: torch.Tensor) -> torch.Tensor:
426        """Apply the UNETR to the input data.
427
428        Args:
429            x: The input tensor.
430
431        Returns:
432            The UNETR output.
433        """
434        features, input_shape = self.encode(x)
435        return self.decode(features, input_shape, tuple(x.shape[2:]))

Apply the UNETR to the input data.

Arguments:
  • x: The input tensor.
Returns:

The UNETR output.

def preprocess_vit_inputs( x: torch.Tensor, use_sam_stats: bool = False, backbone: str = 'sam', use_mae_stats: bool = False, use_dino_stats: bool = False, use_imagenet_stats: bool = False, resize_input: bool = True, img_size: int = 1024, encoder_img_size: int = 1024, perform_range_checks: bool = True) -> Tuple[torch.Tensor, Tuple]:
484def preprocess_vit_inputs(
485    x: torch.Tensor,
486    use_sam_stats: bool = False,
487    backbone: str = "sam",
488    use_mae_stats: bool = False,
489    use_dino_stats: bool = False,
490    use_imagenet_stats: bool = False,
491    resize_input: bool = True,
492    img_size: int = 1024,
493    encoder_img_size: int = 1024,
494    perform_range_checks: bool = True,
495) -> Tuple[torch.Tensor, Tuple]:
496    """Preprocess inputs for ViT-backbones in UNETR models.
497
498    Handles normalization stat selection, input range validation, optional resizing to the longest side,
499    and padding to `encoder_img_size`. Can be used as a standalone function without a model instance.
500
501    Args:
502        x: Input tensor of shape (B, C, H, W) for 2D or (B, C, Z, H, W) for 3D.
503        use_sam_stats: Whether to normalize with SAM/SAM2/SAM3 backbone statistics.
504        backbone: The backbone name - controls which SAM stats are used when `use_sam_stats=True`.
505        use_mae_stats: Whether to normalize with MAE statistics.
506        use_dino_stats: Whether to normalize with DINOv2/DINOv3 statistics.
507        use_imagenet_stats: Whether to normalize with standard ImageNet statistics
508            (mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), range [0, 1]).
509            Use this for torchvision pretrained backbones.
510        resize_input: Whether to resize the input to the longest side before padding.
511        img_size: The model image size, used for 3D resize.
512        encoder_img_size: The encoder image size, used for 2D resize and padding.
513        perform_range_checks: Whether to validate the expected input value range before normalization.
514            You can disable the checks to avoid GPU sync overhead during training when inputs are known to be correct.
515
516    Returns:
517        The preprocessed tensor and the spatial shape after resizing (before padding).
518    """
519    is_3d = (x.ndim == 5)
520    device, dtype = x.device, x.dtype
521    mean, std = (0.0, 0.0, 0.0), (1.0, 1.0, 1.0)
522    expected_range = None
523    unit_scale_max = None
524
525    if use_sam_stats:
526        if backbone == "sam2":
527            mean, std = (0.485, 0.456, 0.406), (0.229, 0.224, 0.225)
528            expected_range = (0.0, 1.0)
529        elif backbone == "sam3":
530            mean, std = (0.5, 0.5, 0.5), (0.5, 0.5, 0.5)
531            expected_range = (0.0, 1.0)
532        else:  # sam1 / default
533            mean, std = (123.675, 116.28, 103.53), (58.395, 57.12, 57.375)
534            expected_range = (0.0, 255.0)
535            unit_scale_max = 1.0
536    elif use_mae_stats:  # TODO: add mean std from mae / scalemae experiments (or open up arguments for this)
537        raise NotImplementedError
538    elif use_dino_stats or use_imagenet_stats:
539        mean, std = (0.485, 0.456, 0.406), (0.229, 0.224, 0.225)
540        expected_range = (0.0, 1.0)
541    else:
542        mean, std = (0.0, 0.0, 0.0), (1.0, 1.0, 1.0)
543        expected_range = None
544
545    if perform_range_checks:
546        _check_input_normalization_range(x, expected_range, unit_scale_max)
547    pixel_mean, pixel_std = _as_stats(mean, std, device=device, dtype=dtype, is_3d=is_3d)
548
549    if resize_input:
550        if x.ndim == 4:
551            target_size = UNETRBase.get_preprocess_shape(x.shape[2], x.shape[3], encoder_img_size)
552            x = F.interpolate(x, target_size, mode="bilinear", align_corners=False, antialias=True)
553        elif x.ndim == 5:
554            B, C, Z, H, W = x.shape
555            target_size = UNETRBase.get_preprocess_shape(H, W, img_size)
556            x = F.interpolate(x, (Z, *target_size), mode="trilinear", align_corners=False)
557
558    input_shape = x.shape[-3:] if is_3d else x.shape[-2:]
559
560    x = (x - pixel_mean) / pixel_std
561    h, w = x.shape[-2:]
562    padh = encoder_img_size - h
563    padw = encoder_img_size - w
564
565    if is_3d:
566        x = F.pad(x, (0, padw, 0, padh, 0, 0))
567    else:
568        x = F.pad(x, (0, padw, 0, padh))
569
570    return x, input_shape

Preprocess inputs for ViT-backbones in UNETR models.

Handles normalization stat selection, input range validation, optional resizing to the longest side, and padding to encoder_img_size. Can be used as a standalone function without a model instance.

Arguments:
  • x: Input tensor of shape (B, C, H, W) for 2D or (B, C, Z, H, W) for 3D.
  • use_sam_stats: Whether to normalize with SAM/SAM2/SAM3 backbone statistics.
  • backbone: The backbone name - controls which SAM stats are used when use_sam_stats=True.
  • use_mae_stats: Whether to normalize with MAE statistics.
  • use_dino_stats: Whether to normalize with DINOv2/DINOv3 statistics.
  • use_imagenet_stats: Whether to normalize with standard ImageNet statistics (mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), range [0, 1]). Use this for torchvision pretrained backbones.
  • resize_input: Whether to resize the input to the longest side before padding.
  • img_size: The model image size, used for 3D resize.
  • encoder_img_size: The encoder image size, used for 2D resize and padding.
  • perform_range_checks: Whether to validate the expected input value range before normalization. You can disable the checks to avoid GPU sync overhead during training when inputs are known to be correct.
Returns:

The preprocessed tensor and the spatial shape after resizing (before padding).

class UNETR(UNETRBase):
573class UNETR(UNETRBase):
574    """A (2d-only) UNet Transformer using a vision transformer as encoder and a convolutional decoder.
575    """
576    def __init__(
577        self,
578        img_size: int = 1024,
579        backbone: Literal[
580            "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3", "torchvision"
581        ] = "sam",
582        encoder: Optional[Union[nn.Module, str]] = "vit_b",
583        decoder: Optional[nn.Module] = None,
584        out_channels: int = 1,
585        use_sam_stats: bool = False,
586        use_mae_stats: bool = False,
587        use_dino_stats: bool = False,
588        use_imagenet_stats: bool = False,
589        resize_input: bool = True,
590        encoder_checkpoint: Optional[Union[str, OrderedDict]] = None,
591        final_activation: Optional[Union[str, nn.Module]] = None,
592        use_skip_connection: bool = True,
593        embed_dim: Optional[int] = None,
594        use_conv_transpose: bool = False,
595        perform_range_checks: bool = True,
596        **kwargs
597    ) -> None:
598
599        super().__init__(
600            img_size=img_size,
601            backbone=backbone,
602            encoder=encoder,
603            decoder=decoder,
604            out_channels=out_channels,
605            use_sam_stats=use_sam_stats,
606            use_mae_stats=use_mae_stats,
607            use_dino_stats=use_dino_stats,
608            use_imagenet_stats=use_imagenet_stats,
609            resize_input=resize_input,
610            encoder_checkpoint=encoder_checkpoint,
611            final_activation=final_activation,
612            use_skip_connection=use_skip_connection,
613            embed_dim=embed_dim,
614            use_conv_transpose=use_conv_transpose,
615            perform_range_checks=perform_range_checks,
616            **kwargs,
617        )
618
619        encoder = self.encoder
620
621        if backbone == "sam2" and hasattr(encoder, "trunk"):
622            in_chans = encoder.trunk.patch_embed.proj.in_channels
623        elif hasattr(encoder, "in_chans"):
624            in_chans = encoder.in_chans
625        else:  # `nn.Module` ViT backbone.
626            try:
627                in_chans = encoder.patch_embed.proj.in_channels
628            except AttributeError:  # for getting the input channels while using 'vit_t' from MobileSam
629                in_chans = encoder.patch_embed.seq[0].c.in_channels
630
631        # parameters for the decoder network
632        depth = 3
633        gain = 2
634        features_decoder = [self.initial_features * gain ** i for i in range(depth + 1)][::-1]
635        scale_factors = depth * [2]
636        self.out_channels = out_channels
637
638        # choice of upsampler - to use (bilinear interpolation + conv) or conv transpose
639        _upsampler = SingleDeconv2DBlock if use_conv_transpose else Upsampler2d
640
641        self.decoder = decoder or Decoder(
642            features=features_decoder,
643            scale_factors=scale_factors[::-1],
644            conv_block_impl=ConvBlock2d,
645            sampler_impl=_upsampler,
646        )
647
648        if use_skip_connection:
649            self.deconv1 = Deconv2DBlock(
650                in_channels=self.embed_dim,
651                out_channels=features_decoder[0],
652                use_conv_transpose=use_conv_transpose,
653            )
654            self.deconv2 = nn.Sequential(
655                Deconv2DBlock(
656                    in_channels=self.embed_dim,
657                    out_channels=features_decoder[0],
658                    use_conv_transpose=use_conv_transpose,
659                ),
660                Deconv2DBlock(
661                    in_channels=features_decoder[0],
662                    out_channels=features_decoder[1],
663                    use_conv_transpose=use_conv_transpose,
664                )
665            )
666            self.deconv3 = nn.Sequential(
667                Deconv2DBlock(
668                    in_channels=self.embed_dim,
669                    out_channels=features_decoder[0],
670                    use_conv_transpose=use_conv_transpose,
671                ),
672                Deconv2DBlock(
673                    in_channels=features_decoder[0],
674                    out_channels=features_decoder[1],
675                    use_conv_transpose=use_conv_transpose,
676                ),
677                Deconv2DBlock(
678                    in_channels=features_decoder[1],
679                    out_channels=features_decoder[2],
680                    use_conv_transpose=use_conv_transpose,
681                )
682            )
683            self.deconv4 = ConvBlock2d(in_chans, features_decoder[-1])
684        else:
685            self.deconv1 = Deconv2DBlock(
686                in_channels=self.embed_dim,
687                out_channels=features_decoder[0],
688                use_conv_transpose=use_conv_transpose,
689            )
690            self.deconv2 = Deconv2DBlock(
691                in_channels=features_decoder[0],
692                out_channels=features_decoder[1],
693                use_conv_transpose=use_conv_transpose,
694            )
695            self.deconv3 = Deconv2DBlock(
696                in_channels=features_decoder[1],
697                out_channels=features_decoder[2],
698                use_conv_transpose=use_conv_transpose,
699            )
700            self.deconv4 = Deconv2DBlock(
701                in_channels=features_decoder[2],
702                out_channels=features_decoder[3],
703                use_conv_transpose=use_conv_transpose,
704            )
705
706        self.base = ConvBlock2d(self.embed_dim, features_decoder[0])
707        self.out_conv = nn.Conv2d(features_decoder[-1], out_channels, 1)
708        self.deconv_out = _upsampler(
709            scale_factor=2, in_channels=features_decoder[-1], out_channels=features_decoder[-1]
710        )
711        self.decoder_head = ConvBlock2d(2 * features_decoder[-1], features_decoder[-1])
712
713    def encode(self, x: torch.Tensor):
714        """Preprocess the input and run the image encoder.
715
716        Args:
717            x: The input tensor of shape (B, C, Y, X).
718
719        Returns:
720            The features as a tuple of the image embeddings, the list of intermediate encoder outputs
721            (None if the encoder returns only the embeddings) and the preprocessed input, which the
722            skip connections consume, and the spatial shape after preprocessing.
723        """
724        # Reshape the inputs to the shape expected by the encoder
725        # and normalize the inputs if normalization is part of the model.
726        x, input_shape = self.preprocess(x)
727
728        encoder_outputs = self.encoder(x)
729
730        if isinstance(encoder_outputs[-1], list):
731            # `encoder_outputs` can be arranged in only two forms:
732            #   - either we only return the image embeddings
733            #   - or, we return the image embeddings and the "list" of global attention layers
734            z12, from_encoder = encoder_outputs
735        else:
736            z12, from_encoder = encoder_outputs, None
737
738        return (z12, from_encoder, x), input_shape
739
740    def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor:
741        """Run the convolutional decoder on the encoder features.
742
743        Args:
744            features: The tuple returned by `encode`.
745            input_shape: The spatial shape (Y, X) after preprocessing.
746            original_shape: The spatial shape (Y, X) of the original input.
747
748        Returns:
749            The UNETR output, resized to `original_shape`.
750        """
751        z12, from_encoder, x = features
752
753        if self.use_skip_connection:
754            from_encoder = from_encoder[::-1]
755            z9 = self.deconv1(from_encoder[0])
756            z6 = self.deconv2(from_encoder[1])
757            z3 = self.deconv3(from_encoder[2])
758            z0 = self.deconv4(x)
759
760        else:
761            z9 = self.deconv1(z12)
762            z6 = self.deconv2(z9)
763            z3 = self.deconv3(z6)
764            z0 = self.deconv4(z3)
765
766        updated_from_encoder = [z9, z6, z3]
767
768        x = self.base(z12)
769        x = self.decoder(x, encoder_inputs=updated_from_encoder)
770        x = self.deconv_out(x)
771
772        x = torch.cat([x, z0], dim=1)
773        x = self.decoder_head(x)
774
775        x = self.out_conv(x)
776        if self.final_activation is not None:
777            x = self.final_activation(x)
778
779        return self.postprocess_masks(x, input_shape, original_shape)

A (2d-only) UNet Transformer using a vision transformer as encoder and a convolutional decoder.

UNETR( img_size: int = 1024, backbone: Literal['sam', 'sam2', 'sam3', 'cellpose_sam', 'mae', 'scalemae', 'dinov2', 'dinov3', 'torchvision'] = 'sam', encoder: Union[torch.nn.modules.module.Module, str, NoneType] = 'vit_b', decoder: Optional[torch.nn.modules.module.Module] = None, out_channels: int = 1, use_sam_stats: bool = False, use_mae_stats: bool = False, use_dino_stats: bool = False, use_imagenet_stats: bool = False, resize_input: bool = True, encoder_checkpoint: Union[str, collections.OrderedDict, NoneType] = None, final_activation: Union[torch.nn.modules.module.Module, str, NoneType] = None, use_skip_connection: bool = True, embed_dim: Optional[int] = None, use_conv_transpose: bool = False, perform_range_checks: bool = True, **kwargs)
576    def __init__(
577        self,
578        img_size: int = 1024,
579        backbone: Literal[
580            "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3", "torchvision"
581        ] = "sam",
582        encoder: Optional[Union[nn.Module, str]] = "vit_b",
583        decoder: Optional[nn.Module] = None,
584        out_channels: int = 1,
585        use_sam_stats: bool = False,
586        use_mae_stats: bool = False,
587        use_dino_stats: bool = False,
588        use_imagenet_stats: bool = False,
589        resize_input: bool = True,
590        encoder_checkpoint: Optional[Union[str, OrderedDict]] = None,
591        final_activation: Optional[Union[str, nn.Module]] = None,
592        use_skip_connection: bool = True,
593        embed_dim: Optional[int] = None,
594        use_conv_transpose: bool = False,
595        perform_range_checks: bool = True,
596        **kwargs
597    ) -> None:
598
599        super().__init__(
600            img_size=img_size,
601            backbone=backbone,
602            encoder=encoder,
603            decoder=decoder,
604            out_channels=out_channels,
605            use_sam_stats=use_sam_stats,
606            use_mae_stats=use_mae_stats,
607            use_dino_stats=use_dino_stats,
608            use_imagenet_stats=use_imagenet_stats,
609            resize_input=resize_input,
610            encoder_checkpoint=encoder_checkpoint,
611            final_activation=final_activation,
612            use_skip_connection=use_skip_connection,
613            embed_dim=embed_dim,
614            use_conv_transpose=use_conv_transpose,
615            perform_range_checks=perform_range_checks,
616            **kwargs,
617        )
618
619        encoder = self.encoder
620
621        if backbone == "sam2" and hasattr(encoder, "trunk"):
622            in_chans = encoder.trunk.patch_embed.proj.in_channels
623        elif hasattr(encoder, "in_chans"):
624            in_chans = encoder.in_chans
625        else:  # `nn.Module` ViT backbone.
626            try:
627                in_chans = encoder.patch_embed.proj.in_channels
628            except AttributeError:  # for getting the input channels while using 'vit_t' from MobileSam
629                in_chans = encoder.patch_embed.seq[0].c.in_channels
630
631        # parameters for the decoder network
632        depth = 3
633        gain = 2
634        features_decoder = [self.initial_features * gain ** i for i in range(depth + 1)][::-1]
635        scale_factors = depth * [2]
636        self.out_channels = out_channels
637
638        # choice of upsampler - to use (bilinear interpolation + conv) or conv transpose
639        _upsampler = SingleDeconv2DBlock if use_conv_transpose else Upsampler2d
640
641        self.decoder = decoder or Decoder(
642            features=features_decoder,
643            scale_factors=scale_factors[::-1],
644            conv_block_impl=ConvBlock2d,
645            sampler_impl=_upsampler,
646        )
647
648        if use_skip_connection:
649            self.deconv1 = Deconv2DBlock(
650                in_channels=self.embed_dim,
651                out_channels=features_decoder[0],
652                use_conv_transpose=use_conv_transpose,
653            )
654            self.deconv2 = nn.Sequential(
655                Deconv2DBlock(
656                    in_channels=self.embed_dim,
657                    out_channels=features_decoder[0],
658                    use_conv_transpose=use_conv_transpose,
659                ),
660                Deconv2DBlock(
661                    in_channels=features_decoder[0],
662                    out_channels=features_decoder[1],
663                    use_conv_transpose=use_conv_transpose,
664                )
665            )
666            self.deconv3 = nn.Sequential(
667                Deconv2DBlock(
668                    in_channels=self.embed_dim,
669                    out_channels=features_decoder[0],
670                    use_conv_transpose=use_conv_transpose,
671                ),
672                Deconv2DBlock(
673                    in_channels=features_decoder[0],
674                    out_channels=features_decoder[1],
675                    use_conv_transpose=use_conv_transpose,
676                ),
677                Deconv2DBlock(
678                    in_channels=features_decoder[1],
679                    out_channels=features_decoder[2],
680                    use_conv_transpose=use_conv_transpose,
681                )
682            )
683            self.deconv4 = ConvBlock2d(in_chans, features_decoder[-1])
684        else:
685            self.deconv1 = Deconv2DBlock(
686                in_channels=self.embed_dim,
687                out_channels=features_decoder[0],
688                use_conv_transpose=use_conv_transpose,
689            )
690            self.deconv2 = Deconv2DBlock(
691                in_channels=features_decoder[0],
692                out_channels=features_decoder[1],
693                use_conv_transpose=use_conv_transpose,
694            )
695            self.deconv3 = Deconv2DBlock(
696                in_channels=features_decoder[1],
697                out_channels=features_decoder[2],
698                use_conv_transpose=use_conv_transpose,
699            )
700            self.deconv4 = Deconv2DBlock(
701                in_channels=features_decoder[2],
702                out_channels=features_decoder[3],
703                use_conv_transpose=use_conv_transpose,
704            )
705
706        self.base = ConvBlock2d(self.embed_dim, features_decoder[0])
707        self.out_conv = nn.Conv2d(features_decoder[-1], out_channels, 1)
708        self.deconv_out = _upsampler(
709            scale_factor=2, in_channels=features_decoder[-1], out_channels=features_decoder[-1]
710        )
711        self.decoder_head = ConvBlock2d(2 * features_decoder[-1], features_decoder[-1])

Initialize internal Module state, shared by both nn.Module and ScriptModule.

out_channels
decoder
base
out_conv
deconv_out
decoder_head
def encode(self, x: torch.Tensor):
713    def encode(self, x: torch.Tensor):
714        """Preprocess the input and run the image encoder.
715
716        Args:
717            x: The input tensor of shape (B, C, Y, X).
718
719        Returns:
720            The features as a tuple of the image embeddings, the list of intermediate encoder outputs
721            (None if the encoder returns only the embeddings) and the preprocessed input, which the
722            skip connections consume, and the spatial shape after preprocessing.
723        """
724        # Reshape the inputs to the shape expected by the encoder
725        # and normalize the inputs if normalization is part of the model.
726        x, input_shape = self.preprocess(x)
727
728        encoder_outputs = self.encoder(x)
729
730        if isinstance(encoder_outputs[-1], list):
731            # `encoder_outputs` can be arranged in only two forms:
732            #   - either we only return the image embeddings
733            #   - or, we return the image embeddings and the "list" of global attention layers
734            z12, from_encoder = encoder_outputs
735        else:
736            z12, from_encoder = encoder_outputs, None
737
738        return (z12, from_encoder, x), input_shape

Preprocess the input and run the image encoder.

Arguments:
  • x: The input tensor of shape (B, C, Y, X).
Returns:

The features as a tuple of the image embeddings, the list of intermediate encoder outputs (None if the encoder returns only the embeddings) and the preprocessed input, which the skip connections consume, and the spatial shape after preprocessing.

def decode( self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor:
740    def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor:
741        """Run the convolutional decoder on the encoder features.
742
743        Args:
744            features: The tuple returned by `encode`.
745            input_shape: The spatial shape (Y, X) after preprocessing.
746            original_shape: The spatial shape (Y, X) of the original input.
747
748        Returns:
749            The UNETR output, resized to `original_shape`.
750        """
751        z12, from_encoder, x = features
752
753        if self.use_skip_connection:
754            from_encoder = from_encoder[::-1]
755            z9 = self.deconv1(from_encoder[0])
756            z6 = self.deconv2(from_encoder[1])
757            z3 = self.deconv3(from_encoder[2])
758            z0 = self.deconv4(x)
759
760        else:
761            z9 = self.deconv1(z12)
762            z6 = self.deconv2(z9)
763            z3 = self.deconv3(z6)
764            z0 = self.deconv4(z3)
765
766        updated_from_encoder = [z9, z6, z3]
767
768        x = self.base(z12)
769        x = self.decoder(x, encoder_inputs=updated_from_encoder)
770        x = self.deconv_out(x)
771
772        x = torch.cat([x, z0], dim=1)
773        x = self.decoder_head(x)
774
775        x = self.out_conv(x)
776        if self.final_activation is not None:
777            x = self.final_activation(x)
778
779        return self.postprocess_masks(x, input_shape, original_shape)

Run the convolutional decoder on the encoder features.

Arguments:
  • features: The tuple returned by encode.
  • input_shape: The spatial shape (Y, X) after preprocessing.
  • original_shape: The spatial shape (Y, X) of the original input.
Returns:

The UNETR output, resized to original_shape.

class UNETR2D(UNETR):
782class UNETR2D(UNETR):
783    """A two-dimensional UNet Transformer using a vision transformer as encoder and a convolutional decoder.
784    """
785    pass

A two-dimensional UNet Transformer using a vision transformer as encoder and a convolutional decoder.

class UNETR3D(UNETRBase):
788class UNETR3D(UNETRBase):
789    """A three dimensional UNet Transformer using a vision transformer as encoder and a convolutional decoder.
790    """
791    def __init__(
792        self,
793        img_size: int = 1024,
794        backbone: Literal[
795            "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3", "torchvision"
796        ] = "sam",
797        encoder: Optional[Union[nn.Module, str]] = "hvit_b",
798        decoder: Optional[nn.Module] = None,
799        out_channels: int = 1,
800        use_sam_stats: bool = False,
801        use_mae_stats: bool = False,
802        use_dino_stats: bool = False,
803        use_imagenet_stats: bool = False,
804        resize_input: bool = True,
805        encoder_checkpoint: Optional[Union[str, OrderedDict]] = None,
806        final_activation: Optional[Union[str, nn.Module]] = None,
807        use_skip_connection: bool = False,
808        embed_dim: Optional[int] = None,
809        use_conv_transpose: bool = False,
810        use_strip_pooling: bool = True,
811        perform_range_checks: bool = True,
812        **kwargs
813    ):
814        if use_skip_connection:
815            raise NotImplementedError("The framework cannot handle skip connections atm.")
816        if use_conv_transpose:
817            raise NotImplementedError("It's not enabled to switch between interpolation and transposed convolutions.")
818
819        # Sort the `embed_dim` out
820        embed_dim = 256 if embed_dim is None else embed_dim
821
822        super().__init__(
823            img_size=img_size,
824            backbone=backbone,
825            encoder=encoder,
826            decoder=decoder,
827            out_channels=out_channels,
828            use_sam_stats=use_sam_stats,
829            use_mae_stats=use_mae_stats,
830            use_dino_stats=use_dino_stats,
831            use_imagenet_stats=use_imagenet_stats,
832            resize_input=resize_input,
833            encoder_checkpoint=encoder_checkpoint,
834            final_activation=final_activation,
835            use_skip_connection=use_skip_connection,
836            embed_dim=embed_dim,
837            use_conv_transpose=use_conv_transpose,
838            perform_range_checks=perform_range_checks,
839            **kwargs,
840        )
841
842        # The 3d convolutional decoder.
843        # First, get the important parameters for the decoder.
844        depth = 3
845        gain = 2
846        features_decoder = [self.initial_features * gain ** i for i in range(depth + 1)][::-1]
847        scale_factors = [1, 2, 2]
848        self.out_channels = out_channels
849
850        # The mapping blocks.
851        self.deconv1 = Deconv3DBlock(
852            in_channels=embed_dim,
853            out_channels=features_decoder[0],
854            scale_factor=scale_factors,
855            use_strip_pooling=use_strip_pooling,
856        )
857        self.deconv2 = Deconv3DBlock(
858            in_channels=features_decoder[0],
859            out_channels=features_decoder[1],
860            scale_factor=scale_factors,
861            use_strip_pooling=use_strip_pooling,
862        )
863        self.deconv3 = Deconv3DBlock(
864            in_channels=features_decoder[1],
865            out_channels=features_decoder[2],
866            scale_factor=scale_factors,
867            use_strip_pooling=use_strip_pooling,
868        )
869        self.deconv4 = Deconv3DBlock(
870            in_channels=features_decoder[2],
871            out_channels=features_decoder[3],
872            scale_factor=scale_factors,
873            use_strip_pooling=use_strip_pooling,
874        )
875
876        # The core decoder block.
877        self.decoder = decoder or Decoder(
878            features=features_decoder,
879            scale_factors=[scale_factors] * depth,
880            conv_block_impl=partial(ConvBlock3dWithStrip, use_strip_pooling=use_strip_pooling),
881            sampler_impl=Upsampler3d,
882        )
883
884        # And the final upsampler to match the expected dimensions.
885        self.deconv_out = Deconv3DBlock(  # NOTE: changed `end_up` to `deconv_out`
886            in_channels=features_decoder[-1],
887            out_channels=features_decoder[-1],
888            scale_factor=scale_factors,
889            use_strip_pooling=use_strip_pooling,
890        )
891
892        # Additional conjunction blocks.
893        self.base = ConvBlock3dWithStrip(
894            in_channels=embed_dim,
895            out_channels=features_decoder[0],
896            use_strip_pooling=use_strip_pooling,
897        )
898
899        # And the output layers.
900        self.decoder_head = ConvBlock3dWithStrip(
901            in_channels=2 * features_decoder[-1],
902            out_channels=features_decoder[-1],
903            use_strip_pooling=use_strip_pooling,
904        )
905        self.out_conv = nn.Conv3d(features_decoder[-1], out_channels, 1)
906
907    def encode(self, x: torch.Tensor):
908        """Preprocess the input and run the image encoder on every z-slice.
909
910        Args:
911            x: Inputs of expected shape (B, C, Z, Y, X), where Z considers flexible inputs.
912
913        Returns:
914            The encoder features of shape (B, D, Z, Y', X') and the spatial shape after preprocessing.
915        """
916        Z = x.shape[2]
917        x, input_shape = self.preprocess(x)
918        features = torch.stack([self.encoder(x[:, :, i])[0] for i in range(Z)], dim=2)
919        return features, input_shape
920
921    def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor:
922        """Run the convolutional decoder on the encoder features.
923
924        Args:
925            features: Encoder features of shape (B, D, Z, Y', X'), see `encode`.
926            input_shape: The spatial shape (Z, Y, X) after preprocessing.
927            original_shape: The spatial shape (Z, Y, X) of the original input.
928
929        Returns:
930            The UNETR output, resized to `original_shape`.
931        """
932        # Prepare the counterparts for the decoder.
933        # NOTE: The section below is sequential, there's no skip connections atm.
934        z9 = self.deconv1(features)
935        z6 = self.deconv2(z9)
936        z3 = self.deconv3(z6)
937        z0 = self.deconv4(z3)
938
939        updated_from_encoder = [z9, z6, z3]
940
941        # Align the features through the base block.
942        x = self.base(features)
943        # Run the decoder
944        x = self.decoder(x, encoder_inputs=updated_from_encoder)
945        x = self.deconv_out(x)  # NOTE before `end_up`
946
947        # And the final output head.
948        x = torch.cat([x, z0], dim=1)
949        x = self.decoder_head(x)
950        x = self.out_conv(x)
951        if self.final_activation is not None:
952            x = self.final_activation(x)
953
954        # Postprocess the output back to original size.
955        return self.postprocess_masks(x, input_shape, original_shape)

A three dimensional UNet Transformer using a vision transformer as encoder and a convolutional decoder.

UNETR3D( img_size: int = 1024, backbone: Literal['sam', 'sam2', 'sam3', 'cellpose_sam', 'mae', 'scalemae', 'dinov2', 'dinov3', 'torchvision'] = 'sam', encoder: Union[torch.nn.modules.module.Module, str, NoneType] = 'hvit_b', decoder: Optional[torch.nn.modules.module.Module] = None, out_channels: int = 1, use_sam_stats: bool = False, use_mae_stats: bool = False, use_dino_stats: bool = False, use_imagenet_stats: bool = False, resize_input: bool = True, encoder_checkpoint: Union[str, collections.OrderedDict, NoneType] = None, final_activation: Union[torch.nn.modules.module.Module, str, NoneType] = None, use_skip_connection: bool = False, embed_dim: Optional[int] = None, use_conv_transpose: bool = False, use_strip_pooling: bool = True, perform_range_checks: bool = True, **kwargs)
791    def __init__(
792        self,
793        img_size: int = 1024,
794        backbone: Literal[
795            "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3", "torchvision"
796        ] = "sam",
797        encoder: Optional[Union[nn.Module, str]] = "hvit_b",
798        decoder: Optional[nn.Module] = None,
799        out_channels: int = 1,
800        use_sam_stats: bool = False,
801        use_mae_stats: bool = False,
802        use_dino_stats: bool = False,
803        use_imagenet_stats: bool = False,
804        resize_input: bool = True,
805        encoder_checkpoint: Optional[Union[str, OrderedDict]] = None,
806        final_activation: Optional[Union[str, nn.Module]] = None,
807        use_skip_connection: bool = False,
808        embed_dim: Optional[int] = None,
809        use_conv_transpose: bool = False,
810        use_strip_pooling: bool = True,
811        perform_range_checks: bool = True,
812        **kwargs
813    ):
814        if use_skip_connection:
815            raise NotImplementedError("The framework cannot handle skip connections atm.")
816        if use_conv_transpose:
817            raise NotImplementedError("It's not enabled to switch between interpolation and transposed convolutions.")
818
819        # Sort the `embed_dim` out
820        embed_dim = 256 if embed_dim is None else embed_dim
821
822        super().__init__(
823            img_size=img_size,
824            backbone=backbone,
825            encoder=encoder,
826            decoder=decoder,
827            out_channels=out_channels,
828            use_sam_stats=use_sam_stats,
829            use_mae_stats=use_mae_stats,
830            use_dino_stats=use_dino_stats,
831            use_imagenet_stats=use_imagenet_stats,
832            resize_input=resize_input,
833            encoder_checkpoint=encoder_checkpoint,
834            final_activation=final_activation,
835            use_skip_connection=use_skip_connection,
836            embed_dim=embed_dim,
837            use_conv_transpose=use_conv_transpose,
838            perform_range_checks=perform_range_checks,
839            **kwargs,
840        )
841
842        # The 3d convolutional decoder.
843        # First, get the important parameters for the decoder.
844        depth = 3
845        gain = 2
846        features_decoder = [self.initial_features * gain ** i for i in range(depth + 1)][::-1]
847        scale_factors = [1, 2, 2]
848        self.out_channels = out_channels
849
850        # The mapping blocks.
851        self.deconv1 = Deconv3DBlock(
852            in_channels=embed_dim,
853            out_channels=features_decoder[0],
854            scale_factor=scale_factors,
855            use_strip_pooling=use_strip_pooling,
856        )
857        self.deconv2 = Deconv3DBlock(
858            in_channels=features_decoder[0],
859            out_channels=features_decoder[1],
860            scale_factor=scale_factors,
861            use_strip_pooling=use_strip_pooling,
862        )
863        self.deconv3 = Deconv3DBlock(
864            in_channels=features_decoder[1],
865            out_channels=features_decoder[2],
866            scale_factor=scale_factors,
867            use_strip_pooling=use_strip_pooling,
868        )
869        self.deconv4 = Deconv3DBlock(
870            in_channels=features_decoder[2],
871            out_channels=features_decoder[3],
872            scale_factor=scale_factors,
873            use_strip_pooling=use_strip_pooling,
874        )
875
876        # The core decoder block.
877        self.decoder = decoder or Decoder(
878            features=features_decoder,
879            scale_factors=[scale_factors] * depth,
880            conv_block_impl=partial(ConvBlock3dWithStrip, use_strip_pooling=use_strip_pooling),
881            sampler_impl=Upsampler3d,
882        )
883
884        # And the final upsampler to match the expected dimensions.
885        self.deconv_out = Deconv3DBlock(  # NOTE: changed `end_up` to `deconv_out`
886            in_channels=features_decoder[-1],
887            out_channels=features_decoder[-1],
888            scale_factor=scale_factors,
889            use_strip_pooling=use_strip_pooling,
890        )
891
892        # Additional conjunction blocks.
893        self.base = ConvBlock3dWithStrip(
894            in_channels=embed_dim,
895            out_channels=features_decoder[0],
896            use_strip_pooling=use_strip_pooling,
897        )
898
899        # And the output layers.
900        self.decoder_head = ConvBlock3dWithStrip(
901            in_channels=2 * features_decoder[-1],
902            out_channels=features_decoder[-1],
903            use_strip_pooling=use_strip_pooling,
904        )
905        self.out_conv = nn.Conv3d(features_decoder[-1], out_channels, 1)

Initialize internal Module state, shared by both nn.Module and ScriptModule.

out_channels
deconv1
deconv2
deconv3
deconv4
decoder
deconv_out
base
decoder_head
out_conv
def encode(self, x: torch.Tensor):
907    def encode(self, x: torch.Tensor):
908        """Preprocess the input and run the image encoder on every z-slice.
909
910        Args:
911            x: Inputs of expected shape (B, C, Z, Y, X), where Z considers flexible inputs.
912
913        Returns:
914            The encoder features of shape (B, D, Z, Y', X') and the spatial shape after preprocessing.
915        """
916        Z = x.shape[2]
917        x, input_shape = self.preprocess(x)
918        features = torch.stack([self.encoder(x[:, :, i])[0] for i in range(Z)], dim=2)
919        return features, input_shape

Preprocess the input and run the image encoder on every z-slice.

Arguments:
  • x: Inputs of expected shape (B, C, Z, Y, X), where Z considers flexible inputs.
Returns:

The encoder features of shape (B, D, Z, Y', X') and the spatial shape after preprocessing.

def decode( self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor:
921    def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor:
922        """Run the convolutional decoder on the encoder features.
923
924        Args:
925            features: Encoder features of shape (B, D, Z, Y', X'), see `encode`.
926            input_shape: The spatial shape (Z, Y, X) after preprocessing.
927            original_shape: The spatial shape (Z, Y, X) of the original input.
928
929        Returns:
930            The UNETR output, resized to `original_shape`.
931        """
932        # Prepare the counterparts for the decoder.
933        # NOTE: The section below is sequential, there's no skip connections atm.
934        z9 = self.deconv1(features)
935        z6 = self.deconv2(z9)
936        z3 = self.deconv3(z6)
937        z0 = self.deconv4(z3)
938
939        updated_from_encoder = [z9, z6, z3]
940
941        # Align the features through the base block.
942        x = self.base(features)
943        # Run the decoder
944        x = self.decoder(x, encoder_inputs=updated_from_encoder)
945        x = self.deconv_out(x)  # NOTE before `end_up`
946
947        # And the final output head.
948        x = torch.cat([x, z0], dim=1)
949        x = self.decoder_head(x)
950        x = self.out_conv(x)
951        if self.final_activation is not None:
952            x = self.final_activation(x)
953
954        # Postprocess the output back to original size.
955        return self.postprocess_masks(x, input_shape, original_shape)

Run the convolutional decoder on the encoder features.

Arguments:
  • features: Encoder features of shape (B, D, Z, Y', X'), see encode.
  • input_shape: The spatial shape (Z, Y, X) after preprocessing.
  • original_shape: The spatial shape (Z, Y, X) of the original input.
Returns:

The UNETR output, resized to original_shape.