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 self.init_kwargs = { 618 "img_size": img_size, "backbone": backbone, "encoder": encoder, "decoder": decoder, 619 "out_channels": out_channels, "use_sam_stats": use_sam_stats, 620 "use_mae_stats": use_mae_stats, "use_dino_stats": use_dino_stats, 621 "use_imagenet_stats": use_imagenet_stats, "resize_input": resize_input, 622 "final_activation": final_activation, 623 "use_skip_connection": use_skip_connection, "embed_dim": embed_dim, 624 "use_conv_transpose": use_conv_transpose, "perform_range_checks": perform_range_checks, 625 **kwargs, 626 } 627 628 encoder = self.encoder 629 630 if backbone == "sam2" and hasattr(encoder, "trunk"): 631 in_chans = encoder.trunk.patch_embed.proj.in_channels 632 elif hasattr(encoder, "in_chans"): 633 in_chans = encoder.in_chans 634 else: # `nn.Module` ViT backbone. 635 try: 636 in_chans = encoder.patch_embed.proj.in_channels 637 except AttributeError: # for getting the input channels while using 'vit_t' from MobileSam 638 in_chans = encoder.patch_embed.seq[0].c.in_channels 639 640 # parameters for the decoder network 641 depth = 3 642 gain = 2 643 features_decoder = [self.initial_features * gain ** i for i in range(depth + 1)][::-1] 644 scale_factors = depth * [2] 645 self.out_channels = out_channels 646 647 # choice of upsampler - to use (bilinear interpolation + conv) or conv transpose 648 _upsampler = SingleDeconv2DBlock if use_conv_transpose else Upsampler2d 649 650 self.decoder = decoder or Decoder( 651 features=features_decoder, 652 scale_factors=scale_factors[::-1], 653 conv_block_impl=ConvBlock2d, 654 sampler_impl=_upsampler, 655 ) 656 657 if use_skip_connection: 658 self.deconv1 = Deconv2DBlock( 659 in_channels=self.embed_dim, 660 out_channels=features_decoder[0], 661 use_conv_transpose=use_conv_transpose, 662 ) 663 self.deconv2 = nn.Sequential( 664 Deconv2DBlock( 665 in_channels=self.embed_dim, 666 out_channels=features_decoder[0], 667 use_conv_transpose=use_conv_transpose, 668 ), 669 Deconv2DBlock( 670 in_channels=features_decoder[0], 671 out_channels=features_decoder[1], 672 use_conv_transpose=use_conv_transpose, 673 ) 674 ) 675 self.deconv3 = nn.Sequential( 676 Deconv2DBlock( 677 in_channels=self.embed_dim, 678 out_channels=features_decoder[0], 679 use_conv_transpose=use_conv_transpose, 680 ), 681 Deconv2DBlock( 682 in_channels=features_decoder[0], 683 out_channels=features_decoder[1], 684 use_conv_transpose=use_conv_transpose, 685 ), 686 Deconv2DBlock( 687 in_channels=features_decoder[1], 688 out_channels=features_decoder[2], 689 use_conv_transpose=use_conv_transpose, 690 ) 691 ) 692 self.deconv4 = ConvBlock2d(in_chans, features_decoder[-1]) 693 else: 694 self.deconv1 = Deconv2DBlock( 695 in_channels=self.embed_dim, 696 out_channels=features_decoder[0], 697 use_conv_transpose=use_conv_transpose, 698 ) 699 self.deconv2 = Deconv2DBlock( 700 in_channels=features_decoder[0], 701 out_channels=features_decoder[1], 702 use_conv_transpose=use_conv_transpose, 703 ) 704 self.deconv3 = Deconv2DBlock( 705 in_channels=features_decoder[1], 706 out_channels=features_decoder[2], 707 use_conv_transpose=use_conv_transpose, 708 ) 709 self.deconv4 = Deconv2DBlock( 710 in_channels=features_decoder[2], 711 out_channels=features_decoder[3], 712 use_conv_transpose=use_conv_transpose, 713 ) 714 715 self.base = ConvBlock2d(self.embed_dim, features_decoder[0]) 716 self.out_conv = nn.Conv2d(features_decoder[-1], out_channels, 1) 717 self.deconv_out = _upsampler( 718 scale_factor=2, in_channels=features_decoder[-1], out_channels=features_decoder[-1] 719 ) 720 self.decoder_head = ConvBlock2d(2 * features_decoder[-1], features_decoder[-1]) 721 722 def encode(self, x: torch.Tensor): 723 """Preprocess the input and run the image encoder. 724 725 Args: 726 x: The input tensor of shape (B, C, Y, X). 727 728 Returns: 729 The features as a tuple of the image embeddings, the list of intermediate encoder outputs 730 (None if the encoder returns only the embeddings) and the preprocessed input, which the 731 skip connections consume, and the spatial shape after preprocessing. 732 """ 733 # Reshape the inputs to the shape expected by the encoder 734 # and normalize the inputs if normalization is part of the model. 735 x, input_shape = self.preprocess(x) 736 737 encoder_outputs = self.encoder(x) 738 739 if isinstance(encoder_outputs[-1], list): 740 # `encoder_outputs` can be arranged in only two forms: 741 # - either we only return the image embeddings 742 # - or, we return the image embeddings and the "list" of global attention layers 743 z12, from_encoder = encoder_outputs 744 else: 745 z12, from_encoder = encoder_outputs, None 746 747 return (z12, from_encoder, x), input_shape 748 749 def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor: 750 """Run the convolutional decoder on the encoder features. 751 752 Args: 753 features: The tuple returned by `encode`. 754 input_shape: The spatial shape (Y, X) after preprocessing. 755 original_shape: The spatial shape (Y, X) of the original input. 756 757 Returns: 758 The UNETR output, resized to `original_shape`. 759 """ 760 z12, from_encoder, x = features 761 762 if self.use_skip_connection: 763 from_encoder = from_encoder[::-1] 764 z9 = self.deconv1(from_encoder[0]) 765 z6 = self.deconv2(from_encoder[1]) 766 z3 = self.deconv3(from_encoder[2]) 767 z0 = self.deconv4(x) 768 769 else: 770 z9 = self.deconv1(z12) 771 z6 = self.deconv2(z9) 772 z3 = self.deconv3(z6) 773 z0 = self.deconv4(z3) 774 775 updated_from_encoder = [z9, z6, z3] 776 777 x = self.base(z12) 778 x = self.decoder(x, encoder_inputs=updated_from_encoder) 779 x = self.deconv_out(x) 780 781 x = torch.cat([x, z0], dim=1) 782 x = self.decoder_head(x) 783 784 x = self.out_conv(x) 785 if self.final_activation is not None: 786 x = self.final_activation(x) 787 788 return self.postprocess_masks(x, input_shape, original_shape) 789 790 791class UNETR2D(UNETR): 792 """A two-dimensional UNet Transformer using a vision transformer as encoder and a convolutional decoder. 793 """ 794 pass 795 796 797class UNETR3D(UNETRBase): 798 """A three dimensional UNet Transformer using a vision transformer as encoder and a convolutional decoder. 799 """ 800 def __init__( 801 self, 802 img_size: int = 1024, 803 backbone: Literal[ 804 "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3", "torchvision" 805 ] = "sam", 806 encoder: Optional[Union[nn.Module, str]] = "hvit_b", 807 decoder: Optional[nn.Module] = None, 808 out_channels: int = 1, 809 use_sam_stats: bool = False, 810 use_mae_stats: bool = False, 811 use_dino_stats: bool = False, 812 use_imagenet_stats: bool = False, 813 resize_input: bool = True, 814 encoder_checkpoint: Optional[Union[str, OrderedDict]] = None, 815 final_activation: Optional[Union[str, nn.Module]] = None, 816 use_skip_connection: bool = False, 817 embed_dim: Optional[int] = None, 818 use_conv_transpose: bool = False, 819 use_strip_pooling: bool = True, 820 perform_range_checks: bool = True, 821 **kwargs 822 ): 823 if use_skip_connection: 824 raise NotImplementedError("The framework cannot handle skip connections atm.") 825 if use_conv_transpose: 826 raise NotImplementedError("It's not enabled to switch between interpolation and transposed convolutions.") 827 828 # Sort the `embed_dim` out 829 embed_dim = 256 if embed_dim is None else embed_dim 830 831 super().__init__( 832 img_size=img_size, 833 backbone=backbone, 834 encoder=encoder, 835 decoder=decoder, 836 out_channels=out_channels, 837 use_sam_stats=use_sam_stats, 838 use_mae_stats=use_mae_stats, 839 use_dino_stats=use_dino_stats, 840 use_imagenet_stats=use_imagenet_stats, 841 resize_input=resize_input, 842 encoder_checkpoint=encoder_checkpoint, 843 final_activation=final_activation, 844 use_skip_connection=use_skip_connection, 845 embed_dim=embed_dim, 846 use_conv_transpose=use_conv_transpose, 847 perform_range_checks=perform_range_checks, 848 **kwargs, 849 ) 850 self.init_kwargs = { 851 "img_size": img_size, "backbone": backbone, "encoder": encoder, "decoder": decoder, 852 "out_channels": out_channels, "use_sam_stats": use_sam_stats, 853 "use_mae_stats": use_mae_stats, "use_dino_stats": use_dino_stats, 854 "use_imagenet_stats": use_imagenet_stats, "resize_input": resize_input, 855 "final_activation": final_activation, 856 "use_skip_connection": use_skip_connection, "embed_dim": embed_dim, 857 "use_conv_transpose": use_conv_transpose, "use_strip_pooling": use_strip_pooling, 858 "perform_range_checks": perform_range_checks, **kwargs, 859 } 860 861 # The 3d convolutional decoder. 862 # First, get the important parameters for the decoder. 863 depth = 3 864 gain = 2 865 features_decoder = [self.initial_features * gain ** i for i in range(depth + 1)][::-1] 866 scale_factors = [1, 2, 2] 867 self.out_channels = out_channels 868 869 # The mapping blocks. 870 self.deconv1 = Deconv3DBlock( 871 in_channels=embed_dim, 872 out_channels=features_decoder[0], 873 scale_factor=scale_factors, 874 use_strip_pooling=use_strip_pooling, 875 ) 876 self.deconv2 = Deconv3DBlock( 877 in_channels=features_decoder[0], 878 out_channels=features_decoder[1], 879 scale_factor=scale_factors, 880 use_strip_pooling=use_strip_pooling, 881 ) 882 self.deconv3 = Deconv3DBlock( 883 in_channels=features_decoder[1], 884 out_channels=features_decoder[2], 885 scale_factor=scale_factors, 886 use_strip_pooling=use_strip_pooling, 887 ) 888 self.deconv4 = Deconv3DBlock( 889 in_channels=features_decoder[2], 890 out_channels=features_decoder[3], 891 scale_factor=scale_factors, 892 use_strip_pooling=use_strip_pooling, 893 ) 894 895 # The core decoder block. 896 self.decoder = decoder or Decoder( 897 features=features_decoder, 898 scale_factors=[scale_factors] * depth, 899 conv_block_impl=partial(ConvBlock3dWithStrip, use_strip_pooling=use_strip_pooling), 900 sampler_impl=Upsampler3d, 901 ) 902 903 # And the final upsampler to match the expected dimensions. 904 self.deconv_out = Deconv3DBlock( # NOTE: changed `end_up` to `deconv_out` 905 in_channels=features_decoder[-1], 906 out_channels=features_decoder[-1], 907 scale_factor=scale_factors, 908 use_strip_pooling=use_strip_pooling, 909 ) 910 911 # Additional conjunction blocks. 912 self.base = ConvBlock3dWithStrip( 913 in_channels=embed_dim, 914 out_channels=features_decoder[0], 915 use_strip_pooling=use_strip_pooling, 916 ) 917 918 # And the output layers. 919 self.decoder_head = ConvBlock3dWithStrip( 920 in_channels=2 * features_decoder[-1], 921 out_channels=features_decoder[-1], 922 use_strip_pooling=use_strip_pooling, 923 ) 924 self.out_conv = nn.Conv3d(features_decoder[-1], out_channels, 1) 925 926 def encode(self, x: torch.Tensor): 927 """Preprocess the input and run the image encoder on every z-slice. 928 929 Args: 930 x: Inputs of expected shape (B, C, Z, Y, X), where Z considers flexible inputs. 931 932 Returns: 933 The encoder features of shape (B, D, Z, Y', X') and the spatial shape after preprocessing. 934 """ 935 Z = x.shape[2] 936 x, input_shape = self.preprocess(x) 937 features = torch.stack([self.encoder(x[:, :, i])[0] for i in range(Z)], dim=2) 938 return features, input_shape 939 940 def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor: 941 """Run the convolutional decoder on the encoder features. 942 943 Args: 944 features: Encoder features of shape (B, D, Z, Y', X'), see `encode`. 945 input_shape: The spatial shape (Z, Y, X) after preprocessing. 946 original_shape: The spatial shape (Z, Y, X) of the original input. 947 948 Returns: 949 The UNETR output, resized to `original_shape`. 950 """ 951 # Prepare the counterparts for the decoder. 952 # NOTE: The section below is sequential, there's no skip connections atm. 953 z9 = self.deconv1(features) 954 z6 = self.deconv2(z9) 955 z3 = self.deconv3(z6) 956 z0 = self.deconv4(z3) 957 958 updated_from_encoder = [z9, z6, z3] 959 960 # Align the features through the base block. 961 x = self.base(features) 962 # Run the decoder 963 x = self.decoder(x, encoder_inputs=updated_from_encoder) 964 x = self.deconv_out(x) # NOTE before `end_up` 965 966 # And the final output head. 967 x = torch.cat([x, z0], dim=1) 968 x = self.decoder_head(x) 969 x = self.out_conv(x) 970 if self.final_activation is not None: 971 x = self.final_activation(x) 972 973 # Postprocess the output back to original size. 974 return self.postprocess_masks(x, input_shape, original_shape) 975 976# 977# ADDITIONAL FUNCTIONALITIES 978# 979 980 981def _strip_pooling_layers(enabled, channels) -> nn.Module: 982 return DepthStripPooling(channels) if enabled else nn.Identity() 983 984 985class DepthStripPooling(nn.Module): 986 """@private 987 """ 988 def __init__(self, channels: int, reduction: int = 4): 989 """Block for strip pooling along the depth dimension (only). 990 991 eg. for 3D (Z > 1) - it aggregates global context across depth by adaptive avg pooling 992 to Z=1, and then passes through a small 1x1x1 MLP, then broadcasts it back to Z to 993 modulate the original features (using a gated residual). 994 995 For 2D (Z == 1): returns input unchanged (no-op). 996 997 Args: 998 channels: The output channels. 999 reduction: The reduction of the hidden layers. 1000 """ 1001 super().__init__() 1002 hidden = max(1, channels // reduction) 1003 self.conv1 = nn.Conv3d(channels, hidden, kernel_size=1) 1004 self.bn1 = nn.BatchNorm3d(hidden) 1005 self.relu = nn.ReLU(inplace=True) 1006 self.conv2 = nn.Conv3d(hidden, channels, kernel_size=1) 1007 1008 def forward(self, x: torch.Tensor) -> torch.Tensor: 1009 if x.dim() != 5: 1010 raise ValueError(f"DepthStripPooling expects 5D tensors as input, got '{x.shape}'.") 1011 1012 B, C, Z, H, W = x.shape 1013 if Z == 1: # i.e. always the case of all 2d. 1014 return x # We simply do nothing there. 1015 1016 # We pool only along the depth dimension: i.e. target shape (B, C, 1, H, W). 1017 # A plain mean over Z is the same operation as adaptive_avg_pool3d to (1, H, W), but its 1018 # reduction kernel is several times faster at full resolution. 1019 feat = x.mean(dim=2, keepdim=True) 1020 feat = self.conv1(feat) 1021 feat = self.bn1(feat) 1022 feat = self.relu(feat) 1023 feat = self.conv2(feat) 1024 gate = torch.sigmoid(feat).expand(B, C, Z, H, W) # Broadcast the collapsed depth context back to all slices 1025 1026 # Gated residual fusion 1027 return x * gate + x 1028 1029 1030class Deconv3DBlock(nn.Module): 1031 """@private 1032 """ 1033 def __init__( 1034 self, 1035 scale_factor, 1036 in_channels, 1037 out_channels, 1038 kernel_size=3, 1039 anisotropic_kernel=True, 1040 use_strip_pooling=True, 1041 ): 1042 super().__init__() 1043 conv_block_kwargs = { 1044 "in_channels": out_channels, 1045 "out_channels": out_channels, 1046 "kernel_size": kernel_size, 1047 "padding": ((kernel_size - 1) // 2), 1048 } 1049 if anisotropic_kernel: 1050 conv_block_kwargs = _update_conv_kwargs(conv_block_kwargs, scale_factor) 1051 1052 self.block = nn.Sequential( 1053 Upsampler3d(scale_factor, in_channels, out_channels), 1054 nn.Conv3d(**conv_block_kwargs), 1055 nn.BatchNorm3d(out_channels), 1056 nn.ReLU(True), 1057 _strip_pooling_layers(enabled=use_strip_pooling, channels=out_channels), 1058 ) 1059 1060 def forward(self, x): 1061 return self.block(x) 1062 1063 1064class ConvBlock3dWithStrip(nn.Module): 1065 """@private 1066 """ 1067 def __init__( 1068 self, in_channels: int, out_channels: int, use_strip_pooling: bool = True, **kwargs 1069 ): 1070 super().__init__() 1071 self.block = nn.Sequential( 1072 ConvBlock3d(in_channels, out_channels, **kwargs), 1073 _strip_pooling_layers(enabled=use_strip_pooling, channels=out_channels), 1074 ) 1075 1076 def forward(self, x): 1077 return self.block(x) 1078 1079 1080class SingleDeconv2DBlock(nn.Module): 1081 """@private 1082 """ 1083 def __init__(self, scale_factor, in_channels, out_channels): 1084 super().__init__() 1085 self.block = nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2, stride=2, padding=0, output_padding=0) 1086 1087 def forward(self, x): 1088 return self.block(x) 1089 1090 1091class SingleConv2DBlock(nn.Module): 1092 """@private 1093 """ 1094 def __init__(self, in_channels, out_channels, kernel_size): 1095 super().__init__() 1096 self.block = nn.Conv2d( 1097 in_channels, out_channels, kernel_size=kernel_size, stride=1, padding=((kernel_size - 1) // 2) 1098 ) 1099 1100 def forward(self, x): 1101 return self.block(x) 1102 1103 1104class Conv2DBlock(nn.Module): 1105 """@private 1106 """ 1107 def __init__(self, in_channels, out_channels, kernel_size=3): 1108 super().__init__() 1109 self.block = nn.Sequential( 1110 SingleConv2DBlock(in_channels, out_channels, kernel_size), 1111 nn.BatchNorm2d(out_channels), 1112 nn.ReLU(True) 1113 ) 1114 1115 def forward(self, x): 1116 return self.block(x) 1117 1118 1119class Deconv2DBlock(nn.Module): 1120 """@private 1121 """ 1122 def __init__(self, in_channels, out_channels, kernel_size=3, use_conv_transpose=True): 1123 super().__init__() 1124 _upsampler = SingleDeconv2DBlock if use_conv_transpose else Upsampler2d 1125 self.block = nn.Sequential( 1126 _upsampler(scale_factor=2, in_channels=in_channels, out_channels=out_channels), 1127 SingleConv2DBlock(out_channels, out_channels, kernel_size), 1128 nn.BatchNorm2d(out_channels), 1129 nn.ReLU(True) 1130 ) 1131 1132 def forward(self, x): 1133 return self.block(x)
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 theimg_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'
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.
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.
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.
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
decodeand the spatial shape after preprocessing.
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
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.
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).
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 self.init_kwargs = { 619 "img_size": img_size, "backbone": backbone, "encoder": encoder, "decoder": decoder, 620 "out_channels": out_channels, "use_sam_stats": use_sam_stats, 621 "use_mae_stats": use_mae_stats, "use_dino_stats": use_dino_stats, 622 "use_imagenet_stats": use_imagenet_stats, "resize_input": resize_input, 623 "final_activation": final_activation, 624 "use_skip_connection": use_skip_connection, "embed_dim": embed_dim, 625 "use_conv_transpose": use_conv_transpose, "perform_range_checks": perform_range_checks, 626 **kwargs, 627 } 628 629 encoder = self.encoder 630 631 if backbone == "sam2" and hasattr(encoder, "trunk"): 632 in_chans = encoder.trunk.patch_embed.proj.in_channels 633 elif hasattr(encoder, "in_chans"): 634 in_chans = encoder.in_chans 635 else: # `nn.Module` ViT backbone. 636 try: 637 in_chans = encoder.patch_embed.proj.in_channels 638 except AttributeError: # for getting the input channels while using 'vit_t' from MobileSam 639 in_chans = encoder.patch_embed.seq[0].c.in_channels 640 641 # parameters for the decoder network 642 depth = 3 643 gain = 2 644 features_decoder = [self.initial_features * gain ** i for i in range(depth + 1)][::-1] 645 scale_factors = depth * [2] 646 self.out_channels = out_channels 647 648 # choice of upsampler - to use (bilinear interpolation + conv) or conv transpose 649 _upsampler = SingleDeconv2DBlock if use_conv_transpose else Upsampler2d 650 651 self.decoder = decoder or Decoder( 652 features=features_decoder, 653 scale_factors=scale_factors[::-1], 654 conv_block_impl=ConvBlock2d, 655 sampler_impl=_upsampler, 656 ) 657 658 if use_skip_connection: 659 self.deconv1 = Deconv2DBlock( 660 in_channels=self.embed_dim, 661 out_channels=features_decoder[0], 662 use_conv_transpose=use_conv_transpose, 663 ) 664 self.deconv2 = nn.Sequential( 665 Deconv2DBlock( 666 in_channels=self.embed_dim, 667 out_channels=features_decoder[0], 668 use_conv_transpose=use_conv_transpose, 669 ), 670 Deconv2DBlock( 671 in_channels=features_decoder[0], 672 out_channels=features_decoder[1], 673 use_conv_transpose=use_conv_transpose, 674 ) 675 ) 676 self.deconv3 = nn.Sequential( 677 Deconv2DBlock( 678 in_channels=self.embed_dim, 679 out_channels=features_decoder[0], 680 use_conv_transpose=use_conv_transpose, 681 ), 682 Deconv2DBlock( 683 in_channels=features_decoder[0], 684 out_channels=features_decoder[1], 685 use_conv_transpose=use_conv_transpose, 686 ), 687 Deconv2DBlock( 688 in_channels=features_decoder[1], 689 out_channels=features_decoder[2], 690 use_conv_transpose=use_conv_transpose, 691 ) 692 ) 693 self.deconv4 = ConvBlock2d(in_chans, features_decoder[-1]) 694 else: 695 self.deconv1 = Deconv2DBlock( 696 in_channels=self.embed_dim, 697 out_channels=features_decoder[0], 698 use_conv_transpose=use_conv_transpose, 699 ) 700 self.deconv2 = Deconv2DBlock( 701 in_channels=features_decoder[0], 702 out_channels=features_decoder[1], 703 use_conv_transpose=use_conv_transpose, 704 ) 705 self.deconv3 = Deconv2DBlock( 706 in_channels=features_decoder[1], 707 out_channels=features_decoder[2], 708 use_conv_transpose=use_conv_transpose, 709 ) 710 self.deconv4 = Deconv2DBlock( 711 in_channels=features_decoder[2], 712 out_channels=features_decoder[3], 713 use_conv_transpose=use_conv_transpose, 714 ) 715 716 self.base = ConvBlock2d(self.embed_dim, features_decoder[0]) 717 self.out_conv = nn.Conv2d(features_decoder[-1], out_channels, 1) 718 self.deconv_out = _upsampler( 719 scale_factor=2, in_channels=features_decoder[-1], out_channels=features_decoder[-1] 720 ) 721 self.decoder_head = ConvBlock2d(2 * features_decoder[-1], features_decoder[-1]) 722 723 def encode(self, x: torch.Tensor): 724 """Preprocess the input and run the image encoder. 725 726 Args: 727 x: The input tensor of shape (B, C, Y, X). 728 729 Returns: 730 The features as a tuple of the image embeddings, the list of intermediate encoder outputs 731 (None if the encoder returns only the embeddings) and the preprocessed input, which the 732 skip connections consume, and the spatial shape after preprocessing. 733 """ 734 # Reshape the inputs to the shape expected by the encoder 735 # and normalize the inputs if normalization is part of the model. 736 x, input_shape = self.preprocess(x) 737 738 encoder_outputs = self.encoder(x) 739 740 if isinstance(encoder_outputs[-1], list): 741 # `encoder_outputs` can be arranged in only two forms: 742 # - either we only return the image embeddings 743 # - or, we return the image embeddings and the "list" of global attention layers 744 z12, from_encoder = encoder_outputs 745 else: 746 z12, from_encoder = encoder_outputs, None 747 748 return (z12, from_encoder, x), input_shape 749 750 def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor: 751 """Run the convolutional decoder on the encoder features. 752 753 Args: 754 features: The tuple returned by `encode`. 755 input_shape: The spatial shape (Y, X) after preprocessing. 756 original_shape: The spatial shape (Y, X) of the original input. 757 758 Returns: 759 The UNETR output, resized to `original_shape`. 760 """ 761 z12, from_encoder, x = features 762 763 if self.use_skip_connection: 764 from_encoder = from_encoder[::-1] 765 z9 = self.deconv1(from_encoder[0]) 766 z6 = self.deconv2(from_encoder[1]) 767 z3 = self.deconv3(from_encoder[2]) 768 z0 = self.deconv4(x) 769 770 else: 771 z9 = self.deconv1(z12) 772 z6 = self.deconv2(z9) 773 z3 = self.deconv3(z6) 774 z0 = self.deconv4(z3) 775 776 updated_from_encoder = [z9, z6, z3] 777 778 x = self.base(z12) 779 x = self.decoder(x, encoder_inputs=updated_from_encoder) 780 x = self.deconv_out(x) 781 782 x = torch.cat([x, z0], dim=1) 783 x = self.decoder_head(x) 784 785 x = self.out_conv(x) 786 if self.final_activation is not None: 787 x = self.final_activation(x) 788 789 return self.postprocess_masks(x, input_shape, original_shape)
A (2d-only) UNet Transformer using a vision transformer as encoder and a convolutional decoder.
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 self.init_kwargs = { 619 "img_size": img_size, "backbone": backbone, "encoder": encoder, "decoder": decoder, 620 "out_channels": out_channels, "use_sam_stats": use_sam_stats, 621 "use_mae_stats": use_mae_stats, "use_dino_stats": use_dino_stats, 622 "use_imagenet_stats": use_imagenet_stats, "resize_input": resize_input, 623 "final_activation": final_activation, 624 "use_skip_connection": use_skip_connection, "embed_dim": embed_dim, 625 "use_conv_transpose": use_conv_transpose, "perform_range_checks": perform_range_checks, 626 **kwargs, 627 } 628 629 encoder = self.encoder 630 631 if backbone == "sam2" and hasattr(encoder, "trunk"): 632 in_chans = encoder.trunk.patch_embed.proj.in_channels 633 elif hasattr(encoder, "in_chans"): 634 in_chans = encoder.in_chans 635 else: # `nn.Module` ViT backbone. 636 try: 637 in_chans = encoder.patch_embed.proj.in_channels 638 except AttributeError: # for getting the input channels while using 'vit_t' from MobileSam 639 in_chans = encoder.patch_embed.seq[0].c.in_channels 640 641 # parameters for the decoder network 642 depth = 3 643 gain = 2 644 features_decoder = [self.initial_features * gain ** i for i in range(depth + 1)][::-1] 645 scale_factors = depth * [2] 646 self.out_channels = out_channels 647 648 # choice of upsampler - to use (bilinear interpolation + conv) or conv transpose 649 _upsampler = SingleDeconv2DBlock if use_conv_transpose else Upsampler2d 650 651 self.decoder = decoder or Decoder( 652 features=features_decoder, 653 scale_factors=scale_factors[::-1], 654 conv_block_impl=ConvBlock2d, 655 sampler_impl=_upsampler, 656 ) 657 658 if use_skip_connection: 659 self.deconv1 = Deconv2DBlock( 660 in_channels=self.embed_dim, 661 out_channels=features_decoder[0], 662 use_conv_transpose=use_conv_transpose, 663 ) 664 self.deconv2 = nn.Sequential( 665 Deconv2DBlock( 666 in_channels=self.embed_dim, 667 out_channels=features_decoder[0], 668 use_conv_transpose=use_conv_transpose, 669 ), 670 Deconv2DBlock( 671 in_channels=features_decoder[0], 672 out_channels=features_decoder[1], 673 use_conv_transpose=use_conv_transpose, 674 ) 675 ) 676 self.deconv3 = nn.Sequential( 677 Deconv2DBlock( 678 in_channels=self.embed_dim, 679 out_channels=features_decoder[0], 680 use_conv_transpose=use_conv_transpose, 681 ), 682 Deconv2DBlock( 683 in_channels=features_decoder[0], 684 out_channels=features_decoder[1], 685 use_conv_transpose=use_conv_transpose, 686 ), 687 Deconv2DBlock( 688 in_channels=features_decoder[1], 689 out_channels=features_decoder[2], 690 use_conv_transpose=use_conv_transpose, 691 ) 692 ) 693 self.deconv4 = ConvBlock2d(in_chans, features_decoder[-1]) 694 else: 695 self.deconv1 = Deconv2DBlock( 696 in_channels=self.embed_dim, 697 out_channels=features_decoder[0], 698 use_conv_transpose=use_conv_transpose, 699 ) 700 self.deconv2 = Deconv2DBlock( 701 in_channels=features_decoder[0], 702 out_channels=features_decoder[1], 703 use_conv_transpose=use_conv_transpose, 704 ) 705 self.deconv3 = Deconv2DBlock( 706 in_channels=features_decoder[1], 707 out_channels=features_decoder[2], 708 use_conv_transpose=use_conv_transpose, 709 ) 710 self.deconv4 = Deconv2DBlock( 711 in_channels=features_decoder[2], 712 out_channels=features_decoder[3], 713 use_conv_transpose=use_conv_transpose, 714 ) 715 716 self.base = ConvBlock2d(self.embed_dim, features_decoder[0]) 717 self.out_conv = nn.Conv2d(features_decoder[-1], out_channels, 1) 718 self.deconv_out = _upsampler( 719 scale_factor=2, in_channels=features_decoder[-1], out_channels=features_decoder[-1] 720 ) 721 self.decoder_head = ConvBlock2d(2 * features_decoder[-1], features_decoder[-1])
Initialize internal Module state, shared by both nn.Module and ScriptModule.
723 def encode(self, x: torch.Tensor): 724 """Preprocess the input and run the image encoder. 725 726 Args: 727 x: The input tensor of shape (B, C, Y, X). 728 729 Returns: 730 The features as a tuple of the image embeddings, the list of intermediate encoder outputs 731 (None if the encoder returns only the embeddings) and the preprocessed input, which the 732 skip connections consume, and the spatial shape after preprocessing. 733 """ 734 # Reshape the inputs to the shape expected by the encoder 735 # and normalize the inputs if normalization is part of the model. 736 x, input_shape = self.preprocess(x) 737 738 encoder_outputs = self.encoder(x) 739 740 if isinstance(encoder_outputs[-1], list): 741 # `encoder_outputs` can be arranged in only two forms: 742 # - either we only return the image embeddings 743 # - or, we return the image embeddings and the "list" of global attention layers 744 z12, from_encoder = encoder_outputs 745 else: 746 z12, from_encoder = encoder_outputs, None 747 748 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.
750 def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor: 751 """Run the convolutional decoder on the encoder features. 752 753 Args: 754 features: The tuple returned by `encode`. 755 input_shape: The spatial shape (Y, X) after preprocessing. 756 original_shape: The spatial shape (Y, X) of the original input. 757 758 Returns: 759 The UNETR output, resized to `original_shape`. 760 """ 761 z12, from_encoder, x = features 762 763 if self.use_skip_connection: 764 from_encoder = from_encoder[::-1] 765 z9 = self.deconv1(from_encoder[0]) 766 z6 = self.deconv2(from_encoder[1]) 767 z3 = self.deconv3(from_encoder[2]) 768 z0 = self.deconv4(x) 769 770 else: 771 z9 = self.deconv1(z12) 772 z6 = self.deconv2(z9) 773 z3 = self.deconv3(z6) 774 z0 = self.deconv4(z3) 775 776 updated_from_encoder = [z9, z6, z3] 777 778 x = self.base(z12) 779 x = self.decoder(x, encoder_inputs=updated_from_encoder) 780 x = self.deconv_out(x) 781 782 x = torch.cat([x, z0], dim=1) 783 x = self.decoder_head(x) 784 785 x = self.out_conv(x) 786 if self.final_activation is not None: 787 x = self.final_activation(x) 788 789 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.
792class UNETR2D(UNETR): 793 """A two-dimensional UNet Transformer using a vision transformer as encoder and a convolutional decoder. 794 """ 795 pass
A two-dimensional UNet Transformer using a vision transformer as encoder and a convolutional decoder.
Inherited Members
798class UNETR3D(UNETRBase): 799 """A three dimensional UNet Transformer using a vision transformer as encoder and a convolutional decoder. 800 """ 801 def __init__( 802 self, 803 img_size: int = 1024, 804 backbone: Literal[ 805 "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3", "torchvision" 806 ] = "sam", 807 encoder: Optional[Union[nn.Module, str]] = "hvit_b", 808 decoder: Optional[nn.Module] = None, 809 out_channels: int = 1, 810 use_sam_stats: bool = False, 811 use_mae_stats: bool = False, 812 use_dino_stats: bool = False, 813 use_imagenet_stats: bool = False, 814 resize_input: bool = True, 815 encoder_checkpoint: Optional[Union[str, OrderedDict]] = None, 816 final_activation: Optional[Union[str, nn.Module]] = None, 817 use_skip_connection: bool = False, 818 embed_dim: Optional[int] = None, 819 use_conv_transpose: bool = False, 820 use_strip_pooling: bool = True, 821 perform_range_checks: bool = True, 822 **kwargs 823 ): 824 if use_skip_connection: 825 raise NotImplementedError("The framework cannot handle skip connections atm.") 826 if use_conv_transpose: 827 raise NotImplementedError("It's not enabled to switch between interpolation and transposed convolutions.") 828 829 # Sort the `embed_dim` out 830 embed_dim = 256 if embed_dim is None else embed_dim 831 832 super().__init__( 833 img_size=img_size, 834 backbone=backbone, 835 encoder=encoder, 836 decoder=decoder, 837 out_channels=out_channels, 838 use_sam_stats=use_sam_stats, 839 use_mae_stats=use_mae_stats, 840 use_dino_stats=use_dino_stats, 841 use_imagenet_stats=use_imagenet_stats, 842 resize_input=resize_input, 843 encoder_checkpoint=encoder_checkpoint, 844 final_activation=final_activation, 845 use_skip_connection=use_skip_connection, 846 embed_dim=embed_dim, 847 use_conv_transpose=use_conv_transpose, 848 perform_range_checks=perform_range_checks, 849 **kwargs, 850 ) 851 self.init_kwargs = { 852 "img_size": img_size, "backbone": backbone, "encoder": encoder, "decoder": decoder, 853 "out_channels": out_channels, "use_sam_stats": use_sam_stats, 854 "use_mae_stats": use_mae_stats, "use_dino_stats": use_dino_stats, 855 "use_imagenet_stats": use_imagenet_stats, "resize_input": resize_input, 856 "final_activation": final_activation, 857 "use_skip_connection": use_skip_connection, "embed_dim": embed_dim, 858 "use_conv_transpose": use_conv_transpose, "use_strip_pooling": use_strip_pooling, 859 "perform_range_checks": perform_range_checks, **kwargs, 860 } 861 862 # The 3d convolutional decoder. 863 # First, get the important parameters for the decoder. 864 depth = 3 865 gain = 2 866 features_decoder = [self.initial_features * gain ** i for i in range(depth + 1)][::-1] 867 scale_factors = [1, 2, 2] 868 self.out_channels = out_channels 869 870 # The mapping blocks. 871 self.deconv1 = Deconv3DBlock( 872 in_channels=embed_dim, 873 out_channels=features_decoder[0], 874 scale_factor=scale_factors, 875 use_strip_pooling=use_strip_pooling, 876 ) 877 self.deconv2 = Deconv3DBlock( 878 in_channels=features_decoder[0], 879 out_channels=features_decoder[1], 880 scale_factor=scale_factors, 881 use_strip_pooling=use_strip_pooling, 882 ) 883 self.deconv3 = Deconv3DBlock( 884 in_channels=features_decoder[1], 885 out_channels=features_decoder[2], 886 scale_factor=scale_factors, 887 use_strip_pooling=use_strip_pooling, 888 ) 889 self.deconv4 = Deconv3DBlock( 890 in_channels=features_decoder[2], 891 out_channels=features_decoder[3], 892 scale_factor=scale_factors, 893 use_strip_pooling=use_strip_pooling, 894 ) 895 896 # The core decoder block. 897 self.decoder = decoder or Decoder( 898 features=features_decoder, 899 scale_factors=[scale_factors] * depth, 900 conv_block_impl=partial(ConvBlock3dWithStrip, use_strip_pooling=use_strip_pooling), 901 sampler_impl=Upsampler3d, 902 ) 903 904 # And the final upsampler to match the expected dimensions. 905 self.deconv_out = Deconv3DBlock( # NOTE: changed `end_up` to `deconv_out` 906 in_channels=features_decoder[-1], 907 out_channels=features_decoder[-1], 908 scale_factor=scale_factors, 909 use_strip_pooling=use_strip_pooling, 910 ) 911 912 # Additional conjunction blocks. 913 self.base = ConvBlock3dWithStrip( 914 in_channels=embed_dim, 915 out_channels=features_decoder[0], 916 use_strip_pooling=use_strip_pooling, 917 ) 918 919 # And the output layers. 920 self.decoder_head = ConvBlock3dWithStrip( 921 in_channels=2 * features_decoder[-1], 922 out_channels=features_decoder[-1], 923 use_strip_pooling=use_strip_pooling, 924 ) 925 self.out_conv = nn.Conv3d(features_decoder[-1], out_channels, 1) 926 927 def encode(self, x: torch.Tensor): 928 """Preprocess the input and run the image encoder on every z-slice. 929 930 Args: 931 x: Inputs of expected shape (B, C, Z, Y, X), where Z considers flexible inputs. 932 933 Returns: 934 The encoder features of shape (B, D, Z, Y', X') and the spatial shape after preprocessing. 935 """ 936 Z = x.shape[2] 937 x, input_shape = self.preprocess(x) 938 features = torch.stack([self.encoder(x[:, :, i])[0] for i in range(Z)], dim=2) 939 return features, input_shape 940 941 def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor: 942 """Run the convolutional decoder on the encoder features. 943 944 Args: 945 features: Encoder features of shape (B, D, Z, Y', X'), see `encode`. 946 input_shape: The spatial shape (Z, Y, X) after preprocessing. 947 original_shape: The spatial shape (Z, Y, X) of the original input. 948 949 Returns: 950 The UNETR output, resized to `original_shape`. 951 """ 952 # Prepare the counterparts for the decoder. 953 # NOTE: The section below is sequential, there's no skip connections atm. 954 z9 = self.deconv1(features) 955 z6 = self.deconv2(z9) 956 z3 = self.deconv3(z6) 957 z0 = self.deconv4(z3) 958 959 updated_from_encoder = [z9, z6, z3] 960 961 # Align the features through the base block. 962 x = self.base(features) 963 # Run the decoder 964 x = self.decoder(x, encoder_inputs=updated_from_encoder) 965 x = self.deconv_out(x) # NOTE before `end_up` 966 967 # And the final output head. 968 x = torch.cat([x, z0], dim=1) 969 x = self.decoder_head(x) 970 x = self.out_conv(x) 971 if self.final_activation is not None: 972 x = self.final_activation(x) 973 974 # Postprocess the output back to original size. 975 return self.postprocess_masks(x, input_shape, original_shape)
A three dimensional UNet Transformer using a vision transformer as encoder and a convolutional decoder.
801 def __init__( 802 self, 803 img_size: int = 1024, 804 backbone: Literal[ 805 "sam", "sam2", "sam3", "cellpose_sam", "mae", "scalemae", "dinov2", "dinov3", "torchvision" 806 ] = "sam", 807 encoder: Optional[Union[nn.Module, str]] = "hvit_b", 808 decoder: Optional[nn.Module] = None, 809 out_channels: int = 1, 810 use_sam_stats: bool = False, 811 use_mae_stats: bool = False, 812 use_dino_stats: bool = False, 813 use_imagenet_stats: bool = False, 814 resize_input: bool = True, 815 encoder_checkpoint: Optional[Union[str, OrderedDict]] = None, 816 final_activation: Optional[Union[str, nn.Module]] = None, 817 use_skip_connection: bool = False, 818 embed_dim: Optional[int] = None, 819 use_conv_transpose: bool = False, 820 use_strip_pooling: bool = True, 821 perform_range_checks: bool = True, 822 **kwargs 823 ): 824 if use_skip_connection: 825 raise NotImplementedError("The framework cannot handle skip connections atm.") 826 if use_conv_transpose: 827 raise NotImplementedError("It's not enabled to switch between interpolation and transposed convolutions.") 828 829 # Sort the `embed_dim` out 830 embed_dim = 256 if embed_dim is None else embed_dim 831 832 super().__init__( 833 img_size=img_size, 834 backbone=backbone, 835 encoder=encoder, 836 decoder=decoder, 837 out_channels=out_channels, 838 use_sam_stats=use_sam_stats, 839 use_mae_stats=use_mae_stats, 840 use_dino_stats=use_dino_stats, 841 use_imagenet_stats=use_imagenet_stats, 842 resize_input=resize_input, 843 encoder_checkpoint=encoder_checkpoint, 844 final_activation=final_activation, 845 use_skip_connection=use_skip_connection, 846 embed_dim=embed_dim, 847 use_conv_transpose=use_conv_transpose, 848 perform_range_checks=perform_range_checks, 849 **kwargs, 850 ) 851 self.init_kwargs = { 852 "img_size": img_size, "backbone": backbone, "encoder": encoder, "decoder": decoder, 853 "out_channels": out_channels, "use_sam_stats": use_sam_stats, 854 "use_mae_stats": use_mae_stats, "use_dino_stats": use_dino_stats, 855 "use_imagenet_stats": use_imagenet_stats, "resize_input": resize_input, 856 "final_activation": final_activation, 857 "use_skip_connection": use_skip_connection, "embed_dim": embed_dim, 858 "use_conv_transpose": use_conv_transpose, "use_strip_pooling": use_strip_pooling, 859 "perform_range_checks": perform_range_checks, **kwargs, 860 } 861 862 # The 3d convolutional decoder. 863 # First, get the important parameters for the decoder. 864 depth = 3 865 gain = 2 866 features_decoder = [self.initial_features * gain ** i for i in range(depth + 1)][::-1] 867 scale_factors = [1, 2, 2] 868 self.out_channels = out_channels 869 870 # The mapping blocks. 871 self.deconv1 = Deconv3DBlock( 872 in_channels=embed_dim, 873 out_channels=features_decoder[0], 874 scale_factor=scale_factors, 875 use_strip_pooling=use_strip_pooling, 876 ) 877 self.deconv2 = Deconv3DBlock( 878 in_channels=features_decoder[0], 879 out_channels=features_decoder[1], 880 scale_factor=scale_factors, 881 use_strip_pooling=use_strip_pooling, 882 ) 883 self.deconv3 = Deconv3DBlock( 884 in_channels=features_decoder[1], 885 out_channels=features_decoder[2], 886 scale_factor=scale_factors, 887 use_strip_pooling=use_strip_pooling, 888 ) 889 self.deconv4 = Deconv3DBlock( 890 in_channels=features_decoder[2], 891 out_channels=features_decoder[3], 892 scale_factor=scale_factors, 893 use_strip_pooling=use_strip_pooling, 894 ) 895 896 # The core decoder block. 897 self.decoder = decoder or Decoder( 898 features=features_decoder, 899 scale_factors=[scale_factors] * depth, 900 conv_block_impl=partial(ConvBlock3dWithStrip, use_strip_pooling=use_strip_pooling), 901 sampler_impl=Upsampler3d, 902 ) 903 904 # And the final upsampler to match the expected dimensions. 905 self.deconv_out = Deconv3DBlock( # NOTE: changed `end_up` to `deconv_out` 906 in_channels=features_decoder[-1], 907 out_channels=features_decoder[-1], 908 scale_factor=scale_factors, 909 use_strip_pooling=use_strip_pooling, 910 ) 911 912 # Additional conjunction blocks. 913 self.base = ConvBlock3dWithStrip( 914 in_channels=embed_dim, 915 out_channels=features_decoder[0], 916 use_strip_pooling=use_strip_pooling, 917 ) 918 919 # And the output layers. 920 self.decoder_head = ConvBlock3dWithStrip( 921 in_channels=2 * features_decoder[-1], 922 out_channels=features_decoder[-1], 923 use_strip_pooling=use_strip_pooling, 924 ) 925 self.out_conv = nn.Conv3d(features_decoder[-1], out_channels, 1)
Initialize internal Module state, shared by both nn.Module and ScriptModule.
927 def encode(self, x: torch.Tensor): 928 """Preprocess the input and run the image encoder on every z-slice. 929 930 Args: 931 x: Inputs of expected shape (B, C, Z, Y, X), where Z considers flexible inputs. 932 933 Returns: 934 The encoder features of shape (B, D, Z, Y', X') and the spatial shape after preprocessing. 935 """ 936 Z = x.shape[2] 937 x, input_shape = self.preprocess(x) 938 features = torch.stack([self.encoder(x[:, :, i])[0] for i in range(Z)], dim=2) 939 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.
941 def decode(self, features, input_shape: Tuple[int, ...], original_shape: Tuple[int, ...]) -> torch.Tensor: 942 """Run the convolutional decoder on the encoder features. 943 944 Args: 945 features: Encoder features of shape (B, D, Z, Y', X'), see `encode`. 946 input_shape: The spatial shape (Z, Y, X) after preprocessing. 947 original_shape: The spatial shape (Z, Y, X) of the original input. 948 949 Returns: 950 The UNETR output, resized to `original_shape`. 951 """ 952 # Prepare the counterparts for the decoder. 953 # NOTE: The section below is sequential, there's no skip connections atm. 954 z9 = self.deconv1(features) 955 z6 = self.deconv2(z9) 956 z3 = self.deconv3(z6) 957 z0 = self.deconv4(z3) 958 959 updated_from_encoder = [z9, z6, z3] 960 961 # Align the features through the base block. 962 x = self.base(features) 963 # Run the decoder 964 x = self.decoder(x, encoder_inputs=updated_from_encoder) 965 x = self.deconv_out(x) # NOTE before `end_up` 966 967 # And the final output head. 968 x = torch.cat([x, z0], dim=1) 969 x = self.decoder_head(x) 970 x = self.out_conv(x) 971 if self.final_activation is not None: 972 x = self.final_activation(x) 973 974 # Postprocess the output back to original size. 975 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.