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.8from typing import Callable, Dict, Iterable, List, NamedTuple, Optional, Tuple, Union9 10import torch11from torch import nn12 13from timm.models import create_model, VisionTransformer14 15from .enable_cpe_support import enable_cpe16from .input_conditioner import InputConditioner17from .adaptor_base import AdaptorBase, RadioOutput, AdaptorInput18from . import eradio_model19from .enable_spectral_reparam import configure_spectral_reparam_from_args20from .feature_normalizer import FeatureNormalizer, IntermediateFeatureNormalizer21from . import dual_hybrid_vit22 23 24class Resolution(NamedTuple):25 height: int26 width: int27 28 29class RADIOModel(nn.Module):30 def __init__(31 self,32 model: nn.Module,33 input_conditioner: InputConditioner,34 patch_size: int,35 max_resolution: int,36 preferred_resolution: Resolution,37 summary_idxs: Optional[torch.Tensor] = None,38 window_size: int = None,39 adaptors: Dict[str, AdaptorBase] = None,40 feature_normalizer: Optional[FeatureNormalizer] = None,41 inter_feature_normalizer: Optional[IntermediateFeatureNormalizer] = None,42 ):43 super().__init__()44 45 self.model = model46 self.input_conditioner = input_conditioner47 if summary_idxs is not None:48 self.register_buffer('summary_idxs', summary_idxs)49 else:50 self.summary_idxs = None51 52 self._preferred_resolution = preferred_resolution53 self._patch_size = patch_size54 self._max_resolution = max_resolution55 self._window_size = window_size56 57 adaptors = adaptors or dict()58 self.adaptors = nn.ModuleDict(adaptors)59 60 if feature_normalizer is None:61 feature_normalizer = nn.Identity()62 self.feature_normalizer = feature_normalizer63 self.inter_feature_normalizer = inter_feature_normalizer64 65 @property66 def num_summary_tokens(self) -> int:67 if hasattr(self.model, 'num_summary_tokens'):68 return self.model.num_summary_tokens69 70 patch_gen = getattr(self.model, "patch_generator", None)71 if patch_gen is not None:72 return patch_gen.num_skip73 elif getattr(self.model, 'global_pool', None) == 'avg':74 return 075 return 176 77 @property78 def num_cls_tokens(self) -> int:79 if hasattr(self.model, 'num_cls_tokens'):80 return self.model.num_cls_tokens81 82 patch_gen = getattr(self.model, 'patch_generator', None)83 if patch_gen is not None:84 return patch_gen.num_cls_tokens85 elif getattr(self.model, 'global_pool', None) == 'avg':86 return 087 return 188 89 @property90 def patch_size(self) -> int:91 if self._patch_size is not None:92 return self._patch_size93 if hasattr(self.model, "patch_size"):94 return self.model.patch_size95 patch_gen = getattr(self.model, "patch_generator", None)96 if patch_gen is not None:97 return patch_gen.patch_size98 return None99 100 @property101 def max_resolution(self) -> int:102 return self._max_resolution103 104 @property105 def preferred_resolution(self) -> Resolution:106 return self._preferred_resolution107 108 @property109 def window_size(self) -> int:110 return self._window_size111 112 @property113 def min_resolution_step(self) -> int:114 res = self.patch_size115 if self.window_size is not None:116 res *= self.window_size117 return res118 119 @property120 def blocks(self) -> Iterable[nn.Module]:121 blocks = getattr(self.model, 'blocks', None)122 if blocks is not None:123 return blocks124 return None125 126 @property127 def embed_dim(self) -> int:128 return self.model.embed_dim129 130 @property131 def summary_dim(self) -> int:132 embed_dim = self.embed_dim133 if self.summary_idxs is not None:134 embed_dim *= self.summary_idxs.shape[0]135 return embed_dim136 137 def make_preprocessor_external(self) -> Callable[[torch.Tensor], torch.Tensor]:138 ret = self.input_conditioner139 self.input_conditioner = nn.Identity()140 return ret141 142 def get_nearest_supported_resolution(self, height: int, width: int) -> Resolution:143 height = int(round(height / self.min_resolution_step) * self.min_resolution_step)144 width = int(round(width / self.min_resolution_step) * self.min_resolution_step)145 146 height = max(height, self.min_resolution_step)147 width = max(width, self.min_resolution_step)148 149 return Resolution(height=height, width=width)150 151 def switch_to_deploy(self):152 fn = getattr(self.model, 'switch_to_deploy', None)153 if fn is not None:154 fn()155 156 def cpe_video_mode(self, t: int):157 '''158 Context Manager.159 160 Puts the patch generator into video mode, with the specified number of temporal frames.161 In video mode, the expectation is that the input buffer is of shape `(B*T, C, H, W)`.162 Video mode means that the same position viewport will be used for every frame in the temporal sequence, while keeping163 distinct viewports for each video in the batch.164 165 Usage:166 with radio_model.cpe_video_mode(t=t):167 y = radio_model(x)168 '''169 return self.model.cpe_video_mode(t)170 171 def forward(self, x: torch.Tensor, feature_fmt: str = 'NLC') -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:172 '''173 Forward process for model.174 Args:175 x: Input tensor. Unless `make_preprocessor_external` has been called, then the dynamic range of `x` is expected to be `[0, 1]`,176 otherwise `x` is expected to be mean centered with unit standard deviation.177 feature_format: ['NLC', 'NCHW'] - The output format for the features.178 '''179 res_step = self.min_resolution_step180 if res_step is not None and (x.shape[-2] % res_step != 0 or x.shape[-1] % res_step != 0):181 raise ValueError('The input resolution must be a multiple of `self.min_resolution_step`. '182 '`self.get_nearest_supported_resolution(<height>, <width>) is provided as a convenience API. '183 f'Input: {x.shape[-2:]}, Nearest: {self.get_nearest_supported_resolution(*x.shape[-2:])}')184 185 x = self.input_conditioner(x)186 y = self.model.forward_features(x)187 ret = self._extract_final(x, y, feature_fmt=feature_fmt)188 return ret189 190 def _extract_final(self, x: torch.Tensor, y: torch.Tensor, feature_fmt: str = 'NLC'):191 if isinstance(self.model, VisionTransformer):192 patch_gen = getattr(self.model, "patch_generator", None)193 if patch_gen is not None:194 all_summary = y[:, : patch_gen.num_cls_tokens]195 if self.summary_idxs is not None:196 bb_summary = all_summary[:, self.summary_idxs]197 else:198 bb_summary = all_summary199 all_feat = y[:, patch_gen.num_skip :]200 elif self.model.global_pool == "avg":201 all_summary = y[:, self.model.num_prefix_tokens :].mean(dim=1)202 bb_summary = all_summary203 all_feat = y204 else:205 all_summary = y[:, 0]206 bb_summary = all_summary207 all_feat = y[:, 1:]208 elif isinstance(self.model, eradio_model.ERADIO):209 _, f = y210 all_feat = f.flatten(2).transpose(1, 2)211 all_summary = all_feat.mean(dim=1)212 bb_summary = all_summary213 elif isinstance(y, (list, tuple)):214 all_summary, all_feat = y215 bb_summary = all_summary216 else:217 all_summary = y[:, :self.num_cls_tokens]218 if self.summary_idxs is not None and all_summary.shape[1] > 1:219 if all_summary.shape[1] == 1:220 # Create dummy duplicates221 all_summary = all_summary.expand(-1, 128, -1)222 bb_summary = all_summary[:, self.summary_idxs]223 else:224 bb_summary = all_summary225 all_feat = y[:, self.num_summary_tokens:]226 227 all_feat = self.feature_normalizer(all_feat)228 229 if feature_fmt == 'NCHW':230 fmt_feat = (all_feat.reshape(all_feat.shape[0], x.shape[-2] // self.patch_size, x.shape[-1] // self.patch_size, all_feat.shape[2])231 .permute(0, 3, 1, 2)232 )233 elif feature_fmt == 'NLC':234 fmt_feat = all_feat235 else:236 raise ValueError(f'Unsupported feature_fmt: {feature_fmt}. Must be one of ["NLC", "NCHW"]')237 238 ret = RadioOutput(bb_summary.flatten(1), fmt_feat)239 240 if self.adaptors:241 ret = dict(backbone=ret)242 for name, adaptor in self.adaptors.items():243 if all_summary.ndim == 3:244 if all_summary.shape[1] == 1:245 summary = all_summary[:, 0]246 else:247 summary = all_summary[:, adaptor.head_idx]248 else:249 summary = all_summary250 ada_input = AdaptorInput(images=x, summary=summary.float(), features=all_feat, feature_fmt=feature_fmt, patch_size=self.patch_size)251 v = adaptor(ada_input).to(torch.float32)252 ret[name] = v253 254 return ret255 256 def forward_intermediates(257 self,258 x: torch.Tensor,259 indices: Optional[Union[int, List[int], Tuple[int]]] = None,260 return_prefix_tokens: bool = False,261 norm: bool = False,262 stop_early: bool = False,263 output_fmt: str = 'NCHW',264 intermediates_only: bool = False,265 aggregation: Optional[str] = "sparse",266 norm_alpha_scheme: Optional[str] = "post-alpha",267 ) -> List[RadioOutput]:268 """ Forward features that returns intermediates.269 Args:270 x: Input image tensor271 indices: Take last n blocks if int, select matching indices if sequence272 return_prefix_tokens: Return both prefix and spatial intermediate tokens273 norm: Apply norm layer to all intermediates274 stop_early: Stop iterating over blocks when last desired intermediate hit275 output_fmt: Shape of intermediate feature outputs. Options: NCHW, NLC276 intermediates_only: Only return intermediate features277 aggregation: intermediate layer aggregation method (sparse or dense).278 Dense accumulation is done by averaging the features in each group.279 norm_alpha_scheme: apply alpha before ("pre-alpha") or after accumulation ("post-alpha"), or don't normalize ("none")280 Only affects dense aggregation281 Returns:282 List of RadioOutput objects.283 """284 x = self.input_conditioner(x)285 intermediates = self.model.forward_intermediates(286 x,287 indices=indices,288 return_prefix_tokens=return_prefix_tokens,289 norm=norm,290 stop_early=stop_early,291 output_fmt=output_fmt,292 intermediates_only=intermediates_only,293 aggregation=aggregation,294 inter_feature_normalizer=self.inter_feature_normalizer,295 norm_alpha_scheme=norm_alpha_scheme,296 )297 298 if not intermediates_only:299 final, intermediates = intermediates300 301 def prepare_summary(summ: Optional[torch.Tensor]):302 if summ is None:303 return summ304 if self.summary_idxs is not None and summ.shape[1] > 1:305 summ = summ[:, self.summary_idxs]306 return summ.flatten(1)307 308 if return_prefix_tokens:309 radio_outputs = [310 RadioOutput(prepare_summary(summary), features)311 for summary, features in intermediates312 ]313 else:314 radio_outputs = intermediates315 316 if intermediates_only:317 return radio_outputs318 else:319 final = self._extract_final(x, final, feature_fmt=output_fmt)320 return final, radio_outputs321 322 323def create_model_from_args(args) -> nn.Module:324 in_chans = 3325 if args.in_chans is not None:326 in_chans = args.in_chans327 elif args.input_size is not None:328 in_chans = args.input_size[0]329 330 # Skip weight initialization unless it's explicitly requested.331 weight_init = args.model_kwargs.pop("weight_init", "skip")332 333 model = create_model(334 args.model,335 pretrained=args.pretrained,336 in_chans=in_chans,337 num_classes=args.num_classes,338 drop_rate=args.drop,339 drop_path_rate=args.drop_path,340 drop_block_rate=args.drop_block,341 global_pool=args.gp,342 bn_momentum=args.bn_momentum,343 bn_eps=args.bn_eps,344 scriptable=args.torchscript,345 checkpoint_path=args.initial_checkpoint,346 weight_init=weight_init,347 **args.model_kwargs,348 )349 350 if hasattr(model, 'norm') and not getattr(args, 'model_norm', False):351 model.norm = nn.Identity()352 353 model.head = nn.Identity()354 355 if args.cpe_max_size is not None:356 uq_teachers = set(t['name'] for t in args.teachers)357 enable_cpe(358 model,359 args.cpe_max_size,360 num_cls_tokens=len(uq_teachers) if args.cls_token_per_teacher else 1,361 register_multiple=getattr(args, 'register_multiple', None),362 num_registers=getattr(args, 'cpe_num_registers', None),363 )364 365 return model366 