CoolFace
Modelpublic

nvidia/C-RADIOv4-H

sourceHugging Faceotherupdated 8mo agoView on Hugging Face
84likes30kdownloads
enable_cpe_support.py188 linesDownload Raw Back to root
1# Copyright (c) 2023-2024, NVIDIA CORPORATION.  All rights reserved.2#3# NVIDIA CORPORATION and its licensors retain all intellectual property4# and proprietary rights in and to this software, related documentation5# and any modifications thereto.  Any use, reproduction, disclosure or6# distribution of this software and related documentation without an express7# license agreement from NVIDIA CORPORATION is strictly prohibited.8 9from contextlib import contextmanager10from typing import List, Optional, Set, Tuple, Union11from types import MethodType12 13import torch14from torch import nn15 16from timm.models import VisionTransformer, checkpoint_seq17 18from .feature_normalizer import IntermediateFeatureNormalizerBase, NullIntermediateFeatureNormalizer19 20from .extra_models import DinoWrapper21from .vit_patch_generator import ViTPatchGenerator22from .forward_intermediates import forward_intermediates23from .dual_hybrid_vit import HybridModel24 25 26def _forward_cpe(self: VisionTransformer, x: torch.Tensor) -> torch.Tensor:27    x = self.patch_generator(x)28    if getattr(self, 'grad_checkpointing', False) and not torch.jit.is_scripting():29        x = checkpoint_seq(self.blocks, x)30    else:31        x = self.blocks(x)32    x = self.norm(x)33    return x34 35 36@contextmanager37def _video_mode(self: VisionTransformer, t: int):38    """39    Context manager to temporarily set the model in video mode.40    This is used to handle models that support both image and video inputs.41    """42    original_num_frames = self.patch_generator.num_video_frames43    self.patch_generator.num_video_frames = t44    try:45        yield46    finally:47        self.patch_generator.num_video_frames = original_num_frames48 49 50def _take_indices(51        num_blocks: int,52        n: Optional[Union[int, List[int], Tuple[int]]],53) -> Tuple[Set[int], int]:54    if isinstance(n, int):55        assert n >= 056        take_indices = {x for x in range(num_blocks - n, num_blocks)}57    else:58        take_indices = {num_blocks + idx if idx < 0 else idx for idx in n}59    return take_indices, max(take_indices)60 61 62def _forward_intermediates_cpe(63        self,64        x: torch.Tensor,65        norm: bool = False,66        **kwargs,67) -> Union[List[torch.Tensor], Tuple[torch.Tensor, List[torch.Tensor]]]:68    return forward_intermediates(69        self,70        patch_extractor=self.patch_generator,71        num_summary_tokens=self.patch_generator.num_skip,72        num_cls_tokens=self.patch_generator.num_cls_tokens,73        norm=self.norm if norm else lambda y: y,74        x=x,75        **kwargs,76    )77 78 79def _forward_cpe_dinov2(self: DinoWrapper, x: torch.Tensor) -> torch.Tensor:80    y = _forward_cpe(self.inner, x)81 82    return y[:, 0], y[:, self.num_summary_tokens:]83 84 85def _forward_intermediates_cpe_dinov2(self: DinoWrapper, *args, **kwargs):86    return _forward_intermediates_cpe(self.inner, *args, **kwargs)87 88 89def _enable_cpe_for_timm_vit(model: VisionTransformer,90                             max_img_size: Union[int, Tuple[int, int]] = 1024,91                             num_cls_tokens: int = 1,92                             pos_dropout: float = 0.1,93                             register_multiple: int = Optional[None],94                             num_registers: int = Optional[None],95):96    if not isinstance(model, VisionTransformer):97        raise ValueError("CPE only support for VisionTransformer models!")98 99    patch_size = model.patch_embed.patch_size[0]100    embed_dim = model.embed_dim101    input_dims = model.patch_embed.img_size102    normalize_patches = not isinstance(model.patch_embed.norm, nn.Identity)103    cls_token = model.cls_token is not None104 105    max_img_size = int(round(max_img_size / patch_size) * patch_size)106 107    patch_generator = ViTPatchGenerator(108        patch_size=patch_size,109        embed_dim=embed_dim,110        input_dims=input_dims,111        normalize_patches=normalize_patches,112        cls_token=cls_token,113        max_input_dims=max_img_size,114        pos_dropout=pos_dropout,115        num_cls_tokens=num_cls_tokens,116        register_multiple=register_multiple,117        num_registers=num_registers,118    )119 120    model.patch_generator = patch_generator121    model.patch_embed = None122    model.cls_token = None123    model.pos_embed = None124    model.pos_drop = None125    model.patch_size = patch_size126    model.num_cls_tokens = num_cls_tokens127    model.num_registers = patch_generator.num_registers128 129    model.forward_features = MethodType(_forward_cpe, model)130    model.forward_intermediates = MethodType(_forward_intermediates_cpe, model)131 132 133def _enable_cpe_for_dv2_reg_vit(model: DinoWrapper,134                                max_img_size: Union[int, Tuple[int, int]] = 1024,135                                num_cls_tokens: int = 1,136                                pos_dropout: float = 0.1,137                                register_multiple: int = Optional[None],138                                num_registers: int = Optional[None],139):140    patch_size = model.patch_size141    embed_dim = model.embed_dim142    input_dims = model.inner.patch_embed.patches_resolution143    normalize_patches = not isinstance(model.inner.patch_embed.norm, nn.Identity)144    cls_token = True145 146    max_img_size = int(round(max_img_size / patch_size) * patch_size)147 148    patch_generator = ViTPatchGenerator(149        patch_size=patch_size,150        embed_dim=embed_dim,151        input_dims=input_dims,152        normalize_patches=normalize_patches,153        cls_token=cls_token,154        max_input_dims=max_img_size,155        pos_dropout=pos_dropout,156        num_cls_tokens=num_cls_tokens,157        register_multiple=register_multiple,158        num_registers=num_registers,159        patch_bias=True,160    )161 162    inner = model.inner163    inner.patch_generator = patch_generator164    inner.patch_embed = None165    inner.cls_token = None166    inner.pos_embed = None167    inner.register_tokens = None168    inner.patch_size = patch_size169 170    model.forward_features = MethodType(_forward_cpe_dinov2, model)171    model.forward_intermediates = MethodType(_forward_intermediates_cpe_dinov2, model)172 173 174def enable_cpe(model: nn.Module,175               *args,176               **kwargs,177):178    if isinstance(model, VisionTransformer):179        _enable_cpe_for_timm_vit(model, *args, **kwargs)180    elif isinstance(model, DinoWrapper):181        _enable_cpe_for_dv2_reg_vit(model, *args, **kwargs)182    elif isinstance(model, HybridModel):183        _enable_cpe_for_timm_vit(model.vit, *args, **kwargs)184    else:185        raise ValueError(f'CPE not supported for this model type: {type(model)}')186 187    model.cpe_video_mode = MethodType(_video_mode, model)188