nvidia/C-RADIOv4-H
8527k
1from distutils.version import LooseVersion2from types import MethodType3from typing import List, Optional, Tuple, Union4import warnings5 6import torch7from torch import nn8import torch.nn.functional as F9 10try:11 from timm.models import register_model12except ImportError:13 from timm.models.registry import register_model14from timm.data.constants import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD15 16from .forward_intermediates import forward_intermediates17from .input_conditioner import InputConditioner18 19_has_torch_sdpa = hasattr(F, 'scaled_dot_product_attention')20 21 22class PaliGemmaWrapper(nn.Module):23 def __init__(self, vis_model: nn.Module, embed_dim: int):24 super().__init__()25 26 self.vis_model = vis_model27 self.embed_dim = embed_dim28 29 @property30 def patch_size(self):31 return self.vis_model.embeddings.patch_size32 33 @property34 def blocks(self):35 return self.vis_model.encoder.layers36 37 @property38 def embed_dim(self):39 return self.vis_model.embeddings.embed_dim40 41 def forward(self, x: torch.Tensor):42 outputs = self.vis_model(43 x,44 return_dict=False,45 interpolate_pos_encoding=True,46 )47 48 features = outputs[0].to(torch.float32)49 50 summary = features.mean(dim=1)51 52 return summary, features53 54 def forward_features(self, x: torch.Tensor):55 return self(x)56 57 58def _get_paligemma_model(repo: str, embed_dim: int = None, dtype: torch.dtype = torch.bfloat16):59 from transformers import PaliGemmaForConditionalGeneration, __version__ as tx_version60 61 if LooseVersion(tx_version) > LooseVersion('4.44.2'):62 warnings.warn(f'Your transformers version "{tx_version}" is higher than 4.44.2, and for whatever reason, PaliGemma might be broken.')63 64 extra_args = dict()65 66 if dtype is not None:67 extra_args['torch_dtype'] = dtype68 rev = str(dtype).split('.')[-1]69 extra_args['revision'] = rev70 71 model = PaliGemmaForConditionalGeneration.from_pretrained(repo, **extra_args)72 73 vis_model = model.vision_tower.vision_model74 75 vis_model = PaliGemmaWrapper(vis_model, embed_dim)76 77 return vis_model78 79@register_model80def paligemma_896_student(**kwargs):81 model = _get_paligemma_model('google/paligemma-3b-pt-896', embed_dim=1152, dtype=None)82 83 return model84 85 86def dv2_sdpa(self, x: torch.Tensor) -> torch.Tensor:87 B, N, C = x.shape88 qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)89 90 q, k, v = qkv[0], qkv[1], qkv[2]91 x = F.scaled_dot_product_attention(92 q, k, v,93 is_causal=False,94 dropout_p=self.attn_drop.p if self.training else 0.,95 scale=self.scale,96 )97 x = x.transpose(1, 2).reshape(B, N, C)98 x = self.proj(x)99 x = self.proj_drop(x)100 return x101 102def _load_dino_v2(dino_v2_model, cache_dir: Optional[str] = None, pretrained=True, **kwargs):103 if cache_dir:104 torch.hub.set_dir(cache_dir)105 model: nn.Module = torch.hub.load(106 'facebookresearch/dinov2',107 dino_v2_model,108 pretrained=pretrained,109 # **kwargs,110 )111 112 if _has_torch_sdpa:113 for n, m in model.named_modules():114 if n.endswith('.attn'):115 m.forward = MethodType(dv2_sdpa, m)116 117 return model118 119class DinoWrapper(nn.Module):120 def __init__(self, dino_model: nn.Module):121 super().__init__()122 123 self.inner = dino_model124 dino_model.blocks = nn.Sequential(*dino_model.blocks)125 126 @property127 def embed_dim(self):128 return self.inner.embed_dim129 130 @property131 def patch_size(self):132 return self.inner.patch_size133 134 @property135 def num_cls_tokens(self):136 return getattr(self.inner, 'num_tokens', 1)137 138 @property139 def num_registers(self):140 return getattr(self.inner, 'num_register_tokens', 0)141 142 @property143 def num_summary_tokens(self):144 return self.num_cls_tokens + self.num_registers145 146 @property147 def blocks(self):148 return self.inner.blocks149 150 def forward(self, *args, **kwargs) -> Tuple[torch.Tensor, torch.Tensor]:151 parts = self.inner.forward_features(*args, **kwargs)152 153 cls_token = parts['x_norm_clstoken']154 features = parts['x_norm_patchtokens']155 156 return cls_token, features157 158 def forward_features(self, x: torch.Tensor):159 x = self.inner.prepare_tokens_with_masks(x)160 x = self.inner.blocks(x)161 x_norm = self.inner.norm(x)162 163 return x_norm[:, 0], x_norm[:, self.num_summary_tokens:]164 165 def patchify(self, x: torch.Tensor) -> torch.Tensor:166 return self.inner.prepare_tokens_with_masks(x)167 168 def forward_intermediates(self,169 x: torch.Tensor,170 norm: bool = False,171 **kwargs,172 ) -> Union[List[torch.Tensor], Tuple[torch.Tensor, List[torch.Tensor]]]:173 return forward_intermediates(174 self,175 patch_extractor=self.inner.prepare_tokens_with_masks,176 num_summary_tokens=self.num_summary_tokens,177 num_cls_tokens=self.num_cls_tokens,178 norm=self.inner.norm if norm else lambda y: y,179 x=x,180 **kwargs,181 )182 183 184def _dino_student(arch: str, **kwargs):185 from . import dinov2_arch186 187 factory = getattr(dinov2_arch, arch)188 model = factory()189 190 model = DinoWrapper(model)191 192 conditioner = InputConditioner(193 input_scale=1.0,194 norm_mean=IMAGENET_DEFAULT_MEAN,195 norm_std=IMAGENET_DEFAULT_STD,196 )197 198 model.input_conditioner = conditioner199 200 return model201 202 203@register_model204def dino_v2_l_student(**kwargs):205 return _dino_student('dinov2_vitl14_reg', **kwargs)206 207@register_model208def dino_v2_g_student(**kwargs):209 return _dino_student('dinov2_vitg14_reg', **kwargs)210 