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
class ViT_Sam:
 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).

Based on: https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/image_encoder.py

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.
ViT_Sam( in_chans: int = 3, embed_dim: int = 768, global_attn_indexes: Tuple[int, ...] = [2, 5, 8, 11], apply_neck: bool = False, **kwargs)
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
chunks_for_projection
in_chans
embed_dim
apply_neck
def forward(self, x: torch.Tensor) -> torch.Tensor:
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.

class ViT_CellposeSAM(torch.nn.modules.module.Module):
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.
ViT_CellposeSAM(ps: int = 8, bsize: int = 256, apply_neck: bool = True)
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.

patch_embed
pos_embed
blocks
neck
embed_dim
img_size
in_chans
apply_neck
chunks_for_projection
def forward(self, x: torch.Tensor) -> torch.Tensor:
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.

class ViT_MAE:
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.
ViT_MAE(img_size: int = 1024, in_chans: int = 3, depth: int = 12, **kwargs)
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
img_size
in_chans
depth
def forward(self, x: torch.Tensor) -> torch.Tensor:
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.

class ViT_Sam2:
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.
ViT_Sam2( backbone_channel_list: List[int], img_size: int = 1024, embed_dim: int = 96, num_heads: int = 1, stages: Tuple[int, ...] = (2, 3, 16, 3), global_att_blocks: Tuple[int, ...] = (12, 16, 20), window_pos_embed_bkg_spatial_size: Tuple[int, int] = (14, 14), window_spec: Tuple[int, ...] = (8, 4, 14, 7), scalp: int = 1, **kwargs)
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
scalp
embed_dim
img_size
def forward(self, x: torch.Tensor):
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
class ViT_Sam3:
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.
ViT_Sam3(img_size: int = 1024, embed_dim: int = 768, **kwargs)
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
img_size
embed_dim
def forward(self, x: torch.Tensor):
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.

class CustomCompose:
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
CustomCompose(rescale_transform, other_transforms, src_transform)
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
rescale_transform
other_transforms
src_transform
def get_2d_sincos_pos_embed_with_resolution(embed_dim, grid_size, res, cls_token=False, device='cpu'):
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)

def get_2d_sincos_pos_embed_from_grid_torch(embed_dim, grid):
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
def get_1d_sincos_pos_embed_from_grid_torch(embed_dim, pos):
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)

class PatchEmbedUnSafe:
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

def forward(self, x):
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
class ViT_ScaleMAE:
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.

ViT_ScaleMAE( img_size=224, patch_size=16, in_chans=3, embed_dim=1024, depth=12, base_resolution=2.5, **kwargs)
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        )
img_size
in_chans
depth
base_resolution
patch_embed
def transform_inputs(self, x):
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
def convert_to_expected_dim(self, x):
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_
def forward_features(self, x):
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
def forward(self, x):
651    def forward(self, x):
652        x, list_from_encoder = self.forward_features(x)
653        return x, list_from_encoder
class ViT_DINOv2:
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.
ViT_DINOv2( img_size: int = 224, patch_size: int = 16, depth: int = 12, num_register_tokens: int = 0, **kwargs)
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]
img_size
num_register_tokens
patch_size
attn_outs
def forward(self, x, masks=None) -> torch.Tensor:
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]
class ViT_DINOv3:
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.
ViT_DINOv3( in_chans: int = 3, img_size: int = 224, patch_size: int = 16, embed_dim: int = 768, depth: int = 12, num_heads: int = 12, ffn_ratio: float = 4.0, n_storage_tokens: int = 0, **kwargs)
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]
in_chans
img_size
n_storage_tokens
attn_outs
def forward(self, x) -> torch.Tensor:
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]
class ViT_Torchvision(torch.nn.modules.module.Module):
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.
ViT_Torchvision( model_name: str, img_size: int = 224, in_chans: int = 3, pretrained: bool = True)
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.

conv_proj
class_token
encoder
img_size
in_chans
embed_dim
chunks_for_projection
input_proj
def forward(self, x: torch.Tensor) -> torch.Tensor:
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.

def get_vision_transformer( backbone: str, model: str, img_size: int = 1024, **kwargs) -> torch.nn.modules.module.Module:
 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.