CoolFace
Modelpublic

nvidia/C-RADIOv4-H

sourceHugging Faceotherupdated 8mo agoView on Hugging Face
84likes30kdownloads
radio_model.py366 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.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