torch_em.model.vit
1import math 2from functools import partial 3from typing import Tuple, List 4 5import torch 6import torch.nn as nn 7import torch.nn.functional as F 8 9# we catch ImportErrors here because segment_anything, micro_sam, scale_mae and timm should 10# only be optional dependencies for torch_em 11try: 12 from segment_anything.modeling import ImageEncoderViT 13 _sam_import_success = True 14except ImportError: 15 ImageEncoderViT = object 16 _sam_import_success = False 17 18try: 19 from timm.models.vision_transformer import VisionTransformer, PatchEmbed 20 _timm_import_success = True 21except ImportError: 22 VisionTransformer = object 23 PatchEmbed = object 24 _timm_import_success = False 25 26try: 27 from sam2.modeling.backbones.hieradet import Hiera 28 from sam2.modeling.position_encoding import PositionEmbeddingSine 29 from sam2.modeling.backbones.image_encoder import ImageEncoder, FpnNeck 30 _sam2_import_success = True 31except ImportError: 32 ImageEncoder = object 33 _sam2_import_success = False 34 35try: 36 from dinov2.models.vision_transformer import DinoVisionTransformer as DinoV2VisionTransformer 37 from dinov2.layers import MemEffAttention, NestedTensorBlock as Block 38 _dinov2_import_success = True 39except ImportError: 40 DinoV2VisionTransformer = object 41 _dinov2_import_success = False 42 43try: 44 from dinov3.models.vision_transformer import DinoVisionTransformer as DinoV3VisionTransformer 45 _dinov3_import_success = True 46except ImportError: 47 DinoV3VisionTransformer = object 48 _dinov3_import_success = False 49 50 51try: 52 from sam3.model.vitdet import ViT as SAM3ViT, get_abs_pos 53 _sam3_import_success = True 54except ImportError: 55 SAM3ViT = object 56 _sam3_import_success = False 57 58try: 59 import torchvision.models as _tv_models 60 _torchvision_import_success = True 61except ImportError: 62 _tv_models = None 63 _torchvision_import_success = False 64 65 66class ViT_Sam(ImageEncoderViT): 67 """Vision Transformer derived from the Segment Anything Codebase (https://arxiv.org/abs/2304.02643). 68 69 Based on: 70 https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/image_encoder.py 71 72 Args: 73 in_chans: The number of input channels. 74 embed_dim: The embedding dimension, corresponding to the number of output channels of the vision transformer. 75 global_attn_indexes: The global attention indices. 76 apply_neck: Whether to apply the convolutional bottleneck after outputs of the last attention head. 77 kwargs: Keyword arguments for the image encoder base class. 78 """ 79 def __init__( 80 self, 81 in_chans: int = 3, 82 embed_dim: int = 768, 83 global_attn_indexes: Tuple[int, ...] = [2, 5, 8, 11], 84 apply_neck: bool = False, 85 **kwargs, 86 ) -> None: 87 if not _sam_import_success: 88 raise RuntimeError( 89 "The vision transformer backend can only be initialized if segment anything is installed. " 90 "Please install segment anything from https://github.com/facebookresearch/segment-anything " 91 "and then rerun your code." 92 ) 93 94 super().__init__(embed_dim=embed_dim, global_attn_indexes=global_attn_indexes, **kwargs) 95 self.chunks_for_projection = global_attn_indexes 96 self.in_chans = in_chans 97 self.embed_dim = embed_dim 98 self.apply_neck = apply_neck 99 100 def forward(self, x: torch.Tensor) -> torch.Tensor: 101 """Apply the vision transformer to input data. 102 103 Args: 104 x: The input data. 105 106 Returns: 107 The vision transformer output. 108 """ 109 x = self.patch_embed(x) 110 if self.pos_embed is not None: 111 x = x + self.pos_embed 112 113 list_from_encoder = [] 114 for i, blk in enumerate(self.blocks): 115 x = blk(x) 116 if i in self.chunks_for_projection: 117 list_from_encoder.append(x) 118 119 x = x.permute(0, 3, 1, 2) 120 121 if self.apply_neck: 122 x = self.neck(x) 123 124 list_from_encoder = [e.permute(0, 3, 1, 2) for e in list_from_encoder] 125 return x, list_from_encoder[:3] 126 127 128class ViT_CellposeSAM(nn.Module): 129 """Vision Transformer derived from the CellposeSAM Codebase (https://doi.org/10.1038/s41592-025-02595-x). 130 131 This replicates CellposeSAM's actual initialization: instantiate SAM's ``ImageEncoderViT`` via 132 ``sam_model_registry``, then modify the patch embedding, position embeddings, and set global attention. 133 This preserves SAM's original relative position bias sizes, enabling direct checkpoint loading 134 without any interpolation. 135 136 Based on: https://github.com/MouseLand/cellpose/blob/main/cellpose/vit_sam.py 137 138 NOTE: The pretrained CellposeSAM model uses ``vit_l`` exclusively. 139 140 Args: 141 ps: The patch size (default for CellposeSAM is 8). 142 bsize: The input image size (default for CellposeSAM is 256). 143 apply_neck: Whether to apply the convolutional bottleneck after outputs of the last attention head. 144 """ 145 def __init__(self, ps: int = 8, bsize: int = 256, apply_neck: bool = True) -> None: 146 super().__init__() 147 148 if not _sam_import_success: 149 raise RuntimeError( 150 "The vision transformer backend can only be initialized if segment anything is installed. " 151 "Please install segment anything from https://github.com/facebookresearch/segment-anything " 152 "and then rerun your code." 153 ) 154 155 from segment_anything import sam_model_registry 156 157 # Creates the SAM vit_l encoder and applies CellposeSAM's modifications (same as cellpose.vit_sam.Transformer). 158 encoder = sam_model_registry["vit_l"](None).image_encoder 159 160 w = encoder.patch_embed.proj.weight.detach() 161 nchan = w.shape[0] 162 163 # CellPoseSAM changes the patch size from 16 to 'ps'. 164 encoder.patch_embed.proj = nn.Conv2d(3, nchan, stride=ps, kernel_size=ps) 165 encoder.patch_embed.proj.weight.data = w[:, :, ::16 // ps, ::16 // ps] 166 167 # Next, they subsample position embeddings for the new patch size and input resolution. 168 ds = (1024 // 16) // (bsize // ps) 169 encoder.pos_embed = nn.Parameter(encoder.pos_embed[:, ::ds, ::ds], requires_grad=True) 170 171 # Finally, they set all blocks to global attention. 172 for blk in encoder.blocks: 173 blk.window_size = 0 174 175 # Store encoder submodules directly ('state_dict' keys match CellposeSAM after prefix stripping). 176 self.patch_embed = encoder.patch_embed 177 self.pos_embed = encoder.pos_embed 178 self.blocks = encoder.blocks 179 self.neck = encoder.neck 180 181 # Additional attributes expected by UNETR. 182 self.embed_dim = nchan 183 self.img_size = bsize 184 self.in_chans = 3 185 self.apply_neck = apply_neck 186 187 # Feature extraction at evenly-spaced depths. 188 depth = len(self.blocks) 189 _chunks = depth // 4 190 self.chunks_for_projection = [_chunks - 1, 2 * _chunks - 1, 3 * _chunks - 1, 4 * _chunks - 1] 191 192 def forward(self, x: torch.Tensor) -> torch.Tensor: 193 """Apply the vision transformer to input data. 194 195 Args: 196 x: The input data. 197 198 Returns: 199 The vision transformer output. 200 """ 201 x = self.patch_embed(x) 202 if self.pos_embed is not None: 203 x = x + self.pos_embed 204 205 list_from_encoder = [] 206 for i, blk in enumerate(self.blocks): 207 x = blk(x) 208 if i in self.chunks_for_projection: 209 list_from_encoder.append(x) 210 211 x = x.permute(0, 3, 1, 2) 212 213 if self.apply_neck: 214 x = self.neck(x) 215 216 list_from_encoder = [e.permute(0, 3, 1, 2) for e in list_from_encoder] 217 return x, list_from_encoder[:3] 218 219 220class ViT_MAE(VisionTransformer): 221 """Vision Transformer derived from the Masked Auto Encoder Codebase (https://arxiv.org/abs/2111.06377). 222 223 Based on: 224 https://github.com/facebookresearch/mae/blob/main/models_vit.py#L20-L53 225 226 Args: 227 img_size: The size of the input for the image encoder. Input images will be resized to match this size. 228 in_chans: The number of input channels. 229 depth: The depth of the vision transformer. 230 kwargs: Additional keyword arguments for the vision transformer base class. 231 """ 232 def __init__( 233 self, 234 img_size: int = 1024, # chosen to match our experiments with segment anything 235 in_chans: int = 3, 236 depth: int = 12, 237 **kwargs 238 ): 239 if not _timm_import_success: 240 raise RuntimeError( 241 "The vision transformer backend can only be initialized if timm is installed. " 242 "Please install timm (using conda/mamba) for using https://github.com/facebookresearch/mae/ " 243 "and then rerun your code" 244 ) 245 super().__init__(img_size=img_size, depth=depth, **kwargs) 246 self.img_size = img_size 247 self.in_chans = in_chans 248 self.depth = depth 249 250 def convert_to_expected_dim(self, inputs_): 251 """@private 252 """ 253 inputs_ = inputs_[:, 1:, :] # removing the class tokens 254 # reshape the outputs to desired shape (N x H*W X C -> N x H x W x C) 255 rdim = inputs_.shape[1] 256 dshape = int(rdim ** 0.5) # finding the square root of the outputs for obtaining the patch shape 257 inputs_ = torch.unflatten(inputs_, 1, (dshape, dshape)) 258 inputs_ = inputs_.permute(0, 3, 1, 2) 259 return inputs_ 260 261 def forward_features(self, x): 262 """@private 263 """ 264 B = x.shape[0] 265 x = self.patch_embed(x) 266 267 cls_tokens = self.cls_token.expand(B, -1, -1) 268 x = torch.cat((cls_tokens, x), dim=1) 269 270 x = x + self.pos_embed 271 x = self.pos_drop(x) 272 273 # chunks obtained for getting the projections for conjuctions with upsampling blocks 274 _chunks = int(self.depth / 4) 275 chunks_for_projection = [_chunks - 1, 2*_chunks - 1, 3*_chunks - 1, 4*_chunks - 1] 276 277 list_from_encoder = [] 278 for i, blk in enumerate(self.blocks): 279 x = blk(x) 280 if i in chunks_for_projection: 281 list_from_encoder.append(self.convert_to_expected_dim(x)) 282 283 x = self.convert_to_expected_dim(x) 284 return x, list_from_encoder[:3] 285 286 def forward(self, x: torch.Tensor) -> torch.Tensor: 287 """Apply the vision transformer to input data. 288 289 Args: 290 x: The input data. 291 292 Returns: 293 The vision transformer output. 294 """ 295 x, list_from_encoder = self.forward_features(x) 296 return x, list_from_encoder 297 298 299class ViT_Sam2(ImageEncoder): 300 """Vision Transformer derived from the Segment Anything 2 Codebase (https://arxiv.org/abs/2408.00714). 301 302 Based on https://github.com/facebookresearch/sam2/blob/main/sam2/modeling/backbones/image_encoder.py. 303 304 Args: 305 backbone_channel_list: The channels throughout the entire backbone. 306 embed_dim: The initial embedding dimension. 307 num_heads: The initial number of heads. 308 stages: The number of blocks per stage. 309 global_att_blocks: The parameter to decide which blocks have global attention. 310 window_pos_embed_bkg_spatial_size: The spatial size of windowed positional embedding. 311 window_spec: The window size per stage, when not using global attention. 312 scalp: The count of lowest resolution features to discard. 313 """ 314 def __init__( 315 self, 316 backbone_channel_list: List[int], 317 img_size: int = 1024, 318 embed_dim: int = 96, 319 num_heads: int = 1, 320 stages: Tuple[int, ...] = (2, 3, 16, 3), 321 global_att_blocks: Tuple[int, ...] = (12, 16, 20), 322 window_pos_embed_bkg_spatial_size: Tuple[int, int] = (14, 14), 323 window_spec: Tuple[int, ...] = (8, 4, 14, 7), 324 scalp: int = 1, 325 **kwargs 326 ): 327 if not _sam2_import_success: 328 raise RuntimeError( 329 "The vision transformer backend can only be initialized if segment anything 2 is installed. " 330 "Please install segment anything 2 from https://github.com/facebookresearch/sam2 " 331 "and then rerun your code" 332 ) 333 334 trunk = Hiera( 335 embed_dim=embed_dim, 336 num_heads=num_heads, 337 stages=stages, 338 global_att_blocks=global_att_blocks, 339 window_pos_embed_bkg_spatial_size=window_pos_embed_bkg_spatial_size, 340 window_spec=window_spec, 341 ) 342 neck = FpnNeck( 343 position_encoding=PositionEmbeddingSine(num_pos_feats=256), 344 d_model=256, 345 backbone_channel_list=backbone_channel_list, 346 fpn_top_down_levels=[2, 3], 347 fpn_interp_model="nearest", 348 ) 349 350 super().__init__(trunk=trunk, neck=neck, scalp=scalp, **kwargs) 351 self.scalp = scalp 352 self.embed_dim = embed_dim 353 self.img_size = img_size 354 355 def forward(self, x: torch.Tensor): 356 # The forward pass throught the backbone. 357 features, pos = self.neck(self.trunk(x)) 358 if self.scalp > 0: # This discard the "n" lowest resolution features. 359 features, pos = features[:-self.scalp], pos[:-self.scalp] 360 361 return features[-1], features 362 363 364class ViT_Sam3(SAM3ViT): 365 """Vision Transformer derived from the Segment Anything 3 Codebase (https://arxiv.org/abs/2511.16719). 366 367 Based on https://github.com/facebookresearch/sam3/blob/main/sam3/model/vitdet.py. 368 369 Args: 370 img_size: The input image size. 371 embed_dim: The embedding dimension, corresponding to the number of output channels of the vision transformer. 372 kwargs: Keyword arguments for the image encoder base class. 373 """ 374 def __init__(self, img_size: int = 1024, embed_dim: int = 768, **kwargs): 375 if not _sam3_import_success: 376 raise RuntimeError( 377 "The vision transformer backend can only be initialized if segment anything 3 is installed. " 378 "Please install segment anything 3 from https://github.com/facebookresearch/sam3 " 379 "and then rerun your code" 380 ) 381 382 super().__init__(img_size=img_size, embed_dim=embed_dim, **kwargs) 383 self.img_size = img_size 384 self.embed_dim = embed_dim 385 386 def forward_features(self, x): 387 """@private 388 """ 389 x = self.patch_embed(x) 390 h, w = x.shape[1], x.shape[2] 391 392 s = 0 393 if self.retain_cls_token: 394 # If the 'cls_token' is retained, we don't maintain the spatial shape. 395 x = torch.cat([self.class_embedding, x.flatten(1, 2)], dim=1) 396 s = 1 397 398 if self.pos_embed is not None: 399 x = x + get_abs_pos( 400 self.pos_embed, self.pretrain_use_cls_token, (h, w), self.retain_cls_token, tiling=self.tile_abs_pos, 401 ) 402 403 x = self.ln_pre(x) 404 405 list_from_encoder = [] 406 for i, blk in enumerate(self.blocks): 407 if self.use_act_checkpoint and self.training: 408 x = torch.utils.checkpoint.checkpoint(blk, x, use_reentrant=False) 409 else: 410 x = blk(x) 411 412 x = self._convert_to_expected_dim(x, i, s) 413 414 if i in self.full_attn_ids: 415 list_from_encoder.append(x) 416 417 return x, list_from_encoder 418 419 def _convert_to_expected_dim(self, x, i, s): 420 if (i == self.full_attn_ids[-1]) or ( 421 self.return_interm_layers and i in self.full_attn_ids 422 ): 423 if i == self.full_attn_ids[-1]: 424 x = self.ln_post(x) 425 426 feats = x[:, s:] 427 if feats.ndim == 4: 428 feats = feats.permute(0, 3, 1, 2) 429 else: 430 assert feats.ndim == 3 431 h = w = math.sqrt(feats.shape[1]) 432 feats = feats.reshape(feats.shape[0], h, w, feats.shape[-1]).permute(0, 3, 1, 2) 433 return feats 434 435 else: 436 return x 437 438 def forward(self, x: torch.Tensor): 439 """Apply the vision transformer to input data. 440 441 Args: 442 x: The input data. 443 444 Returns: 445 The vision transformer output. 446 """ 447 x, list_from_encoder = self.forward_features(x) 448 return x, list_from_encoder 449 450# 451# Utilities for ScaleMAE's ViT 452# 453 454 455class CustomCompose: 456 def __init__(self, rescale_transform, other_transforms, src_transform): 457 self.rescale_transform = rescale_transform 458 self.other_transforms = other_transforms 459 self.src_transform = src_transform 460 461 def __call__(self, x, valid_masks=None): 462 if valid_masks is not None: 463 nodata = (x * (1 - valid_masks.float())).max() 464 x_aug = self.rescale_transform(x) 465 parms = self.rescale_transform._params 466 467 # sanity check, comment if this is working 468 # valid_masks = self.rescale_transform(valid_masks.float(), params=parms) 469 # assert (x_aug==self.rescale_transform(x, params=parms)).all() # 470 471 if valid_masks is not None: 472 valid_masks = x_aug != nodata 473 _, c, h, w = x_aug.shape 474 zero_ratio = ((valid_masks == 0).sum((1, 2, 3)) / (h * w * c)).cpu().numpy() 475 else: 476 zero_ratio = -1 477 478 if self.other_transforms: 479 x_aug = self.other_transforms(x_aug) 480 x_src = self.src_transform(x_aug) 481 dx = parms["src"][:, 1, 0] - parms["src"][:, 0, 0] 482 483 # dy = (parms['src'][:,2,1] - parms['src'][:,1,1]) 484 # assert (dx == dy).all() 485 486 h, w = x_aug.shape[-2:] 487 # assert h == w 488 489 return x_aug, x_src, dx / h, zero_ratio, valid_masks 490 491 492def get_2d_sincos_pos_embed_with_resolution(embed_dim, grid_size, res, cls_token=False, device="cpu"): 493 """ 494 grid_size: int of the grid height and width 495 res: array of size n, representing the resolution of a pixel (say, in meters), 496 return: 497 pos_embed: [n,grid_size*grid_size, embed_dim] or [n,1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token) 498 """ 499 # res = torch.FloatTensor(res).to(device) 500 res = res.to(device) 501 grid_h = torch.arange(grid_size, dtype=torch.float32, device=device) 502 grid_w = torch.arange(grid_size, dtype=torch.float32, device=device) 503 grid = torch.meshgrid(grid_w, grid_h, indexing="xy") # here h goes first,direction reversed for numpy 504 grid = torch.stack(grid, dim=0) # 2 x h x w 505 506 # grid = grid.reshape([2, 1, grid_size, grid_size]) 507 grid = torch.einsum("chw,n->cnhw", grid, res) # 2 x n x h x w 508 _, n, h, w = grid.shape 509 pos_embed = get_2d_sincos_pos_embed_from_grid_torch(embed_dim, grid) # (nxH*W, D/2) 510 pos_embed = pos_embed.reshape(n, h * w, embed_dim) 511 if cls_token: 512 pos_embed = torch.cat( 513 [torch.zeros([n, 1, embed_dim], dtype=torch.float32, device=pos_embed.device), pos_embed], dim=1 514 ) 515 516 return pos_embed 517 518 519def get_2d_sincos_pos_embed_from_grid_torch(embed_dim, grid): 520 assert embed_dim % 2 == 0 521 522 # use half of dimensions to encode grid_h 523 emb_h = get_1d_sincos_pos_embed_from_grid_torch(embed_dim // 2, grid[0]) # (H*W, D/2) 524 emb_w = get_1d_sincos_pos_embed_from_grid_torch(embed_dim // 2, grid[1]) # (H*W, D/2) 525 526 emb = torch.cat([emb_h, emb_w], dim=1) # (H*W, D) 527 return emb 528 529 530def get_1d_sincos_pos_embed_from_grid_torch(embed_dim, pos): 531 """ 532 embed_dim: output dimension for each position 533 pos: a list of positions to be encoded: size (M,) 534 out: (M, D) 535 """ 536 assert embed_dim % 2 == 0 537 # old_shape = pos 538 omega = torch.arange(embed_dim // 2, dtype=torch.float32, device=pos.device) 539 omega /= embed_dim / 2.0 540 omega = 1.0 / 10000**omega # (D/2,) 541 542 pos = pos.reshape(-1) # (M,) 543 out = torch.einsum("m,d->md", pos, omega) # (M, D/2), outer product 544 545 emb_sin = torch.sin(out) # (M, D/2) 546 emb_cos = torch.cos(out) # (M, D/2) 547 548 emb = torch.cat([emb_sin, emb_cos], dim=1) # (M, D) 549 return emb 550 551 552class PatchEmbedUnSafe(PatchEmbed): 553 """Image to Patch Embedding""" 554 555 def forward(self, x): 556 B, C, H, W = x.shape 557 558 # NOTE: Comment code from ScaleMAE: Dropped size check in timm 559 # assert H == self.img_size[0] and W == self.img_size[1], \ 560 # f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]})." 561 562 x = self.proj(x).flatten(2).transpose(1, 2) 563 return x 564 565 566class ViT_ScaleMAE(VisionTransformer): 567 """Vision Transformer dervied from the Scale Masked Auto Encoder codebase (TODO: paper and github link). 568 569 NOTE: For downstream tasks, the "base_resoulution" parameter needs to be adjusted manually when using 570 the model on a different zoom factor dataset. 571 """ 572 573 def __init__( 574 self, img_size=224, patch_size=16, in_chans=3, embed_dim=1024, depth=12, base_resolution=2.5, **kwargs 575 ): 576 super().__init__(img_size=img_size, embed_dim=embed_dim, **kwargs) 577 self.img_size = img_size 578 self.in_chans = in_chans 579 self.depth = depth 580 self.base_resolution = base_resolution 581 582 self.patch_embed = PatchEmbedUnSafe( 583 img_size=img_size, 584 patch_size=patch_size, 585 in_chans=in_chans, 586 embed_dim=embed_dim, 587 ) 588 589 def transform_inputs(self, x): 590 import kornia.augmentation as K 591 from kornia.constants import Resample 592 593 self._transforms = CustomCompose( 594 rescale_transform=K.RandomResizedCrop( 595 (448, 448), 596 ratio=(1.0, 1.0), 597 scale=(1.0, 1.0), 598 resample=Resample.BICUBIC.name, 599 ), 600 other_transforms=None, 601 src_transform=K.Resize((224, 224)), 602 ) 603 x, _, ratios, _, _ = self._transforms(x) 604 input_res = ratios * self.base_resolution 605 return x, input_res 606 607 def convert_to_expected_dim(self, x): 608 inputs_ = x[:, 1:, :] # removing the class tokens 609 # reshape the outputs to desired shape (N X H*W X C -> N X H X W X C) 610 rdim = inputs_.shape[1] 611 dshape = int(rdim ** 0.5) # finding square root of the outputs for obtaining the patch shape 612 inputs_ = torch.unflatten(inputs_, 1, (dshape, dshape)) 613 inputs_ = inputs_.permute(0, 3, 1, 2) 614 return inputs_ 615 616 def forward_features(self, x): 617 x, input_res = self.transform_inputs(x) 618 619 B, _, h, w = x.shape 620 x = self.patch_embed(x) 621 622 num_patches = int((h * w) / (self.patch_embed.patch_size[0] * self.patch_embed.patch_size[1])) 623 pos_embed = get_2d_sincos_pos_embed_with_resolution( 624 x.shape[-1], 625 int(num_patches ** 0.5), 626 input_res, 627 cls_token=True, 628 device=x.device, 629 ) 630 631 cls_tokens = self.cls_token.expand(B, -1, -1) # stole cls_tokens impl from Phil Wang, thanks 632 x = torch.cat((cls_tokens, x), dim=1) 633 x = x + pos_embed 634 x = self.pos_drop(x) 635 636 # chunks obtained for getting the projections for conjuctions with upsampling blocks 637 _chunks = int(self.depth / 4) 638 chunks_for_projection = [_chunks - 1, 2*_chunks - 1, 3*_chunks - 1, 4*_chunks - 1] 639 640 list_from_encoder = [] 641 for i, blk in enumerate(self.blocks): 642 x = blk(x) 643 if i in chunks_for_projection: 644 list_from_encoder.append(self.convert_to_expected_dim(x)) 645 646 x = self.convert_to_expected_dim(x) 647 648 return x, list_from_encoder 649 650 def forward(self, x): 651 x, list_from_encoder = self.forward_features(x) 652 return x, list_from_encoder 653 654 655class ViT_DINOv2(DinoV2VisionTransformer): 656 """Vision Transformer derived from the DINOv2 Codebase (https://arxiv.org/abs/2304.07193). 657 658 Based on: 659 https://github.com/facebookresearch/dinov2/blob/main/dinov2/models/vision_transformer.py. 660 661 Args: 662 img_size: The input image size. 663 patch_size: The patch size. 664 depth: The depth of the network. 665 num_register_tokens: The number of registers added (in addition to the class tokens). 666 It's important to know for ViTs trained with registers, to remove them at the end. 667 """ 668 def __init__( 669 self, 670 img_size: int = 224, 671 patch_size: int = 16, 672 depth: int = 12, 673 num_register_tokens: int = 0, 674 **kwargs 675 ): 676 if not _dinov2_import_success: 677 raise RuntimeError( 678 "The vision transformer backend can only be initialized if DINOv2 is installed. " 679 "Please install DINOv2 from https://github.com/facebookresearch/dinov2 " 680 "and then rerun your code." 681 ) 682 683 super().__init__( 684 img_size=img_size, 685 depth=depth, 686 patch_size=patch_size, 687 num_register_tokens=num_register_tokens, 688 **kwargs 689 ) 690 691 self.img_size = img_size 692 self.num_register_tokens = num_register_tokens 693 self.patch_size = patch_size 694 self.attn_outs = [i for i in range(depth) if i % 3 == 2] 695 696 def forward(self, x, masks=None) -> torch.Tensor: 697 698 B = x.shape[0] 699 700 x = self.prepare_tokens_with_masks(x) 701 702 list_of_encoder = [] 703 for i, blk in enumerate(self.blocks): 704 x = blk(x) 705 if i in self.attn_outs: 706 list_of_encoder.append(x) 707 708 x = self.norm(x) 709 x = x[:, self.num_register_tokens + 1:].reshape( 710 B, self.img_size // self.patch_size, self.img_size // self.patch_size, -1 711 ).permute(0, 3, 1, 2).contiguous() 712 713 list_of_encoder = [ 714 o[:, self.num_register_tokens + 1:].reshape( 715 B, self.img_size // self.patch_size, self.img_size // self.patch_size, -1 716 ).permute(0, 3, 1, 2).contiguous() for o in list_of_encoder 717 ] 718 719 return x, list_of_encoder[:3] 720 721 722class ViT_DINOv3(DinoV3VisionTransformer): 723 """Vision Transformer derived from the DINOv3 Codebase (https://arxiv.org/abs/2508.10104). 724 725 Based on: 726 https://github.com/facebookresearch/dinov3/blob/main/dinov3/models/vision_transformer.py. 727 728 Args: 729 img_size: The input image size. 730 patch_size: The patch size. 731 embed_dim: The embedding dimension. 732 depth: The depth of the network. 733 num_heads: The number of heads. 734 ffn_ratio: The FFN rato. 735 n_storage_tokens: The number of storage (class) tokens to remove. 736 kwargs: Keyword arguments for the image encoder base class. 737 """ 738 def __init__( 739 self, 740 in_chans: int = 3, 741 img_size: int = 224, 742 patch_size: int = 16, 743 embed_dim: int = 768, 744 depth: int = 12, 745 num_heads: int = 12, 746 ffn_ratio: float = 4.0, 747 n_storage_tokens: int = 0, 748 **kwargs 749 ): 750 if not _dinov3_import_success: 751 raise RuntimeError( 752 "The vision transformer backend can only be initialized if DINOv3 is installed. " 753 "Please install DINOv3 from https://github.com/facebookresearch/dinov3 " 754 "and then rerun your code." 755 ) 756 757 super().__init__( 758 in_chans=in_chans, 759 img_size=img_size, 760 patch_size=patch_size, 761 embed_dim=embed_dim, 762 depth=depth, 763 num_heads=num_heads, 764 ffn_ratio=ffn_ratio, 765 n_storage_tokens=n_storage_tokens, 766 **kwargs 767 ) 768 769 self.in_chans = in_chans 770 self.img_size = img_size 771 self.n_storage_tokens = n_storage_tokens 772 self.attn_outs = [i for i in range(depth) if i % 3 == 2] 773 774 def forward(self, x) -> torch.Tensor: 775 776 B = x.shape[0] 777 778 x, hw_tuple = self.prepare_tokens_with_masks(x) 779 780 list_of_encoder = [] 781 for i, blk in enumerate(self.blocks): 782 rope_sincos = self.rope_embed(H=hw_tuple[0], W=hw_tuple[1]) 783 x = blk(x, rope_sincos) 784 if i in self.attn_outs: 785 list_of_encoder.append(x) 786 787 x = self.norm(x) 788 x = x[:, self.n_storage_tokens + 1:].reshape( 789 B, self.img_size // self.patch_size, self.img_size // self.patch_size, -1 790 ).permute(0, 3, 1, 2).contiguous() 791 792 list_of_encoder = [ 793 o[:, self.n_storage_tokens + 1:].reshape( 794 B, self.img_size // self.patch_size, self.img_size // self.patch_size, -1 795 ).permute(0, 3, 1, 2).contiguous() for o in list_of_encoder 796 ] 797 798 return x, list_of_encoder[:3] 799 800 801class ViT_Torchvision(nn.Module): 802 """Vision Transformer from torchvision (https://arxiv.org/abs/2010.11929). 803 804 Wraps torchvision ViT models for use as UNETR encoders. Intermediate patch-token grids 805 are collected at quarter-depth intervals and returned as spatial feature maps alongside 806 the final encoder output, matching the `(x, list_from_encoder)` contract of other ViT classes. 807 808 Supported models: vit_b_16, vit_b_32, vit_l_16, vit_l_32, vit_h_14. 809 810 All five models can be used with UNETR. vit_b_16 and vit_l_16 (patch_size=16) work with 811 any skip-connection setting. vit_b_32, vit_l_32, and vit_h_14 require use_skip_connection=False 812 in UNETR; the decoder's internal cropping handles the spatial size difference, and postprocess_masks 813 resizes the output back to the input resolution. 814 815 Args: 816 model_name: Torchvision ViT model name (e.g. 'vit_b_16'). 817 img_size: Expected input image size used by UNETR preprocessing. 818 in_chans: Number of input channels. If != 3, a 1x1 conv projects to 3 channels. 819 pretrained: Whether to load ImageNet-pretrained weights. 820 """ 821 def __init__( 822 self, 823 model_name: str, 824 img_size: int = 224, 825 in_chans: int = 3, 826 pretrained: bool = True, 827 ): 828 super().__init__() 829 if not _torchvision_import_success: 830 raise RuntimeError( 831 "The vision transformer backend can only be initialized if torchvision is installed. " 832 "Please install torchvision from https://github.com/pytorch/vision and then rerun your code." 833 ) 834 835 fn = getattr(_tv_models, model_name) 836 backbone = fn(weights="DEFAULT" if pretrained else None) 837 838 self.conv_proj = backbone.conv_proj 839 self.class_token = backbone.class_token 840 self.encoder = backbone.encoder # pos_embedding, dropout, layers, ln 841 842 self.img_size = img_size 843 self.in_chans = in_chans 844 self.embed_dim = backbone.hidden_dim 845 846 depth = len(backbone.encoder.layers) 847 _c = depth // 4 848 self.chunks_for_projection = [_c - 1, 2 * _c - 1, 3 * _c - 1] 849 850 self.input_proj = nn.Conv2d(in_chans, 3, kernel_size=1) if in_chans != 3 else None 851 852 def _load_from_state_dict( 853 self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs, 854 ): 855 pos_embed = state_dict.get(prefix + "encoder.pos_embedding") 856 current = self.encoder.pos_embedding 857 if ( 858 pos_embed is not None and pos_embed.ndim == 3 859 and pos_embed.shape[0] == current.shape[0] and pos_embed.shape[2] == current.shape[2] 860 and pos_embed.shape[1] != current.shape[1] 861 ): 862 self.encoder.pos_embedding = nn.Parameter( 863 current.new_empty(pos_embed.shape), requires_grad=current.requires_grad, 864 ) 865 super()._load_from_state_dict( 866 state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs, 867 ) 868 869 def _interpolate_pos_embed(self, pos_embed: torch.Tensor, H_p: int, W_p: int) -> torch.Tensor: 870 cls_pos, patch_pos = pos_embed[:, :1], pos_embed[:, 1:] 871 N = patch_pos.shape[1] 872 H_t = W_t = int(N ** 0.5) 873 patch_pos = patch_pos.reshape(1, H_t, W_t, -1).permute(0, 3, 1, 2) 874 patch_pos = F.interpolate(patch_pos, size=(H_p, W_p), mode="bicubic", align_corners=False) 875 patch_pos = patch_pos.permute(0, 2, 3, 1).reshape(1, H_p * W_p, -1) 876 return torch.cat([cls_pos, patch_pos], dim=1) 877 878 def forward(self, x: torch.Tensor) -> torch.Tensor: 879 """Apply the vision transformer to input data. 880 881 Args: 882 x: The input data. 883 884 Returns: 885 The vision transformer output. 886 """ 887 if self.input_proj is not None: 888 x = self.input_proj(x) 889 890 x = self.conv_proj(x) # (B, D, H_p, W_p) 891 B, D, H_p, W_p = x.shape 892 x = x.reshape(B, D, H_p * W_p).permute(0, 2, 1) # (B, N, D) 893 894 cls = self.class_token.expand(B, -1, -1) 895 x = torch.cat([cls, x], dim=1) # (B, 1+N, D) 896 897 pos = self.encoder.pos_embedding 898 if pos.shape[1] != x.shape[1]: 899 pos = self._interpolate_pos_embed(pos, H_p, W_p) 900 x = x + pos 901 x = self.encoder.dropout(x) 902 903 list_from_encoder = [] 904 for i, blk in enumerate(self.encoder.layers): 905 x = blk(x) 906 if i in self.chunks_for_projection: 907 feat = x[:, 1:].reshape(B, H_p, W_p, D).permute(0, 3, 1, 2).contiguous() 908 list_from_encoder.append(feat) 909 910 x = self.encoder.ln(x) 911 x = x[:, 1:].reshape(B, H_p, W_p, D).permute(0, 3, 1, 2).contiguous() 912 return x, list_from_encoder 913 914 915def get_vision_transformer(backbone: str, model: str, img_size: int = 1024, **kwargs) -> nn.Module: 916 """Get vision transformer encoder. 917 918 Args: 919 backbone: The name of the vision transformer implementation. 920 One of "sam" / "cellpose_sam" / "sam2" / "sam3" / "mae" / "scalemae" / "dinov2" / "dinov3" / "torchvision". 921 model: The name of the model. One of "vit_b", "vit_l" or "vit_h". 922 img_size: The size of the input for the image encoder. Input images will be resized to match this size. 923 kwargs: Additional kwargs which can be expected by the vision transformer, 924 e.g. 'base_resolution' for `ViT_ScaleMAE`. 925 926 Returns: 927 The vision transformer. 928 """ 929 if backbone == "sam": 930 if model == "vit_b": 931 encoder = ViT_Sam( 932 depth=12, embed_dim=768, img_size=img_size, mlp_ratio=4, 933 norm_layer=partial(torch.nn.LayerNorm, eps=1e-6), 934 num_heads=12, patch_size=16, qkv_bias=True, use_rel_pos=True, 935 global_attn_indexes=[2, 5, 8, 11], 936 window_size=14, out_chans=256, 937 ) 938 elif model == "vit_l": 939 encoder = ViT_Sam( 940 depth=24, embed_dim=1024, img_size=img_size, mlp_ratio=4, 941 norm_layer=partial(torch.nn.LayerNorm, eps=1e-6), 942 num_heads=16, patch_size=16, qkv_bias=True, use_rel_pos=True, 943 global_attn_indexes=[5, 11, 17, 23], 944 window_size=14, out_chans=256, 945 ) 946 elif model == "vit_h": 947 encoder = ViT_Sam( 948 depth=32, embed_dim=1280, img_size=img_size, mlp_ratio=4, 949 norm_layer=partial(torch.nn.LayerNorm, eps=1e-6), 950 num_heads=16, patch_size=16, qkv_bias=True, use_rel_pos=True, 951 global_attn_indexes=[7, 15, 23, 31], 952 window_size=14, out_chans=256, 953 ) 954 else: 955 raise ValueError(f"'{model}' is not supported by SAM. Currently, 'vit_b', 'vit_l', 'vit_h' are supported.") 956 957 elif backbone == "cellpose_sam": 958 if model != "vit_l": 959 raise ValueError(f"'{model}' is not supported by CellposeSAM. Only 'vit_l' is supported.") 960 encoder = ViT_CellposeSAM(ps=8, bsize=img_size) 961 962 elif backbone == "sam2": 963 if model == "hvit_t": 964 encoder = ViT_Sam2( 965 img_size=img_size, embed_dim=96, num_heads=1, stages=[1, 2, 7, 2], global_att_blocks=[5, 7, 9], 966 window_pos_embed_bkg_spatial_size=[7, 7], backbone_channel_list=[768, 384, 192, 96], 967 ) 968 elif model == "hvit_s": 969 encoder = ViT_Sam2( 970 img_size=img_size, embed_dim=96, num_heads=1, stages=[1, 2, 11, 2], global_att_blocks=[7, 10, 13], 971 window_pos_embed_bkg_spatial_size=[7, 7], backbone_channel_list=[768, 384, 192, 96], 972 ) 973 elif model == "hvit_b": 974 encoder = ViT_Sam2( 975 img_size=img_size, embed_dim=112, num_heads=2, backbone_channel_list=[896, 448, 224, 112], 976 ) 977 elif model == "hvit_l": 978 encoder = ViT_Sam2( 979 img_size=img_size, embed_dim=144, num_heads=2, stages=[2, 6, 36, 4], global_att_blocks=[23, 33, 43], 980 window_spec=[8, 4, 16, 8], backbone_channel_list=[1152, 576, 288, 144], 981 ) 982 else: 983 raise ValueError( 984 f"'{model}' is not supported by SAM2. Currently, 'hvit_t', 'hvit_s', 'hvit_b', 'hvit_l' are supported." 985 ) 986 987 elif backbone == "sam3": 988 if model != "vit_pe": 989 raise ValueError( 990 "'sam3' does not have multiple model configurations. Please use 'vit_pe' as the model configuration." 991 ) 992 993 encoder = ViT_Sam3( 994 img_size=1008, pretrain_img_size=336, patch_size=14, embed_dim=1024, depth=32, num_heads=16, 995 mlp_ratio=4.625, norm_layer="LayerNorm", drop_path_rate=0.1, qkv_bias=True, use_abs_pos=True, 996 tile_abs_pos=True, global_att_blocks=(7, 15, 23, 31), rel_pos_blocks=(), use_rope=True, 997 use_interp_rope=True, window_size=24, pretrain_use_cls_token=True, retain_cls_token=False, ln_pre=True, 998 ln_post=False, return_interm_layers=False, bias_patch_embed=False, compile_mode=None, 999 ) 1000 1001 elif backbone == "mae": 1002 if model == "vit_b": 1003 encoder = ViT_MAE( 1004 img_size=img_size, patch_size=16, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4, 1005 qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6) 1006 ) 1007 elif model == "vit_l": 1008 encoder = ViT_MAE( 1009 img_size=img_size, patch_size=16, embed_dim=1024, depth=24, num_heads=16, mlp_ratio=4, 1010 qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6) 1011 ) 1012 elif model == "vit_h": 1013 encoder = ViT_MAE( 1014 img_size=img_size, patch_size=14, embed_dim=1280, depth=32, num_heads=16, mlp_ratio=4, 1015 qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6) 1016 ) 1017 else: 1018 raise ValueError(f"'{model}' is not supported by MAE. Currently, 'vit_b', 'vit_l', 'vit_h' are supported.") 1019 1020 elif backbone == "scalemae": 1021 base_resolution = kwargs.get("base_resolution", 2.5) 1022 1023 if model == "vit_b": 1024 encoder = ViT_ScaleMAE( 1025 img_size=img_size, patch_size=8, embed_dim=768, depth=12, num_heads=12, 1026 mlp_ratio=4, qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), 1027 base_resolution=base_resolution, 1028 ) 1029 elif model == "vit_l": 1030 encoder = ViT_ScaleMAE( 1031 img_size=img_size, patch_size=8, embed_dim=1024, depth=24, num_heads=16, 1032 mlp_ratio=4, qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), 1033 base_resolution=base_resolution, 1034 ) 1035 elif model == "vit_h": 1036 encoder = ViT_ScaleMAE( 1037 img_size=img_size, patch_size=8, embed_dim=1280, depth=32, num_heads=16, 1038 mlp_ratio=4, qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), 1039 base_resolution=base_resolution, 1040 ) 1041 else: 1042 raise ValueError( 1043 f"'{model}' is not supported by ScaleMAE. Currently, 'vit_b', 'vit_l' and 'vit_h' are supported." 1044 ) 1045 1046 elif backbone == "dinov2": 1047 block_fn = partial(Block, attn_class=MemEffAttention) 1048 msg = "The model name should be either 'vit_<X>' or 'vit_<X>_reg<Y>." 1049 1050 if model.startswith("vit_s"): 1051 assert model in ["vit_s", "vit_s_reg4"], msg 1052 encoder = ViT_DINOv2( 1053 img_size=img_size, patch_size=14, embed_dim=384, depth=12, num_heads=6, mlp_ratio=4, 1054 block_fn=block_fn, in_chans=3, channel_adaptive=False, init_values=1e-5, block_chunks=0, 1055 num_register_tokens=4 if model.endswith("_reg4") else 0, 1056 ) 1057 elif model.startswith("vit_b"): 1058 assert model in ["vit_b", "vit_b_reg4"], msg 1059 encoder = ViT_DINOv2( 1060 img_size=img_size, patch_size=14, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4, 1061 block_fn=block_fn, in_chans=3, channel_adaptive=False, init_values=1e-5, block_chunks=0, 1062 num_register_tokens=4 if model.endswith("_reg4") else 0, 1063 ) 1064 elif model.startswith("vit_l"): 1065 assert model in ["vit_l", "vit_l_reg4"], msg 1066 encoder = ViT_DINOv2( 1067 img_size=img_size, patch_size=14, embed_dim=1024, depth=24, num_heads=16, mlp_ratio=4, 1068 block_fn=block_fn, in_chans=3, channel_adaptive=False, init_values=1e-5, block_chunks=0, 1069 num_register_tokens=4 if model.endswith("_reg4") else 0, 1070 ) 1071 elif model.startswith("vit_g"): 1072 assert model in ["vit_g", "vit_g_reg4"], msg 1073 encoder = ViT_DINOv2( 1074 img_size=img_size, patch_size=14, embed_dim=1536, depth=40, num_heads=24, mlp_ratio=4, 1075 block_fn=block_fn, in_chans=3, channel_adaptive=False, init_values=1e-5, block_chunks=0, 1076 num_register_tokens=4 if model.endswith("_reg4") else 0, ffn_layer="swiglu", 1077 ) 1078 else: 1079 raise ValueError( 1080 f"'{model}' is not supported by DINOv2. Currently, 'vit_s', 'vit_b', 'vit_l' and 'vit_g' are supported." 1081 ) 1082 1083 elif backbone == "dinov3": 1084 1085 if model == "vit_s": 1086 encoder = ViT_DINOv3( 1087 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", embed_dim=384, 1088 num_heads=6, layerscale_init=1.0e-05, norm_layer="layernormbf16", n_storage_tokens=4, mask_k_bias=True, 1089 ) 1090 elif model == "vit_s+": 1091 encoder = ViT_DINOv3( 1092 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", embed_dim=384, 1093 num_heads=6, ffn_ratio=6, layerscale_init=1.0e-05, norm_layer="layernormbf16", 1094 ffn_layer="swiglu", n_storage_tokens=4, mask_k_bias=True, 1095 ) 1096 1097 elif model == "vit_b": 1098 encoder = ViT_DINOv3( 1099 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", 1100 layerscale_init=1.0e-05, norm_layer="layernormbf16", n_storage_tokens=4, mask_k_bias=True, 1101 ) 1102 elif model == "vit_l": 1103 encoder = ViT_DINOv3( 1104 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", embed_dim=1024, 1105 depth=24, num_heads=16, layerscale_init=1.0e-05, norm_layer="layernormbf16", 1106 n_storage_tokens=4, mask_k_bias=True, 1107 ) 1108 elif model == "vit_l+": 1109 encoder = ViT_DINOv3( 1110 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", embed_dim=1024, 1111 depth=24, num_heads=16, ffn_ratio=6.0, layerscale_init=1.0e-05, norm_layer="layernormbf16", 1112 ffn_layer="swiglu", n_storage_tokens=4, mask_k_bias=True, 1113 ) 1114 elif model == "vit_h+": 1115 encoder = ViT_DINOv3( 1116 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", embed_dim=1280, 1117 depth=32, num_heads=20, ffn_ratio=6.0, layerscale_init=1.0e-05, norm_layer="layernormbf16", 1118 ffn_layer="swiglu", n_storage_tokens=4, mask_k_bias=True, 1119 ) 1120 elif model == "vit_7b": 1121 encoder = ViT_DINOv3( 1122 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", embed_dim=4096, 1123 depth=40, num_heads=32, ffn_ratio=3, qkv_bias=False, drop_path_rate=0.0, layerscale_init=1.0e-05, 1124 norm_layer="layernormbf16", ffn_layer="swiglu64", n_storage_tokens=4, mask_k_bias=True, 1125 untie_global_and_local_cls_norm=True, 1126 ) 1127 else: 1128 raise ValueError( 1129 f"'{model}' is not supported by DINOv3. Currently, " 1130 " 'vit_s', 'vit_s+', 'vit_b', 'vit_l', 'vit_l+', 'vit_h+'. 'vit_7b' are supported." 1131 ) 1132 1133 elif backbone == "torchvision": 1134 supported = ["vit_b_16", "vit_b_32", "vit_l_16", "vit_l_32", "vit_h_14"] 1135 if model not in supported: 1136 raise ValueError(f"'{model}' is not supported by the torchvision backbone. Choose from: {supported}.") 1137 encoder = ViT_Torchvision(model_name=model, img_size=img_size, **kwargs) 1138 1139 else: 1140 raise ValueError( 1141 "The 'UNETR' supported backbones are 'sam', 'cellpose_sam', 'sam2', 'sam3', " 1142 "'mae', 'scalemae', 'dinov2', 'dinov3' or 'torchvision'. Please choose one of them." 1143 ) 1144 1145 return encoder
67class ViT_Sam(ImageEncoderViT): 68 """Vision Transformer derived from the Segment Anything Codebase (https://arxiv.org/abs/2304.02643). 69 70 Based on: 71 https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/image_encoder.py 72 73 Args: 74 in_chans: The number of input channels. 75 embed_dim: The embedding dimension, corresponding to the number of output channels of the vision transformer. 76 global_attn_indexes: The global attention indices. 77 apply_neck: Whether to apply the convolutional bottleneck after outputs of the last attention head. 78 kwargs: Keyword arguments for the image encoder base class. 79 """ 80 def __init__( 81 self, 82 in_chans: int = 3, 83 embed_dim: int = 768, 84 global_attn_indexes: Tuple[int, ...] = [2, 5, 8, 11], 85 apply_neck: bool = False, 86 **kwargs, 87 ) -> None: 88 if not _sam_import_success: 89 raise RuntimeError( 90 "The vision transformer backend can only be initialized if segment anything is installed. " 91 "Please install segment anything from https://github.com/facebookresearch/segment-anything " 92 "and then rerun your code." 93 ) 94 95 super().__init__(embed_dim=embed_dim, global_attn_indexes=global_attn_indexes, **kwargs) 96 self.chunks_for_projection = global_attn_indexes 97 self.in_chans = in_chans 98 self.embed_dim = embed_dim 99 self.apply_neck = apply_neck 100 101 def forward(self, x: torch.Tensor) -> torch.Tensor: 102 """Apply the vision transformer to input data. 103 104 Args: 105 x: The input data. 106 107 Returns: 108 The vision transformer output. 109 """ 110 x = self.patch_embed(x) 111 if self.pos_embed is not None: 112 x = x + self.pos_embed 113 114 list_from_encoder = [] 115 for i, blk in enumerate(self.blocks): 116 x = blk(x) 117 if i in self.chunks_for_projection: 118 list_from_encoder.append(x) 119 120 x = x.permute(0, 3, 1, 2) 121 122 if self.apply_neck: 123 x = self.neck(x) 124 125 list_from_encoder = [e.permute(0, 3, 1, 2) for e in list_from_encoder] 126 return x, list_from_encoder[:3]
Vision Transformer derived from the Segment Anything Codebase (https://arxiv.org/abs/2304.02643).
Arguments:
- in_chans: The number of input channels.
- embed_dim: The embedding dimension, corresponding to the number of output channels of the vision transformer.
- global_attn_indexes: The global attention indices.
- apply_neck: Whether to apply the convolutional bottleneck after outputs of the last attention head.
- kwargs: Keyword arguments for the image encoder base class.
80 def __init__( 81 self, 82 in_chans: int = 3, 83 embed_dim: int = 768, 84 global_attn_indexes: Tuple[int, ...] = [2, 5, 8, 11], 85 apply_neck: bool = False, 86 **kwargs, 87 ) -> None: 88 if not _sam_import_success: 89 raise RuntimeError( 90 "The vision transformer backend can only be initialized if segment anything is installed. " 91 "Please install segment anything from https://github.com/facebookresearch/segment-anything " 92 "and then rerun your code." 93 ) 94 95 super().__init__(embed_dim=embed_dim, global_attn_indexes=global_attn_indexes, **kwargs) 96 self.chunks_for_projection = global_attn_indexes 97 self.in_chans = in_chans 98 self.embed_dim = embed_dim 99 self.apply_neck = apply_neck
101 def forward(self, x: torch.Tensor) -> torch.Tensor: 102 """Apply the vision transformer to input data. 103 104 Args: 105 x: The input data. 106 107 Returns: 108 The vision transformer output. 109 """ 110 x = self.patch_embed(x) 111 if self.pos_embed is not None: 112 x = x + self.pos_embed 113 114 list_from_encoder = [] 115 for i, blk in enumerate(self.blocks): 116 x = blk(x) 117 if i in self.chunks_for_projection: 118 list_from_encoder.append(x) 119 120 x = x.permute(0, 3, 1, 2) 121 122 if self.apply_neck: 123 x = self.neck(x) 124 125 list_from_encoder = [e.permute(0, 3, 1, 2) for e in list_from_encoder] 126 return x, list_from_encoder[:3]
Apply the vision transformer to input data.
Arguments:
- x: The input data.
Returns:
The vision transformer output.
129class ViT_CellposeSAM(nn.Module): 130 """Vision Transformer derived from the CellposeSAM Codebase (https://doi.org/10.1038/s41592-025-02595-x). 131 132 This replicates CellposeSAM's actual initialization: instantiate SAM's ``ImageEncoderViT`` via 133 ``sam_model_registry``, then modify the patch embedding, position embeddings, and set global attention. 134 This preserves SAM's original relative position bias sizes, enabling direct checkpoint loading 135 without any interpolation. 136 137 Based on: https://github.com/MouseLand/cellpose/blob/main/cellpose/vit_sam.py 138 139 NOTE: The pretrained CellposeSAM model uses ``vit_l`` exclusively. 140 141 Args: 142 ps: The patch size (default for CellposeSAM is 8). 143 bsize: The input image size (default for CellposeSAM is 256). 144 apply_neck: Whether to apply the convolutional bottleneck after outputs of the last attention head. 145 """ 146 def __init__(self, ps: int = 8, bsize: int = 256, apply_neck: bool = True) -> None: 147 super().__init__() 148 149 if not _sam_import_success: 150 raise RuntimeError( 151 "The vision transformer backend can only be initialized if segment anything is installed. " 152 "Please install segment anything from https://github.com/facebookresearch/segment-anything " 153 "and then rerun your code." 154 ) 155 156 from segment_anything import sam_model_registry 157 158 # Creates the SAM vit_l encoder and applies CellposeSAM's modifications (same as cellpose.vit_sam.Transformer). 159 encoder = sam_model_registry["vit_l"](None).image_encoder 160 161 w = encoder.patch_embed.proj.weight.detach() 162 nchan = w.shape[0] 163 164 # CellPoseSAM changes the patch size from 16 to 'ps'. 165 encoder.patch_embed.proj = nn.Conv2d(3, nchan, stride=ps, kernel_size=ps) 166 encoder.patch_embed.proj.weight.data = w[:, :, ::16 // ps, ::16 // ps] 167 168 # Next, they subsample position embeddings for the new patch size and input resolution. 169 ds = (1024 // 16) // (bsize // ps) 170 encoder.pos_embed = nn.Parameter(encoder.pos_embed[:, ::ds, ::ds], requires_grad=True) 171 172 # Finally, they set all blocks to global attention. 173 for blk in encoder.blocks: 174 blk.window_size = 0 175 176 # Store encoder submodules directly ('state_dict' keys match CellposeSAM after prefix stripping). 177 self.patch_embed = encoder.patch_embed 178 self.pos_embed = encoder.pos_embed 179 self.blocks = encoder.blocks 180 self.neck = encoder.neck 181 182 # Additional attributes expected by UNETR. 183 self.embed_dim = nchan 184 self.img_size = bsize 185 self.in_chans = 3 186 self.apply_neck = apply_neck 187 188 # Feature extraction at evenly-spaced depths. 189 depth = len(self.blocks) 190 _chunks = depth // 4 191 self.chunks_for_projection = [_chunks - 1, 2 * _chunks - 1, 3 * _chunks - 1, 4 * _chunks - 1] 192 193 def forward(self, x: torch.Tensor) -> torch.Tensor: 194 """Apply the vision transformer to input data. 195 196 Args: 197 x: The input data. 198 199 Returns: 200 The vision transformer output. 201 """ 202 x = self.patch_embed(x) 203 if self.pos_embed is not None: 204 x = x + self.pos_embed 205 206 list_from_encoder = [] 207 for i, blk in enumerate(self.blocks): 208 x = blk(x) 209 if i in self.chunks_for_projection: 210 list_from_encoder.append(x) 211 212 x = x.permute(0, 3, 1, 2) 213 214 if self.apply_neck: 215 x = self.neck(x) 216 217 list_from_encoder = [e.permute(0, 3, 1, 2) for e in list_from_encoder] 218 return x, list_from_encoder[:3]
Vision Transformer derived from the CellposeSAM Codebase (https://doi.org/10.1038/s41592-025-02595-x).
This replicates CellposeSAM's actual initialization: instantiate SAM's ImageEncoderViT via
sam_model_registry, then modify the patch embedding, position embeddings, and set global attention.
This preserves SAM's original relative position bias sizes, enabling direct checkpoint loading
without any interpolation.
Based on: https://github.com/MouseLand/cellpose/blob/main/cellpose/vit_sam.py
NOTE: The pretrained CellposeSAM model uses vit_l exclusively.
Arguments:
- ps: The patch size (default for CellposeSAM is 8).
- bsize: The input image size (default for CellposeSAM is 256).
- apply_neck: Whether to apply the convolutional bottleneck after outputs of the last attention head.
146 def __init__(self, ps: int = 8, bsize: int = 256, apply_neck: bool = True) -> None: 147 super().__init__() 148 149 if not _sam_import_success: 150 raise RuntimeError( 151 "The vision transformer backend can only be initialized if segment anything is installed. " 152 "Please install segment anything from https://github.com/facebookresearch/segment-anything " 153 "and then rerun your code." 154 ) 155 156 from segment_anything import sam_model_registry 157 158 # Creates the SAM vit_l encoder and applies CellposeSAM's modifications (same as cellpose.vit_sam.Transformer). 159 encoder = sam_model_registry["vit_l"](None).image_encoder 160 161 w = encoder.patch_embed.proj.weight.detach() 162 nchan = w.shape[0] 163 164 # CellPoseSAM changes the patch size from 16 to 'ps'. 165 encoder.patch_embed.proj = nn.Conv2d(3, nchan, stride=ps, kernel_size=ps) 166 encoder.patch_embed.proj.weight.data = w[:, :, ::16 // ps, ::16 // ps] 167 168 # Next, they subsample position embeddings for the new patch size and input resolution. 169 ds = (1024 // 16) // (bsize // ps) 170 encoder.pos_embed = nn.Parameter(encoder.pos_embed[:, ::ds, ::ds], requires_grad=True) 171 172 # Finally, they set all blocks to global attention. 173 for blk in encoder.blocks: 174 blk.window_size = 0 175 176 # Store encoder submodules directly ('state_dict' keys match CellposeSAM after prefix stripping). 177 self.patch_embed = encoder.patch_embed 178 self.pos_embed = encoder.pos_embed 179 self.blocks = encoder.blocks 180 self.neck = encoder.neck 181 182 # Additional attributes expected by UNETR. 183 self.embed_dim = nchan 184 self.img_size = bsize 185 self.in_chans = 3 186 self.apply_neck = apply_neck 187 188 # Feature extraction at evenly-spaced depths. 189 depth = len(self.blocks) 190 _chunks = depth // 4 191 self.chunks_for_projection = [_chunks - 1, 2 * _chunks - 1, 3 * _chunks - 1, 4 * _chunks - 1]
Initialize internal Module state, shared by both nn.Module and ScriptModule.
193 def forward(self, x: torch.Tensor) -> torch.Tensor: 194 """Apply the vision transformer to input data. 195 196 Args: 197 x: The input data. 198 199 Returns: 200 The vision transformer output. 201 """ 202 x = self.patch_embed(x) 203 if self.pos_embed is not None: 204 x = x + self.pos_embed 205 206 list_from_encoder = [] 207 for i, blk in enumerate(self.blocks): 208 x = blk(x) 209 if i in self.chunks_for_projection: 210 list_from_encoder.append(x) 211 212 x = x.permute(0, 3, 1, 2) 213 214 if self.apply_neck: 215 x = self.neck(x) 216 217 list_from_encoder = [e.permute(0, 3, 1, 2) for e in list_from_encoder] 218 return x, list_from_encoder[:3]
Apply the vision transformer to input data.
Arguments:
- x: The input data.
Returns:
The vision transformer output.
221class ViT_MAE(VisionTransformer): 222 """Vision Transformer derived from the Masked Auto Encoder Codebase (https://arxiv.org/abs/2111.06377). 223 224 Based on: 225 https://github.com/facebookresearch/mae/blob/main/models_vit.py#L20-L53 226 227 Args: 228 img_size: The size of the input for the image encoder. Input images will be resized to match this size. 229 in_chans: The number of input channels. 230 depth: The depth of the vision transformer. 231 kwargs: Additional keyword arguments for the vision transformer base class. 232 """ 233 def __init__( 234 self, 235 img_size: int = 1024, # chosen to match our experiments with segment anything 236 in_chans: int = 3, 237 depth: int = 12, 238 **kwargs 239 ): 240 if not _timm_import_success: 241 raise RuntimeError( 242 "The vision transformer backend can only be initialized if timm is installed. " 243 "Please install timm (using conda/mamba) for using https://github.com/facebookresearch/mae/ " 244 "and then rerun your code" 245 ) 246 super().__init__(img_size=img_size, depth=depth, **kwargs) 247 self.img_size = img_size 248 self.in_chans = in_chans 249 self.depth = depth 250 251 def convert_to_expected_dim(self, inputs_): 252 """@private 253 """ 254 inputs_ = inputs_[:, 1:, :] # removing the class tokens 255 # reshape the outputs to desired shape (N x H*W X C -> N x H x W x C) 256 rdim = inputs_.shape[1] 257 dshape = int(rdim ** 0.5) # finding the square root of the outputs for obtaining the patch shape 258 inputs_ = torch.unflatten(inputs_, 1, (dshape, dshape)) 259 inputs_ = inputs_.permute(0, 3, 1, 2) 260 return inputs_ 261 262 def forward_features(self, x): 263 """@private 264 """ 265 B = x.shape[0] 266 x = self.patch_embed(x) 267 268 cls_tokens = self.cls_token.expand(B, -1, -1) 269 x = torch.cat((cls_tokens, x), dim=1) 270 271 x = x + self.pos_embed 272 x = self.pos_drop(x) 273 274 # chunks obtained for getting the projections for conjuctions with upsampling blocks 275 _chunks = int(self.depth / 4) 276 chunks_for_projection = [_chunks - 1, 2*_chunks - 1, 3*_chunks - 1, 4*_chunks - 1] 277 278 list_from_encoder = [] 279 for i, blk in enumerate(self.blocks): 280 x = blk(x) 281 if i in chunks_for_projection: 282 list_from_encoder.append(self.convert_to_expected_dim(x)) 283 284 x = self.convert_to_expected_dim(x) 285 return x, list_from_encoder[:3] 286 287 def forward(self, x: torch.Tensor) -> torch.Tensor: 288 """Apply the vision transformer to input data. 289 290 Args: 291 x: The input data. 292 293 Returns: 294 The vision transformer output. 295 """ 296 x, list_from_encoder = self.forward_features(x) 297 return x, list_from_encoder
Vision Transformer derived from the Masked Auto Encoder Codebase (https://arxiv.org/abs/2111.06377).
Based on: https://github.com/facebookresearch/mae/blob/main/models_vit.py#L20-L53
Arguments:
- img_size: The size of the input for the image encoder. Input images will be resized to match this size.
- in_chans: The number of input channels.
- depth: The depth of the vision transformer.
- kwargs: Additional keyword arguments for the vision transformer base class.
233 def __init__( 234 self, 235 img_size: int = 1024, # chosen to match our experiments with segment anything 236 in_chans: int = 3, 237 depth: int = 12, 238 **kwargs 239 ): 240 if not _timm_import_success: 241 raise RuntimeError( 242 "The vision transformer backend can only be initialized if timm is installed. " 243 "Please install timm (using conda/mamba) for using https://github.com/facebookresearch/mae/ " 244 "and then rerun your code" 245 ) 246 super().__init__(img_size=img_size, depth=depth, **kwargs) 247 self.img_size = img_size 248 self.in_chans = in_chans 249 self.depth = depth
287 def forward(self, x: torch.Tensor) -> torch.Tensor: 288 """Apply the vision transformer to input data. 289 290 Args: 291 x: The input data. 292 293 Returns: 294 The vision transformer output. 295 """ 296 x, list_from_encoder = self.forward_features(x) 297 return x, list_from_encoder
Apply the vision transformer to input data.
Arguments:
- x: The input data.
Returns:
The vision transformer output.
300class ViT_Sam2(ImageEncoder): 301 """Vision Transformer derived from the Segment Anything 2 Codebase (https://arxiv.org/abs/2408.00714). 302 303 Based on https://github.com/facebookresearch/sam2/blob/main/sam2/modeling/backbones/image_encoder.py. 304 305 Args: 306 backbone_channel_list: The channels throughout the entire backbone. 307 embed_dim: The initial embedding dimension. 308 num_heads: The initial number of heads. 309 stages: The number of blocks per stage. 310 global_att_blocks: The parameter to decide which blocks have global attention. 311 window_pos_embed_bkg_spatial_size: The spatial size of windowed positional embedding. 312 window_spec: The window size per stage, when not using global attention. 313 scalp: The count of lowest resolution features to discard. 314 """ 315 def __init__( 316 self, 317 backbone_channel_list: List[int], 318 img_size: int = 1024, 319 embed_dim: int = 96, 320 num_heads: int = 1, 321 stages: Tuple[int, ...] = (2, 3, 16, 3), 322 global_att_blocks: Tuple[int, ...] = (12, 16, 20), 323 window_pos_embed_bkg_spatial_size: Tuple[int, int] = (14, 14), 324 window_spec: Tuple[int, ...] = (8, 4, 14, 7), 325 scalp: int = 1, 326 **kwargs 327 ): 328 if not _sam2_import_success: 329 raise RuntimeError( 330 "The vision transformer backend can only be initialized if segment anything 2 is installed. " 331 "Please install segment anything 2 from https://github.com/facebookresearch/sam2 " 332 "and then rerun your code" 333 ) 334 335 trunk = Hiera( 336 embed_dim=embed_dim, 337 num_heads=num_heads, 338 stages=stages, 339 global_att_blocks=global_att_blocks, 340 window_pos_embed_bkg_spatial_size=window_pos_embed_bkg_spatial_size, 341 window_spec=window_spec, 342 ) 343 neck = FpnNeck( 344 position_encoding=PositionEmbeddingSine(num_pos_feats=256), 345 d_model=256, 346 backbone_channel_list=backbone_channel_list, 347 fpn_top_down_levels=[2, 3], 348 fpn_interp_model="nearest", 349 ) 350 351 super().__init__(trunk=trunk, neck=neck, scalp=scalp, **kwargs) 352 self.scalp = scalp 353 self.embed_dim = embed_dim 354 self.img_size = img_size 355 356 def forward(self, x: torch.Tensor): 357 # The forward pass throught the backbone. 358 features, pos = self.neck(self.trunk(x)) 359 if self.scalp > 0: # This discard the "n" lowest resolution features. 360 features, pos = features[:-self.scalp], pos[:-self.scalp] 361 362 return features[-1], features
Vision Transformer derived from the Segment Anything 2 Codebase (https://arxiv.org/abs/2408.00714).
Based on https://github.com/facebookresearch/sam2/blob/main/sam2/modeling/backbones/image_encoder.py.
Arguments:
- backbone_channel_list: The channels throughout the entire backbone.
- embed_dim: The initial embedding dimension.
- num_heads: The initial number of heads.
- stages: The number of blocks per stage.
- global_att_blocks: The parameter to decide which blocks have global attention.
- window_pos_embed_bkg_spatial_size: The spatial size of windowed positional embedding.
- window_spec: The window size per stage, when not using global attention.
- scalp: The count of lowest resolution features to discard.
315 def __init__( 316 self, 317 backbone_channel_list: List[int], 318 img_size: int = 1024, 319 embed_dim: int = 96, 320 num_heads: int = 1, 321 stages: Tuple[int, ...] = (2, 3, 16, 3), 322 global_att_blocks: Tuple[int, ...] = (12, 16, 20), 323 window_pos_embed_bkg_spatial_size: Tuple[int, int] = (14, 14), 324 window_spec: Tuple[int, ...] = (8, 4, 14, 7), 325 scalp: int = 1, 326 **kwargs 327 ): 328 if not _sam2_import_success: 329 raise RuntimeError( 330 "The vision transformer backend can only be initialized if segment anything 2 is installed. " 331 "Please install segment anything 2 from https://github.com/facebookresearch/sam2 " 332 "and then rerun your code" 333 ) 334 335 trunk = Hiera( 336 embed_dim=embed_dim, 337 num_heads=num_heads, 338 stages=stages, 339 global_att_blocks=global_att_blocks, 340 window_pos_embed_bkg_spatial_size=window_pos_embed_bkg_spatial_size, 341 window_spec=window_spec, 342 ) 343 neck = FpnNeck( 344 position_encoding=PositionEmbeddingSine(num_pos_feats=256), 345 d_model=256, 346 backbone_channel_list=backbone_channel_list, 347 fpn_top_down_levels=[2, 3], 348 fpn_interp_model="nearest", 349 ) 350 351 super().__init__(trunk=trunk, neck=neck, scalp=scalp, **kwargs) 352 self.scalp = scalp 353 self.embed_dim = embed_dim 354 self.img_size = img_size
356 def forward(self, x: torch.Tensor): 357 # The forward pass throught the backbone. 358 features, pos = self.neck(self.trunk(x)) 359 if self.scalp > 0: # This discard the "n" lowest resolution features. 360 features, pos = features[:-self.scalp], pos[:-self.scalp] 361 362 return features[-1], features
365class ViT_Sam3(SAM3ViT): 366 """Vision Transformer derived from the Segment Anything 3 Codebase (https://arxiv.org/abs/2511.16719). 367 368 Based on https://github.com/facebookresearch/sam3/blob/main/sam3/model/vitdet.py. 369 370 Args: 371 img_size: The input image size. 372 embed_dim: The embedding dimension, corresponding to the number of output channels of the vision transformer. 373 kwargs: Keyword arguments for the image encoder base class. 374 """ 375 def __init__(self, img_size: int = 1024, embed_dim: int = 768, **kwargs): 376 if not _sam3_import_success: 377 raise RuntimeError( 378 "The vision transformer backend can only be initialized if segment anything 3 is installed. " 379 "Please install segment anything 3 from https://github.com/facebookresearch/sam3 " 380 "and then rerun your code" 381 ) 382 383 super().__init__(img_size=img_size, embed_dim=embed_dim, **kwargs) 384 self.img_size = img_size 385 self.embed_dim = embed_dim 386 387 def forward_features(self, x): 388 """@private 389 """ 390 x = self.patch_embed(x) 391 h, w = x.shape[1], x.shape[2] 392 393 s = 0 394 if self.retain_cls_token: 395 # If the 'cls_token' is retained, we don't maintain the spatial shape. 396 x = torch.cat([self.class_embedding, x.flatten(1, 2)], dim=1) 397 s = 1 398 399 if self.pos_embed is not None: 400 x = x + get_abs_pos( 401 self.pos_embed, self.pretrain_use_cls_token, (h, w), self.retain_cls_token, tiling=self.tile_abs_pos, 402 ) 403 404 x = self.ln_pre(x) 405 406 list_from_encoder = [] 407 for i, blk in enumerate(self.blocks): 408 if self.use_act_checkpoint and self.training: 409 x = torch.utils.checkpoint.checkpoint(blk, x, use_reentrant=False) 410 else: 411 x = blk(x) 412 413 x = self._convert_to_expected_dim(x, i, s) 414 415 if i in self.full_attn_ids: 416 list_from_encoder.append(x) 417 418 return x, list_from_encoder 419 420 def _convert_to_expected_dim(self, x, i, s): 421 if (i == self.full_attn_ids[-1]) or ( 422 self.return_interm_layers and i in self.full_attn_ids 423 ): 424 if i == self.full_attn_ids[-1]: 425 x = self.ln_post(x) 426 427 feats = x[:, s:] 428 if feats.ndim == 4: 429 feats = feats.permute(0, 3, 1, 2) 430 else: 431 assert feats.ndim == 3 432 h = w = math.sqrt(feats.shape[1]) 433 feats = feats.reshape(feats.shape[0], h, w, feats.shape[-1]).permute(0, 3, 1, 2) 434 return feats 435 436 else: 437 return x 438 439 def forward(self, x: torch.Tensor): 440 """Apply the vision transformer to input data. 441 442 Args: 443 x: The input data. 444 445 Returns: 446 The vision transformer output. 447 """ 448 x, list_from_encoder = self.forward_features(x) 449 return x, list_from_encoder
Vision Transformer derived from the Segment Anything 3 Codebase (https://arxiv.org/abs/2511.16719).
Based on https://github.com/facebookresearch/sam3/blob/main/sam3/model/vitdet.py.
Arguments:
- img_size: The input image size.
- embed_dim: The embedding dimension, corresponding to the number of output channels of the vision transformer.
- kwargs: Keyword arguments for the image encoder base class.
375 def __init__(self, img_size: int = 1024, embed_dim: int = 768, **kwargs): 376 if not _sam3_import_success: 377 raise RuntimeError( 378 "The vision transformer backend can only be initialized if segment anything 3 is installed. " 379 "Please install segment anything 3 from https://github.com/facebookresearch/sam3 " 380 "and then rerun your code" 381 ) 382 383 super().__init__(img_size=img_size, embed_dim=embed_dim, **kwargs) 384 self.img_size = img_size 385 self.embed_dim = embed_dim
439 def forward(self, x: torch.Tensor): 440 """Apply the vision transformer to input data. 441 442 Args: 443 x: The input data. 444 445 Returns: 446 The vision transformer output. 447 """ 448 x, list_from_encoder = self.forward_features(x) 449 return x, list_from_encoder
Apply the vision transformer to input data.
Arguments:
- x: The input data.
Returns:
The vision transformer output.
456class CustomCompose: 457 def __init__(self, rescale_transform, other_transforms, src_transform): 458 self.rescale_transform = rescale_transform 459 self.other_transforms = other_transforms 460 self.src_transform = src_transform 461 462 def __call__(self, x, valid_masks=None): 463 if valid_masks is not None: 464 nodata = (x * (1 - valid_masks.float())).max() 465 x_aug = self.rescale_transform(x) 466 parms = self.rescale_transform._params 467 468 # sanity check, comment if this is working 469 # valid_masks = self.rescale_transform(valid_masks.float(), params=parms) 470 # assert (x_aug==self.rescale_transform(x, params=parms)).all() # 471 472 if valid_masks is not None: 473 valid_masks = x_aug != nodata 474 _, c, h, w = x_aug.shape 475 zero_ratio = ((valid_masks == 0).sum((1, 2, 3)) / (h * w * c)).cpu().numpy() 476 else: 477 zero_ratio = -1 478 479 if self.other_transforms: 480 x_aug = self.other_transforms(x_aug) 481 x_src = self.src_transform(x_aug) 482 dx = parms["src"][:, 1, 0] - parms["src"][:, 0, 0] 483 484 # dy = (parms['src'][:,2,1] - parms['src'][:,1,1]) 485 # assert (dx == dy).all() 486 487 h, w = x_aug.shape[-2:] 488 # assert h == w 489 490 return x_aug, x_src, dx / h, zero_ratio, valid_masks
493def get_2d_sincos_pos_embed_with_resolution(embed_dim, grid_size, res, cls_token=False, device="cpu"): 494 """ 495 grid_size: int of the grid height and width 496 res: array of size n, representing the resolution of a pixel (say, in meters), 497 return: 498 pos_embed: [n,grid_size*grid_size, embed_dim] or [n,1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token) 499 """ 500 # res = torch.FloatTensor(res).to(device) 501 res = res.to(device) 502 grid_h = torch.arange(grid_size, dtype=torch.float32, device=device) 503 grid_w = torch.arange(grid_size, dtype=torch.float32, device=device) 504 grid = torch.meshgrid(grid_w, grid_h, indexing="xy") # here h goes first,direction reversed for numpy 505 grid = torch.stack(grid, dim=0) # 2 x h x w 506 507 # grid = grid.reshape([2, 1, grid_size, grid_size]) 508 grid = torch.einsum("chw,n->cnhw", grid, res) # 2 x n x h x w 509 _, n, h, w = grid.shape 510 pos_embed = get_2d_sincos_pos_embed_from_grid_torch(embed_dim, grid) # (nxH*W, D/2) 511 pos_embed = pos_embed.reshape(n, h * w, embed_dim) 512 if cls_token: 513 pos_embed = torch.cat( 514 [torch.zeros([n, 1, embed_dim], dtype=torch.float32, device=pos_embed.device), pos_embed], dim=1 515 ) 516 517 return pos_embed
grid_size: int of the grid height and width res: array of size n, representing the resolution of a pixel (say, in meters), return: pos_embed: [n,grid_sizegrid_size, embed_dim] or [n,1+grid_sizegrid_size, embed_dim] (w/ or w/o cls_token)
520def get_2d_sincos_pos_embed_from_grid_torch(embed_dim, grid): 521 assert embed_dim % 2 == 0 522 523 # use half of dimensions to encode grid_h 524 emb_h = get_1d_sincos_pos_embed_from_grid_torch(embed_dim // 2, grid[0]) # (H*W, D/2) 525 emb_w = get_1d_sincos_pos_embed_from_grid_torch(embed_dim // 2, grid[1]) # (H*W, D/2) 526 527 emb = torch.cat([emb_h, emb_w], dim=1) # (H*W, D) 528 return emb
531def get_1d_sincos_pos_embed_from_grid_torch(embed_dim, pos): 532 """ 533 embed_dim: output dimension for each position 534 pos: a list of positions to be encoded: size (M,) 535 out: (M, D) 536 """ 537 assert embed_dim % 2 == 0 538 # old_shape = pos 539 omega = torch.arange(embed_dim // 2, dtype=torch.float32, device=pos.device) 540 omega /= embed_dim / 2.0 541 omega = 1.0 / 10000**omega # (D/2,) 542 543 pos = pos.reshape(-1) # (M,) 544 out = torch.einsum("m,d->md", pos, omega) # (M, D/2), outer product 545 546 emb_sin = torch.sin(out) # (M, D/2) 547 emb_cos = torch.cos(out) # (M, D/2) 548 549 emb = torch.cat([emb_sin, emb_cos], dim=1) # (M, D) 550 return emb
embed_dim: output dimension for each position pos: a list of positions to be encoded: size (M,) out: (M, D)
553class PatchEmbedUnSafe(PatchEmbed): 554 """Image to Patch Embedding""" 555 556 def forward(self, x): 557 B, C, H, W = x.shape 558 559 # NOTE: Comment code from ScaleMAE: Dropped size check in timm 560 # assert H == self.img_size[0] and W == self.img_size[1], \ 561 # f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]})." 562 563 x = self.proj(x).flatten(2).transpose(1, 2) 564 return x
Image to Patch Embedding
556 def forward(self, x): 557 B, C, H, W = x.shape 558 559 # NOTE: Comment code from ScaleMAE: Dropped size check in timm 560 # assert H == self.img_size[0] and W == self.img_size[1], \ 561 # f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]})." 562 563 x = self.proj(x).flatten(2).transpose(1, 2) 564 return x
567class ViT_ScaleMAE(VisionTransformer): 568 """Vision Transformer dervied from the Scale Masked Auto Encoder codebase (TODO: paper and github link). 569 570 NOTE: For downstream tasks, the "base_resoulution" parameter needs to be adjusted manually when using 571 the model on a different zoom factor dataset. 572 """ 573 574 def __init__( 575 self, img_size=224, patch_size=16, in_chans=3, embed_dim=1024, depth=12, base_resolution=2.5, **kwargs 576 ): 577 super().__init__(img_size=img_size, embed_dim=embed_dim, **kwargs) 578 self.img_size = img_size 579 self.in_chans = in_chans 580 self.depth = depth 581 self.base_resolution = base_resolution 582 583 self.patch_embed = PatchEmbedUnSafe( 584 img_size=img_size, 585 patch_size=patch_size, 586 in_chans=in_chans, 587 embed_dim=embed_dim, 588 ) 589 590 def transform_inputs(self, x): 591 import kornia.augmentation as K 592 from kornia.constants import Resample 593 594 self._transforms = CustomCompose( 595 rescale_transform=K.RandomResizedCrop( 596 (448, 448), 597 ratio=(1.0, 1.0), 598 scale=(1.0, 1.0), 599 resample=Resample.BICUBIC.name, 600 ), 601 other_transforms=None, 602 src_transform=K.Resize((224, 224)), 603 ) 604 x, _, ratios, _, _ = self._transforms(x) 605 input_res = ratios * self.base_resolution 606 return x, input_res 607 608 def convert_to_expected_dim(self, x): 609 inputs_ = x[:, 1:, :] # removing the class tokens 610 # reshape the outputs to desired shape (N X H*W X C -> N X H X W X C) 611 rdim = inputs_.shape[1] 612 dshape = int(rdim ** 0.5) # finding square root of the outputs for obtaining the patch shape 613 inputs_ = torch.unflatten(inputs_, 1, (dshape, dshape)) 614 inputs_ = inputs_.permute(0, 3, 1, 2) 615 return inputs_ 616 617 def forward_features(self, x): 618 x, input_res = self.transform_inputs(x) 619 620 B, _, h, w = x.shape 621 x = self.patch_embed(x) 622 623 num_patches = int((h * w) / (self.patch_embed.patch_size[0] * self.patch_embed.patch_size[1])) 624 pos_embed = get_2d_sincos_pos_embed_with_resolution( 625 x.shape[-1], 626 int(num_patches ** 0.5), 627 input_res, 628 cls_token=True, 629 device=x.device, 630 ) 631 632 cls_tokens = self.cls_token.expand(B, -1, -1) # stole cls_tokens impl from Phil Wang, thanks 633 x = torch.cat((cls_tokens, x), dim=1) 634 x = x + pos_embed 635 x = self.pos_drop(x) 636 637 # chunks obtained for getting the projections for conjuctions with upsampling blocks 638 _chunks = int(self.depth / 4) 639 chunks_for_projection = [_chunks - 1, 2*_chunks - 1, 3*_chunks - 1, 4*_chunks - 1] 640 641 list_from_encoder = [] 642 for i, blk in enumerate(self.blocks): 643 x = blk(x) 644 if i in chunks_for_projection: 645 list_from_encoder.append(self.convert_to_expected_dim(x)) 646 647 x = self.convert_to_expected_dim(x) 648 649 return x, list_from_encoder 650 651 def forward(self, x): 652 x, list_from_encoder = self.forward_features(x) 653 return x, list_from_encoder
Vision Transformer dervied from the Scale Masked Auto Encoder codebase (TODO: paper and github link).
NOTE: For downstream tasks, the "base_resoulution" parameter needs to be adjusted manually when using the model on a different zoom factor dataset.
574 def __init__( 575 self, img_size=224, patch_size=16, in_chans=3, embed_dim=1024, depth=12, base_resolution=2.5, **kwargs 576 ): 577 super().__init__(img_size=img_size, embed_dim=embed_dim, **kwargs) 578 self.img_size = img_size 579 self.in_chans = in_chans 580 self.depth = depth 581 self.base_resolution = base_resolution 582 583 self.patch_embed = PatchEmbedUnSafe( 584 img_size=img_size, 585 patch_size=patch_size, 586 in_chans=in_chans, 587 embed_dim=embed_dim, 588 )
590 def transform_inputs(self, x): 591 import kornia.augmentation as K 592 from kornia.constants import Resample 593 594 self._transforms = CustomCompose( 595 rescale_transform=K.RandomResizedCrop( 596 (448, 448), 597 ratio=(1.0, 1.0), 598 scale=(1.0, 1.0), 599 resample=Resample.BICUBIC.name, 600 ), 601 other_transforms=None, 602 src_transform=K.Resize((224, 224)), 603 ) 604 x, _, ratios, _, _ = self._transforms(x) 605 input_res = ratios * self.base_resolution 606 return x, input_res
608 def convert_to_expected_dim(self, x): 609 inputs_ = x[:, 1:, :] # removing the class tokens 610 # reshape the outputs to desired shape (N X H*W X C -> N X H X W X C) 611 rdim = inputs_.shape[1] 612 dshape = int(rdim ** 0.5) # finding square root of the outputs for obtaining the patch shape 613 inputs_ = torch.unflatten(inputs_, 1, (dshape, dshape)) 614 inputs_ = inputs_.permute(0, 3, 1, 2) 615 return inputs_
617 def forward_features(self, x): 618 x, input_res = self.transform_inputs(x) 619 620 B, _, h, w = x.shape 621 x = self.patch_embed(x) 622 623 num_patches = int((h * w) / (self.patch_embed.patch_size[0] * self.patch_embed.patch_size[1])) 624 pos_embed = get_2d_sincos_pos_embed_with_resolution( 625 x.shape[-1], 626 int(num_patches ** 0.5), 627 input_res, 628 cls_token=True, 629 device=x.device, 630 ) 631 632 cls_tokens = self.cls_token.expand(B, -1, -1) # stole cls_tokens impl from Phil Wang, thanks 633 x = torch.cat((cls_tokens, x), dim=1) 634 x = x + pos_embed 635 x = self.pos_drop(x) 636 637 # chunks obtained for getting the projections for conjuctions with upsampling blocks 638 _chunks = int(self.depth / 4) 639 chunks_for_projection = [_chunks - 1, 2*_chunks - 1, 3*_chunks - 1, 4*_chunks - 1] 640 641 list_from_encoder = [] 642 for i, blk in enumerate(self.blocks): 643 x = blk(x) 644 if i in chunks_for_projection: 645 list_from_encoder.append(self.convert_to_expected_dim(x)) 646 647 x = self.convert_to_expected_dim(x) 648 649 return x, list_from_encoder
656class ViT_DINOv2(DinoV2VisionTransformer): 657 """Vision Transformer derived from the DINOv2 Codebase (https://arxiv.org/abs/2304.07193). 658 659 Based on: 660 https://github.com/facebookresearch/dinov2/blob/main/dinov2/models/vision_transformer.py. 661 662 Args: 663 img_size: The input image size. 664 patch_size: The patch size. 665 depth: The depth of the network. 666 num_register_tokens: The number of registers added (in addition to the class tokens). 667 It's important to know for ViTs trained with registers, to remove them at the end. 668 """ 669 def __init__( 670 self, 671 img_size: int = 224, 672 patch_size: int = 16, 673 depth: int = 12, 674 num_register_tokens: int = 0, 675 **kwargs 676 ): 677 if not _dinov2_import_success: 678 raise RuntimeError( 679 "The vision transformer backend can only be initialized if DINOv2 is installed. " 680 "Please install DINOv2 from https://github.com/facebookresearch/dinov2 " 681 "and then rerun your code." 682 ) 683 684 super().__init__( 685 img_size=img_size, 686 depth=depth, 687 patch_size=patch_size, 688 num_register_tokens=num_register_tokens, 689 **kwargs 690 ) 691 692 self.img_size = img_size 693 self.num_register_tokens = num_register_tokens 694 self.patch_size = patch_size 695 self.attn_outs = [i for i in range(depth) if i % 3 == 2] 696 697 def forward(self, x, masks=None) -> torch.Tensor: 698 699 B = x.shape[0] 700 701 x = self.prepare_tokens_with_masks(x) 702 703 list_of_encoder = [] 704 for i, blk in enumerate(self.blocks): 705 x = blk(x) 706 if i in self.attn_outs: 707 list_of_encoder.append(x) 708 709 x = self.norm(x) 710 x = x[:, self.num_register_tokens + 1:].reshape( 711 B, self.img_size // self.patch_size, self.img_size // self.patch_size, -1 712 ).permute(0, 3, 1, 2).contiguous() 713 714 list_of_encoder = [ 715 o[:, self.num_register_tokens + 1:].reshape( 716 B, self.img_size // self.patch_size, self.img_size // self.patch_size, -1 717 ).permute(0, 3, 1, 2).contiguous() for o in list_of_encoder 718 ] 719 720 return x, list_of_encoder[:3]
Vision Transformer derived from the DINOv2 Codebase (https://arxiv.org/abs/2304.07193).
Based on: https://github.com/facebookresearch/dinov2/blob/main/dinov2/models/vision_transformer.py.
Arguments:
- img_size: The input image size.
- patch_size: The patch size.
- depth: The depth of the network.
- num_register_tokens: The number of registers added (in addition to the class tokens). It's important to know for ViTs trained with registers, to remove them at the end.
669 def __init__( 670 self, 671 img_size: int = 224, 672 patch_size: int = 16, 673 depth: int = 12, 674 num_register_tokens: int = 0, 675 **kwargs 676 ): 677 if not _dinov2_import_success: 678 raise RuntimeError( 679 "The vision transformer backend can only be initialized if DINOv2 is installed. " 680 "Please install DINOv2 from https://github.com/facebookresearch/dinov2 " 681 "and then rerun your code." 682 ) 683 684 super().__init__( 685 img_size=img_size, 686 depth=depth, 687 patch_size=patch_size, 688 num_register_tokens=num_register_tokens, 689 **kwargs 690 ) 691 692 self.img_size = img_size 693 self.num_register_tokens = num_register_tokens 694 self.patch_size = patch_size 695 self.attn_outs = [i for i in range(depth) if i % 3 == 2]
697 def forward(self, x, masks=None) -> torch.Tensor: 698 699 B = x.shape[0] 700 701 x = self.prepare_tokens_with_masks(x) 702 703 list_of_encoder = [] 704 for i, blk in enumerate(self.blocks): 705 x = blk(x) 706 if i in self.attn_outs: 707 list_of_encoder.append(x) 708 709 x = self.norm(x) 710 x = x[:, self.num_register_tokens + 1:].reshape( 711 B, self.img_size // self.patch_size, self.img_size // self.patch_size, -1 712 ).permute(0, 3, 1, 2).contiguous() 713 714 list_of_encoder = [ 715 o[:, self.num_register_tokens + 1:].reshape( 716 B, self.img_size // self.patch_size, self.img_size // self.patch_size, -1 717 ).permute(0, 3, 1, 2).contiguous() for o in list_of_encoder 718 ] 719 720 return x, list_of_encoder[:3]
723class ViT_DINOv3(DinoV3VisionTransformer): 724 """Vision Transformer derived from the DINOv3 Codebase (https://arxiv.org/abs/2508.10104). 725 726 Based on: 727 https://github.com/facebookresearch/dinov3/blob/main/dinov3/models/vision_transformer.py. 728 729 Args: 730 img_size: The input image size. 731 patch_size: The patch size. 732 embed_dim: The embedding dimension. 733 depth: The depth of the network. 734 num_heads: The number of heads. 735 ffn_ratio: The FFN rato. 736 n_storage_tokens: The number of storage (class) tokens to remove. 737 kwargs: Keyword arguments for the image encoder base class. 738 """ 739 def __init__( 740 self, 741 in_chans: int = 3, 742 img_size: int = 224, 743 patch_size: int = 16, 744 embed_dim: int = 768, 745 depth: int = 12, 746 num_heads: int = 12, 747 ffn_ratio: float = 4.0, 748 n_storage_tokens: int = 0, 749 **kwargs 750 ): 751 if not _dinov3_import_success: 752 raise RuntimeError( 753 "The vision transformer backend can only be initialized if DINOv3 is installed. " 754 "Please install DINOv3 from https://github.com/facebookresearch/dinov3 " 755 "and then rerun your code." 756 ) 757 758 super().__init__( 759 in_chans=in_chans, 760 img_size=img_size, 761 patch_size=patch_size, 762 embed_dim=embed_dim, 763 depth=depth, 764 num_heads=num_heads, 765 ffn_ratio=ffn_ratio, 766 n_storage_tokens=n_storage_tokens, 767 **kwargs 768 ) 769 770 self.in_chans = in_chans 771 self.img_size = img_size 772 self.n_storage_tokens = n_storage_tokens 773 self.attn_outs = [i for i in range(depth) if i % 3 == 2] 774 775 def forward(self, x) -> torch.Tensor: 776 777 B = x.shape[0] 778 779 x, hw_tuple = self.prepare_tokens_with_masks(x) 780 781 list_of_encoder = [] 782 for i, blk in enumerate(self.blocks): 783 rope_sincos = self.rope_embed(H=hw_tuple[0], W=hw_tuple[1]) 784 x = blk(x, rope_sincos) 785 if i in self.attn_outs: 786 list_of_encoder.append(x) 787 788 x = self.norm(x) 789 x = x[:, self.n_storage_tokens + 1:].reshape( 790 B, self.img_size // self.patch_size, self.img_size // self.patch_size, -1 791 ).permute(0, 3, 1, 2).contiguous() 792 793 list_of_encoder = [ 794 o[:, self.n_storage_tokens + 1:].reshape( 795 B, self.img_size // self.patch_size, self.img_size // self.patch_size, -1 796 ).permute(0, 3, 1, 2).contiguous() for o in list_of_encoder 797 ] 798 799 return x, list_of_encoder[:3]
Vision Transformer derived from the DINOv3 Codebase (https://arxiv.org/abs/2508.10104).
Based on: https://github.com/facebookresearch/dinov3/blob/main/dinov3/models/vision_transformer.py.
Arguments:
- img_size: The input image size.
- patch_size: The patch size.
- embed_dim: The embedding dimension.
- depth: The depth of the network.
- num_heads: The number of heads.
- ffn_ratio: The FFN rato.
- n_storage_tokens: The number of storage (class) tokens to remove.
- kwargs: Keyword arguments for the image encoder base class.
739 def __init__( 740 self, 741 in_chans: int = 3, 742 img_size: int = 224, 743 patch_size: int = 16, 744 embed_dim: int = 768, 745 depth: int = 12, 746 num_heads: int = 12, 747 ffn_ratio: float = 4.0, 748 n_storage_tokens: int = 0, 749 **kwargs 750 ): 751 if not _dinov3_import_success: 752 raise RuntimeError( 753 "The vision transformer backend can only be initialized if DINOv3 is installed. " 754 "Please install DINOv3 from https://github.com/facebookresearch/dinov3 " 755 "and then rerun your code." 756 ) 757 758 super().__init__( 759 in_chans=in_chans, 760 img_size=img_size, 761 patch_size=patch_size, 762 embed_dim=embed_dim, 763 depth=depth, 764 num_heads=num_heads, 765 ffn_ratio=ffn_ratio, 766 n_storage_tokens=n_storage_tokens, 767 **kwargs 768 ) 769 770 self.in_chans = in_chans 771 self.img_size = img_size 772 self.n_storage_tokens = n_storage_tokens 773 self.attn_outs = [i for i in range(depth) if i % 3 == 2]
775 def forward(self, x) -> torch.Tensor: 776 777 B = x.shape[0] 778 779 x, hw_tuple = self.prepare_tokens_with_masks(x) 780 781 list_of_encoder = [] 782 for i, blk in enumerate(self.blocks): 783 rope_sincos = self.rope_embed(H=hw_tuple[0], W=hw_tuple[1]) 784 x = blk(x, rope_sincos) 785 if i in self.attn_outs: 786 list_of_encoder.append(x) 787 788 x = self.norm(x) 789 x = x[:, self.n_storage_tokens + 1:].reshape( 790 B, self.img_size // self.patch_size, self.img_size // self.patch_size, -1 791 ).permute(0, 3, 1, 2).contiguous() 792 793 list_of_encoder = [ 794 o[:, self.n_storage_tokens + 1:].reshape( 795 B, self.img_size // self.patch_size, self.img_size // self.patch_size, -1 796 ).permute(0, 3, 1, 2).contiguous() for o in list_of_encoder 797 ] 798 799 return x, list_of_encoder[:3]
802class ViT_Torchvision(nn.Module): 803 """Vision Transformer from torchvision (https://arxiv.org/abs/2010.11929). 804 805 Wraps torchvision ViT models for use as UNETR encoders. Intermediate patch-token grids 806 are collected at quarter-depth intervals and returned as spatial feature maps alongside 807 the final encoder output, matching the `(x, list_from_encoder)` contract of other ViT classes. 808 809 Supported models: vit_b_16, vit_b_32, vit_l_16, vit_l_32, vit_h_14. 810 811 All five models can be used with UNETR. vit_b_16 and vit_l_16 (patch_size=16) work with 812 any skip-connection setting. vit_b_32, vit_l_32, and vit_h_14 require use_skip_connection=False 813 in UNETR; the decoder's internal cropping handles the spatial size difference, and postprocess_masks 814 resizes the output back to the input resolution. 815 816 Args: 817 model_name: Torchvision ViT model name (e.g. 'vit_b_16'). 818 img_size: Expected input image size used by UNETR preprocessing. 819 in_chans: Number of input channels. If != 3, a 1x1 conv projects to 3 channels. 820 pretrained: Whether to load ImageNet-pretrained weights. 821 """ 822 def __init__( 823 self, 824 model_name: str, 825 img_size: int = 224, 826 in_chans: int = 3, 827 pretrained: bool = True, 828 ): 829 super().__init__() 830 if not _torchvision_import_success: 831 raise RuntimeError( 832 "The vision transformer backend can only be initialized if torchvision is installed. " 833 "Please install torchvision from https://github.com/pytorch/vision and then rerun your code." 834 ) 835 836 fn = getattr(_tv_models, model_name) 837 backbone = fn(weights="DEFAULT" if pretrained else None) 838 839 self.conv_proj = backbone.conv_proj 840 self.class_token = backbone.class_token 841 self.encoder = backbone.encoder # pos_embedding, dropout, layers, ln 842 843 self.img_size = img_size 844 self.in_chans = in_chans 845 self.embed_dim = backbone.hidden_dim 846 847 depth = len(backbone.encoder.layers) 848 _c = depth // 4 849 self.chunks_for_projection = [_c - 1, 2 * _c - 1, 3 * _c - 1] 850 851 self.input_proj = nn.Conv2d(in_chans, 3, kernel_size=1) if in_chans != 3 else None 852 853 def _load_from_state_dict( 854 self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs, 855 ): 856 pos_embed = state_dict.get(prefix + "encoder.pos_embedding") 857 current = self.encoder.pos_embedding 858 if ( 859 pos_embed is not None and pos_embed.ndim == 3 860 and pos_embed.shape[0] == current.shape[0] and pos_embed.shape[2] == current.shape[2] 861 and pos_embed.shape[1] != current.shape[1] 862 ): 863 self.encoder.pos_embedding = nn.Parameter( 864 current.new_empty(pos_embed.shape), requires_grad=current.requires_grad, 865 ) 866 super()._load_from_state_dict( 867 state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs, 868 ) 869 870 def _interpolate_pos_embed(self, pos_embed: torch.Tensor, H_p: int, W_p: int) -> torch.Tensor: 871 cls_pos, patch_pos = pos_embed[:, :1], pos_embed[:, 1:] 872 N = patch_pos.shape[1] 873 H_t = W_t = int(N ** 0.5) 874 patch_pos = patch_pos.reshape(1, H_t, W_t, -1).permute(0, 3, 1, 2) 875 patch_pos = F.interpolate(patch_pos, size=(H_p, W_p), mode="bicubic", align_corners=False) 876 patch_pos = patch_pos.permute(0, 2, 3, 1).reshape(1, H_p * W_p, -1) 877 return torch.cat([cls_pos, patch_pos], dim=1) 878 879 def forward(self, x: torch.Tensor) -> torch.Tensor: 880 """Apply the vision transformer to input data. 881 882 Args: 883 x: The input data. 884 885 Returns: 886 The vision transformer output. 887 """ 888 if self.input_proj is not None: 889 x = self.input_proj(x) 890 891 x = self.conv_proj(x) # (B, D, H_p, W_p) 892 B, D, H_p, W_p = x.shape 893 x = x.reshape(B, D, H_p * W_p).permute(0, 2, 1) # (B, N, D) 894 895 cls = self.class_token.expand(B, -1, -1) 896 x = torch.cat([cls, x], dim=1) # (B, 1+N, D) 897 898 pos = self.encoder.pos_embedding 899 if pos.shape[1] != x.shape[1]: 900 pos = self._interpolate_pos_embed(pos, H_p, W_p) 901 x = x + pos 902 x = self.encoder.dropout(x) 903 904 list_from_encoder = [] 905 for i, blk in enumerate(self.encoder.layers): 906 x = blk(x) 907 if i in self.chunks_for_projection: 908 feat = x[:, 1:].reshape(B, H_p, W_p, D).permute(0, 3, 1, 2).contiguous() 909 list_from_encoder.append(feat) 910 911 x = self.encoder.ln(x) 912 x = x[:, 1:].reshape(B, H_p, W_p, D).permute(0, 3, 1, 2).contiguous() 913 return x, list_from_encoder
Vision Transformer from torchvision (https://arxiv.org/abs/2010.11929).
Wraps torchvision ViT models for use as UNETR encoders. Intermediate patch-token grids
are collected at quarter-depth intervals and returned as spatial feature maps alongside
the final encoder output, matching the (x, list_from_encoder) contract of other ViT classes.
Supported models: vit_b_16, vit_b_32, vit_l_16, vit_l_32, vit_h_14.
All five models can be used with UNETR. vit_b_16 and vit_l_16 (patch_size=16) work with any skip-connection setting. vit_b_32, vit_l_32, and vit_h_14 require use_skip_connection=False in UNETR; the decoder's internal cropping handles the spatial size difference, and postprocess_masks resizes the output back to the input resolution.
Arguments:
- model_name: Torchvision ViT model name (e.g. 'vit_b_16').
- img_size: Expected input image size used by UNETR preprocessing.
- in_chans: Number of input channels. If != 3, a 1x1 conv projects to 3 channels.
- pretrained: Whether to load ImageNet-pretrained weights.
822 def __init__( 823 self, 824 model_name: str, 825 img_size: int = 224, 826 in_chans: int = 3, 827 pretrained: bool = True, 828 ): 829 super().__init__() 830 if not _torchvision_import_success: 831 raise RuntimeError( 832 "The vision transformer backend can only be initialized if torchvision is installed. " 833 "Please install torchvision from https://github.com/pytorch/vision and then rerun your code." 834 ) 835 836 fn = getattr(_tv_models, model_name) 837 backbone = fn(weights="DEFAULT" if pretrained else None) 838 839 self.conv_proj = backbone.conv_proj 840 self.class_token = backbone.class_token 841 self.encoder = backbone.encoder # pos_embedding, dropout, layers, ln 842 843 self.img_size = img_size 844 self.in_chans = in_chans 845 self.embed_dim = backbone.hidden_dim 846 847 depth = len(backbone.encoder.layers) 848 _c = depth // 4 849 self.chunks_for_projection = [_c - 1, 2 * _c - 1, 3 * _c - 1] 850 851 self.input_proj = nn.Conv2d(in_chans, 3, kernel_size=1) if in_chans != 3 else None
Initialize internal Module state, shared by both nn.Module and ScriptModule.
879 def forward(self, x: torch.Tensor) -> torch.Tensor: 880 """Apply the vision transformer to input data. 881 882 Args: 883 x: The input data. 884 885 Returns: 886 The vision transformer output. 887 """ 888 if self.input_proj is not None: 889 x = self.input_proj(x) 890 891 x = self.conv_proj(x) # (B, D, H_p, W_p) 892 B, D, H_p, W_p = x.shape 893 x = x.reshape(B, D, H_p * W_p).permute(0, 2, 1) # (B, N, D) 894 895 cls = self.class_token.expand(B, -1, -1) 896 x = torch.cat([cls, x], dim=1) # (B, 1+N, D) 897 898 pos = self.encoder.pos_embedding 899 if pos.shape[1] != x.shape[1]: 900 pos = self._interpolate_pos_embed(pos, H_p, W_p) 901 x = x + pos 902 x = self.encoder.dropout(x) 903 904 list_from_encoder = [] 905 for i, blk in enumerate(self.encoder.layers): 906 x = blk(x) 907 if i in self.chunks_for_projection: 908 feat = x[:, 1:].reshape(B, H_p, W_p, D).permute(0, 3, 1, 2).contiguous() 909 list_from_encoder.append(feat) 910 911 x = self.encoder.ln(x) 912 x = x[:, 1:].reshape(B, H_p, W_p, D).permute(0, 3, 1, 2).contiguous() 913 return x, list_from_encoder
Apply the vision transformer to input data.
Arguments:
- x: The input data.
Returns:
The vision transformer output.
916def get_vision_transformer(backbone: str, model: str, img_size: int = 1024, **kwargs) -> nn.Module: 917 """Get vision transformer encoder. 918 919 Args: 920 backbone: The name of the vision transformer implementation. 921 One of "sam" / "cellpose_sam" / "sam2" / "sam3" / "mae" / "scalemae" / "dinov2" / "dinov3" / "torchvision". 922 model: The name of the model. One of "vit_b", "vit_l" or "vit_h". 923 img_size: The size of the input for the image encoder. Input images will be resized to match this size. 924 kwargs: Additional kwargs which can be expected by the vision transformer, 925 e.g. 'base_resolution' for `ViT_ScaleMAE`. 926 927 Returns: 928 The vision transformer. 929 """ 930 if backbone == "sam": 931 if model == "vit_b": 932 encoder = ViT_Sam( 933 depth=12, embed_dim=768, img_size=img_size, mlp_ratio=4, 934 norm_layer=partial(torch.nn.LayerNorm, eps=1e-6), 935 num_heads=12, patch_size=16, qkv_bias=True, use_rel_pos=True, 936 global_attn_indexes=[2, 5, 8, 11], 937 window_size=14, out_chans=256, 938 ) 939 elif model == "vit_l": 940 encoder = ViT_Sam( 941 depth=24, embed_dim=1024, img_size=img_size, mlp_ratio=4, 942 norm_layer=partial(torch.nn.LayerNorm, eps=1e-6), 943 num_heads=16, patch_size=16, qkv_bias=True, use_rel_pos=True, 944 global_attn_indexes=[5, 11, 17, 23], 945 window_size=14, out_chans=256, 946 ) 947 elif model == "vit_h": 948 encoder = ViT_Sam( 949 depth=32, embed_dim=1280, img_size=img_size, mlp_ratio=4, 950 norm_layer=partial(torch.nn.LayerNorm, eps=1e-6), 951 num_heads=16, patch_size=16, qkv_bias=True, use_rel_pos=True, 952 global_attn_indexes=[7, 15, 23, 31], 953 window_size=14, out_chans=256, 954 ) 955 else: 956 raise ValueError(f"'{model}' is not supported by SAM. Currently, 'vit_b', 'vit_l', 'vit_h' are supported.") 957 958 elif backbone == "cellpose_sam": 959 if model != "vit_l": 960 raise ValueError(f"'{model}' is not supported by CellposeSAM. Only 'vit_l' is supported.") 961 encoder = ViT_CellposeSAM(ps=8, bsize=img_size) 962 963 elif backbone == "sam2": 964 if model == "hvit_t": 965 encoder = ViT_Sam2( 966 img_size=img_size, embed_dim=96, num_heads=1, stages=[1, 2, 7, 2], global_att_blocks=[5, 7, 9], 967 window_pos_embed_bkg_spatial_size=[7, 7], backbone_channel_list=[768, 384, 192, 96], 968 ) 969 elif model == "hvit_s": 970 encoder = ViT_Sam2( 971 img_size=img_size, embed_dim=96, num_heads=1, stages=[1, 2, 11, 2], global_att_blocks=[7, 10, 13], 972 window_pos_embed_bkg_spatial_size=[7, 7], backbone_channel_list=[768, 384, 192, 96], 973 ) 974 elif model == "hvit_b": 975 encoder = ViT_Sam2( 976 img_size=img_size, embed_dim=112, num_heads=2, backbone_channel_list=[896, 448, 224, 112], 977 ) 978 elif model == "hvit_l": 979 encoder = ViT_Sam2( 980 img_size=img_size, embed_dim=144, num_heads=2, stages=[2, 6, 36, 4], global_att_blocks=[23, 33, 43], 981 window_spec=[8, 4, 16, 8], backbone_channel_list=[1152, 576, 288, 144], 982 ) 983 else: 984 raise ValueError( 985 f"'{model}' is not supported by SAM2. Currently, 'hvit_t', 'hvit_s', 'hvit_b', 'hvit_l' are supported." 986 ) 987 988 elif backbone == "sam3": 989 if model != "vit_pe": 990 raise ValueError( 991 "'sam3' does not have multiple model configurations. Please use 'vit_pe' as the model configuration." 992 ) 993 994 encoder = ViT_Sam3( 995 img_size=1008, pretrain_img_size=336, patch_size=14, embed_dim=1024, depth=32, num_heads=16, 996 mlp_ratio=4.625, norm_layer="LayerNorm", drop_path_rate=0.1, qkv_bias=True, use_abs_pos=True, 997 tile_abs_pos=True, global_att_blocks=(7, 15, 23, 31), rel_pos_blocks=(), use_rope=True, 998 use_interp_rope=True, window_size=24, pretrain_use_cls_token=True, retain_cls_token=False, ln_pre=True, 999 ln_post=False, return_interm_layers=False, bias_patch_embed=False, compile_mode=None, 1000 ) 1001 1002 elif backbone == "mae": 1003 if model == "vit_b": 1004 encoder = ViT_MAE( 1005 img_size=img_size, patch_size=16, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4, 1006 qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6) 1007 ) 1008 elif model == "vit_l": 1009 encoder = ViT_MAE( 1010 img_size=img_size, patch_size=16, embed_dim=1024, depth=24, num_heads=16, mlp_ratio=4, 1011 qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6) 1012 ) 1013 elif model == "vit_h": 1014 encoder = ViT_MAE( 1015 img_size=img_size, patch_size=14, embed_dim=1280, depth=32, num_heads=16, mlp_ratio=4, 1016 qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6) 1017 ) 1018 else: 1019 raise ValueError(f"'{model}' is not supported by MAE. Currently, 'vit_b', 'vit_l', 'vit_h' are supported.") 1020 1021 elif backbone == "scalemae": 1022 base_resolution = kwargs.get("base_resolution", 2.5) 1023 1024 if model == "vit_b": 1025 encoder = ViT_ScaleMAE( 1026 img_size=img_size, patch_size=8, embed_dim=768, depth=12, num_heads=12, 1027 mlp_ratio=4, qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), 1028 base_resolution=base_resolution, 1029 ) 1030 elif model == "vit_l": 1031 encoder = ViT_ScaleMAE( 1032 img_size=img_size, patch_size=8, embed_dim=1024, depth=24, num_heads=16, 1033 mlp_ratio=4, qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), 1034 base_resolution=base_resolution, 1035 ) 1036 elif model == "vit_h": 1037 encoder = ViT_ScaleMAE( 1038 img_size=img_size, patch_size=8, embed_dim=1280, depth=32, num_heads=16, 1039 mlp_ratio=4, qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), 1040 base_resolution=base_resolution, 1041 ) 1042 else: 1043 raise ValueError( 1044 f"'{model}' is not supported by ScaleMAE. Currently, 'vit_b', 'vit_l' and 'vit_h' are supported." 1045 ) 1046 1047 elif backbone == "dinov2": 1048 block_fn = partial(Block, attn_class=MemEffAttention) 1049 msg = "The model name should be either 'vit_<X>' or 'vit_<X>_reg<Y>." 1050 1051 if model.startswith("vit_s"): 1052 assert model in ["vit_s", "vit_s_reg4"], msg 1053 encoder = ViT_DINOv2( 1054 img_size=img_size, patch_size=14, embed_dim=384, depth=12, num_heads=6, mlp_ratio=4, 1055 block_fn=block_fn, in_chans=3, channel_adaptive=False, init_values=1e-5, block_chunks=0, 1056 num_register_tokens=4 if model.endswith("_reg4") else 0, 1057 ) 1058 elif model.startswith("vit_b"): 1059 assert model in ["vit_b", "vit_b_reg4"], msg 1060 encoder = ViT_DINOv2( 1061 img_size=img_size, patch_size=14, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4, 1062 block_fn=block_fn, in_chans=3, channel_adaptive=False, init_values=1e-5, block_chunks=0, 1063 num_register_tokens=4 if model.endswith("_reg4") else 0, 1064 ) 1065 elif model.startswith("vit_l"): 1066 assert model in ["vit_l", "vit_l_reg4"], msg 1067 encoder = ViT_DINOv2( 1068 img_size=img_size, patch_size=14, embed_dim=1024, depth=24, num_heads=16, mlp_ratio=4, 1069 block_fn=block_fn, in_chans=3, channel_adaptive=False, init_values=1e-5, block_chunks=0, 1070 num_register_tokens=4 if model.endswith("_reg4") else 0, 1071 ) 1072 elif model.startswith("vit_g"): 1073 assert model in ["vit_g", "vit_g_reg4"], msg 1074 encoder = ViT_DINOv2( 1075 img_size=img_size, patch_size=14, embed_dim=1536, depth=40, num_heads=24, mlp_ratio=4, 1076 block_fn=block_fn, in_chans=3, channel_adaptive=False, init_values=1e-5, block_chunks=0, 1077 num_register_tokens=4 if model.endswith("_reg4") else 0, ffn_layer="swiglu", 1078 ) 1079 else: 1080 raise ValueError( 1081 f"'{model}' is not supported by DINOv2. Currently, 'vit_s', 'vit_b', 'vit_l' and 'vit_g' are supported." 1082 ) 1083 1084 elif backbone == "dinov3": 1085 1086 if model == "vit_s": 1087 encoder = ViT_DINOv3( 1088 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", embed_dim=384, 1089 num_heads=6, layerscale_init=1.0e-05, norm_layer="layernormbf16", n_storage_tokens=4, mask_k_bias=True, 1090 ) 1091 elif model == "vit_s+": 1092 encoder = ViT_DINOv3( 1093 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", embed_dim=384, 1094 num_heads=6, ffn_ratio=6, layerscale_init=1.0e-05, norm_layer="layernormbf16", 1095 ffn_layer="swiglu", n_storage_tokens=4, mask_k_bias=True, 1096 ) 1097 1098 elif model == "vit_b": 1099 encoder = ViT_DINOv3( 1100 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", 1101 layerscale_init=1.0e-05, norm_layer="layernormbf16", n_storage_tokens=4, mask_k_bias=True, 1102 ) 1103 elif model == "vit_l": 1104 encoder = ViT_DINOv3( 1105 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", embed_dim=1024, 1106 depth=24, num_heads=16, layerscale_init=1.0e-05, norm_layer="layernormbf16", 1107 n_storage_tokens=4, mask_k_bias=True, 1108 ) 1109 elif model == "vit_l+": 1110 encoder = ViT_DINOv3( 1111 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", embed_dim=1024, 1112 depth=24, num_heads=16, ffn_ratio=6.0, layerscale_init=1.0e-05, norm_layer="layernormbf16", 1113 ffn_layer="swiglu", n_storage_tokens=4, mask_k_bias=True, 1114 ) 1115 elif model == "vit_h+": 1116 encoder = ViT_DINOv3( 1117 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", embed_dim=1280, 1118 depth=32, num_heads=20, ffn_ratio=6.0, layerscale_init=1.0e-05, norm_layer="layernormbf16", 1119 ffn_layer="swiglu", n_storage_tokens=4, mask_k_bias=True, 1120 ) 1121 elif model == "vit_7b": 1122 encoder = ViT_DINOv3( 1123 img_size=img_size, pos_embed_rope_rescale_coords=2, pos_embed_rope_dtype="fp32", embed_dim=4096, 1124 depth=40, num_heads=32, ffn_ratio=3, qkv_bias=False, drop_path_rate=0.0, layerscale_init=1.0e-05, 1125 norm_layer="layernormbf16", ffn_layer="swiglu64", n_storage_tokens=4, mask_k_bias=True, 1126 untie_global_and_local_cls_norm=True, 1127 ) 1128 else: 1129 raise ValueError( 1130 f"'{model}' is not supported by DINOv3. Currently, " 1131 " 'vit_s', 'vit_s+', 'vit_b', 'vit_l', 'vit_l+', 'vit_h+'. 'vit_7b' are supported." 1132 ) 1133 1134 elif backbone == "torchvision": 1135 supported = ["vit_b_16", "vit_b_32", "vit_l_16", "vit_l_32", "vit_h_14"] 1136 if model not in supported: 1137 raise ValueError(f"'{model}' is not supported by the torchvision backbone. Choose from: {supported}.") 1138 encoder = ViT_Torchvision(model_name=model, img_size=img_size, **kwargs) 1139 1140 else: 1141 raise ValueError( 1142 "The 'UNETR' supported backbones are 'sam', 'cellpose_sam', 'sam2', 'sam3', " 1143 "'mae', 'scalemae', 'dinov2', 'dinov3' or 'torchvision'. Please choose one of them." 1144 ) 1145 1146 return encoder
Get vision transformer encoder.
Arguments:
- backbone: The name of the vision transformer implementation. One of "sam" / "cellpose_sam" / "sam2" / "sam3" / "mae" / "scalemae" / "dinov2" / "dinov3" / "torchvision".
- model: The name of the model. One of "vit_b", "vit_l" or "vit_h".
- img_size: The size of the input for the image encoder. Input images will be resized to match this size.
- kwargs: Additional kwargs which can be expected by the vision transformer,
e.g. 'base_resolution' for
ViT_ScaleMAE.
Returns:
The vision transformer.