Team Ai
Modelpublic

nvidia/C-RADIOv4-H

sourceHugging Faceotherupdated 8mo agoView on Hugging Face
85likes27kdownloads
hf_model.py209 linesDownload Raw Back to root
1# Copyright (c) 2023-2024, NVIDIA CORPORATION.  All rights reserved.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7#     http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14from collections import namedtuple15from typing import Callable, Dict, Optional, List, Union16 17from timm.models import VisionTransformer18import torch19from torch import nn20from transformers import PretrainedConfig, PreTrainedModel21 22 23from .common import RESOURCE_MAP, DEFAULT_VERSION24 25# Import all required modules.26from .adaptor_base import AdaptorBase, RadioOutput, AdaptorInput27from .adaptor_generic import GenericAdaptor, AdaptorBase28from .adaptor_module_factory import create_mlp_from_config29from .adaptor_mlp import MLP, MLP230from .adaptor_attn import AttnFDHead31from .adaptor_registry import adaptor_registry32from .cls_token import ClsToken33from .dinov2_arch import dinov2_vitg14_reg34from .enable_cpe_support import enable_cpe35from .enable_damp import configure_damp_from_args36from .enable_spectral_reparam import configure_spectral_reparam_from_args37from .eradio_model import eradio38from .feature_normalizer import FeatureNormalizer, IntermediateFeatureNormalizer39from .forward_intermediates import forward_intermediates40from .radio_model import create_model_from_args41from .radio_model import RADIOModel as RADIOModelBase, Resolution42from .input_conditioner import get_default_conditioner, InputConditioner43from .open_clip_adaptor import OpenCLIP_RADIO44from .siglip2_adaptor import SigLIP2Adaptor45from .vit_patch_generator import ViTPatchGenerator46from .vitdet import apply_vitdet_arch, VitDetArgs47 48# Register extra models49from .extra_timm_models import *50from .extra_models import *51 52 53class RADIOConfig(PretrainedConfig):54    """Pretrained Hugging Face configuration for RADIO models."""55 56    def __init__(57        self,58        args: Optional[dict] = None,59        version: Optional[str] = DEFAULT_VERSION,60        patch_size: Optional[int] = None,61        max_resolution: Optional[int] = None,62        preferred_resolution: Optional[Resolution] = None,63        adaptor_names: Union[str, List[str]] = None,64        adaptor_configs: Dict[str, Dict[str, int]] = None,65        vitdet_window_size: Optional[int] = None,66        feature_normalizer_config: Optional[dict] = None,67        inter_feature_normalizer_config: Optional[dict] = None,68        **kwargs,69    ):70        self.args = args71        for field in ["dtype", "amp_dtype"]:72            if self.args is not None and field in self.args:73                # Convert to a string in order to make it serializable.74                # For example for torch.float32 we will store "float32",75                # for "bfloat16" we will store "bfloat16".76                self.args[field] = str(args[field]).split(".")[-1]77        self.version = version78        resource = RESOURCE_MAP[version]79        self.patch_size = patch_size or resource.patch_size80        self.max_resolution = max_resolution or resource.max_resolution81        self.preferred_resolution = (82            preferred_resolution or resource.preferred_resolution83        )84        self.adaptor_names = adaptor_names85        self.adaptor_configs = adaptor_configs86        self.vitdet_window_size = vitdet_window_size87        self.feature_normalizer_config = feature_normalizer_config88        self.inter_feature_normalizer_config = inter_feature_normalizer_config89        super().__init__(**kwargs)90 91 92 93class RADIOModel(PreTrainedModel):94    """Pretrained Hugging Face model for RADIO.95 96    This class inherits from PreTrainedModel, which provides97    HuggingFace's functionality for loading and saving models.98    """99 100    config_class = RADIOConfig101 102    def __init__(self, config: RADIOConfig):103        super().__init__(config)104        if hasattr(super(), "post_init"):105            super().post_init()106 107        RADIOArgs = namedtuple("RADIOArgs", config.args.keys())108        args = RADIOArgs(**config.args)109        self.config = config110 111        model = create_model_from_args(args)112        input_conditioner: InputConditioner = get_default_conditioner()113 114        dtype = getattr(args, "dtype", torch.float32)115        if isinstance(dtype, str):116            # Convert the dtype's string representation back to a dtype.117            dtype = getattr(torch, dtype)118        model.to(dtype=dtype)119        input_conditioner.dtype = dtype120 121        summary_idxs = torch.tensor(122            [i for i, t in enumerate(args.teachers) if t.get("use_summary", True)],123            dtype=torch.int64,124        )125 126        adaptor_configs = config.adaptor_configs127        adaptor_names = config.adaptor_names or []128 129        adaptors = dict()130        for adaptor_name in adaptor_names:131            mlp_config = adaptor_configs[adaptor_name]132            adaptor = GenericAdaptor(args, None, None, mlp_config)133            adaptor.head_idx = mlp_config["head_idx"]134            adaptors[adaptor_name] = adaptor135 136        feature_normalizer = None137        if config.feature_normalizer_config is not None:138            # Actual normalization values will be restored when loading checkpoint weights.139            feature_normalizer = FeatureNormalizer(config.feature_normalizer_config["embed_dim"])140 141        inter_feature_normalizer = None142        if config.inter_feature_normalizer_config is not None:143            inter_feature_normalizer = IntermediateFeatureNormalizer(144                config.inter_feature_normalizer_config["num_intermediates"],145                config.inter_feature_normalizer_config["embed_dim"],146                rot_per_layer=config.inter_feature_normalizer_config["rot_per_layer"],147                dtype=dtype)148 149        self.radio_model = RADIOModelBase(150            model,151            input_conditioner,152            summary_idxs=summary_idxs,153            patch_size=config.patch_size,154            max_resolution=config.max_resolution,155            window_size=config.vitdet_window_size,156            preferred_resolution=config.preferred_resolution,157            adaptors=adaptors,158            feature_normalizer=feature_normalizer,159            inter_feature_normalizer=inter_feature_normalizer,160        )161 162    @property163    def adaptors(self) -> nn.ModuleDict:164        return self.radio_model.adaptors165 166    @property167    def model(self) -> VisionTransformer:168        return self.radio_model.model169 170    @property171    def input_conditioner(self) -> InputConditioner:172        return self.radio_model.input_conditioner173 174    @property175    def num_summary_tokens(self) -> int:176        return self.radio_model.num_summary_tokens177 178    @property179    def patch_size(self) -> int:180        return self.radio_model.patch_size181 182    @property183    def max_resolution(self) -> int:184        return self.radio_model.max_resolution185 186    @property187    def preferred_resolution(self) -> Resolution:188        return self.radio_model.preferred_resolution189 190    @property191    def window_size(self) -> int:192        return self.radio_model.window_size193 194    @property195    def min_resolution_step(self) -> int:196        return self.radio_model.min_resolution_step197 198    def make_preprocessor_external(self) -> Callable[[torch.Tensor], torch.Tensor]:199        return self.radio_model.make_preprocessor_external()200 201    def get_nearest_supported_resolution(self, height: int, width: int) -> Resolution:202        return self.radio_model.get_nearest_supported_resolution(height, width)203 204    def switch_to_deploy(self):205        return self.radio_model.switch_to_deploy()206 207    def forward(self, x: torch.Tensor):208        return self.radio_model.forward(x)209