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