nvidia/C-RADIOv4-H
8430k
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 