Team Ai
Modelpublic

nvidia/C-RADIOv4-H

sourceHugging Faceotherupdated 8mo agoView on Hugging Face
85likes27kdownloads
extra_models.py210 linesDownload Raw Back to root
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