nvidia/C-RADIO
3012k
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 typing import Union, Tuple10from types import MethodType11 12import torch13from torch import nn14 15from timm.models import VisionTransformer, checkpoint_seq16 17from .vit_patch_generator import ViTPatchGenerator18 19 20def _forward_cpe(self: VisionTransformer, x: torch.Tensor) -> torch.Tensor:21 x = self.patch_generator(x)22 if self.grad_checkpointing and not torch.jit.is_scripting():23 x = checkpoint_seq(self.blocks, x)24 else:25 x = self.blocks(x)26 x = self.norm(x)27 return x28 29 30def enable_cpe(model: nn.Module,31 max_img_size: Union[int, Tuple[int, int]] = 1024,32 num_cls_tokens: int = 1,33 pos_dropout: float = 0.1,34 register_multiple: int = 0,35):36 if not isinstance(model, VisionTransformer):37 raise ValueError("CPE only support for VisionTransformer models!")38 39 patch_size = model.patch_embed.patch_size[0]40 embed_dim = model.embed_dim41 input_dims = model.patch_embed.img_size42 normalize_patches = not isinstance(model.patch_embed.norm, nn.Identity)43 cls_token = model.cls_token is not None44 45 max_img_size = int(round(max_img_size / patch_size) * patch_size)46 47 patch_generator = ViTPatchGenerator(48 patch_size=patch_size,49 embed_dim=embed_dim,50 input_dims=input_dims,51 normalize_patches=normalize_patches,52 cls_token=cls_token,53 max_input_dims=max_img_size,54 pos_dropout=pos_dropout,55 num_cls_tokens=num_cls_tokens,56 register_multiple=register_multiple,57 )58 59 model.patch_generator = patch_generator60 model.patch_embed = None61 model.cls_token = None62 model.pos_embed = None63 model.pos_drop = None64 model.num_cls_tokens = num_cls_tokens65 model.num_registers = patch_generator.num_registers66 67 model.forward_features = MethodType(_forward_cpe, model)68 