nvidia/C-RADIOv4-H
8527k
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 9import math10from typing import Union, Tuple, Optional11 12import torch13import torch.nn.functional as F14from torch import nn15from einops import rearrange16 17from .cls_token import ClsToken18 19input_dim_t = Union[int, Tuple[int, int]]20 21try:22 # raise ImportError()23 from indirect_grid_sample import indirect_grid_sample24except ImportError:25 indirect_grid_sample = None26 27class ViTPatchGenerator(nn.Module):28 def __init__(self,29 patch_size: int,30 embed_dim: int,31 input_dims: input_dim_t,32 abs_pos: bool = True,33 normalize_patches: bool = False,34 cls_token: bool = False,35 max_input_dims: Optional[input_dim_t] = None,36 pos_dropout: float = 0.0,37 return_pos_enc: bool = False,38 num_cls_tokens: int = 1,39 register_multiple: Optional[int] = None,40 num_registers: Optional[int] = None,41 patch_bias: bool = False,42 device=None, dtype=None,43 ):44 super().__init__()45 46 if isinstance(input_dims, int):47 input_dims = (input_dims, input_dims)48 49 if max_input_dims is None:50 max_input_dims = input_dims51 if isinstance(max_input_dims, int):52 max_input_dims = (max_input_dims, max_input_dims)53 54 max_input_dims = tuple(55 int(math.ceil(d / patch_size) * patch_size)56 for d in max_input_dims57 )58 59 self.cpe_mode = max_input_dims != input_dims60 self.pos_dropout = pos_dropout61 self.return_pos_enc = return_pos_enc62 63 factory = dict(device=device, dtype=dtype)64 65 self.patch_size = patch_size66 self.abs_pos = abs_pos67 self.embed_dim = embed_dim68 69 self.num_rows = max_input_dims[0] // patch_size70 self.num_cols = max_input_dims[1] // patch_size71 self.input_dims = tuple(d // patch_size for d in input_dims)72 self.num_patches = self.num_rows * self.num_cols73 self.max_input_dims = max_input_dims74 75 self.im_to_patches = Im2Patches(patch_size)76 self.embedder = ViTPatchLinear(patch_size, embed_dim, bias=patch_bias, **factory)77 78 if abs_pos:79 scale = embed_dim ** -0.580 self.pos_embed = nn.Parameter(torch.randn(1, self.num_patches, embed_dim, **factory) * scale)81 82 self.cls_token = ClsToken(83 embed_dim,84 num_tokens=num_cls_tokens,85 enabled=cls_token,86 register_multiple=register_multiple,87 num_registers=num_registers,88 )89 90 self.patch_normalizer = nn.LayerNorm(embed_dim) if normalize_patches else nn.Identity()91 92 self.num_video_frames = None93 94 def forward(self, x: torch.Tensor) -> torch.Tensor:95 patches = self.embed_patches(x)96 patches, pos_enc = self.apply_pos_enc(patches, input_size=x.shape[2:])97 patches = self.cls_token(patches)98 patches = self.patch_normalizer(patches)99 if self.return_pos_enc:100 return patches, pos_enc101 return patches102 103 @property104 def apply_cls_token(self):105 return self.cls_token.enabled106 107 @property108 def num_cls_tokens(self):109 return self.cls_token.num_tokens110 111 @property112 def num_cls_patches(self):113 return self.cls_token.num_patches114 115 @property116 def num_registers(self):117 return self.cls_token.num_registers118 119 @property120 def num_skip(self):121 return self.num_cls_tokens + self.num_registers122 123 def no_weight_decay(self):124 return [125 'pos_embed',126 ]127 128 def _load_embed(self, src_embed: torch.Tensor, targ_embed: nn.Parameter):129 if src_embed.shape != targ_embed.shape:130 src_size = int(math.sqrt(src_embed.shape[1]))131 132 assert src_size ** 2 == src_embed.shape[1], 'Unable to interpolate non-square embedding'133 134 src_embed = rearrange(src_embed, 'b (h w) c -> b c h w', h=src_size, w=src_size)135 src_embed = F.interpolate(src_embed, size=(self.num_rows, self.num_cols), mode='bicubic', align_corners=True, antialias=False)136 src_embed = rearrange(src_embed, 'b c h w -> b (h w) c')137 targ_embed.data.copy_(src_embed)138 139 def _load_projection(self, src_proj_weight: torch.Tensor, targ_proj_weight: torch.Tensor):140 if src_proj_weight.shape != targ_proj_weight.shape:141 src_patch_size = int(math.sqrt(src_proj_weight.shape[1] // 3))142 143 assert (src_patch_size ** 2) * 3 == src_proj_weight.shape[1], 'Unable to interpolate non-square patch size'144 145 src_proj_weight = rearrange(src_proj_weight, 'b (c h w) -> b c h w', c=3, h=src_patch_size, w=src_patch_size)146 src_proj_weight = F.interpolate(src_proj_weight, size=(self.patch_size, self.patch_size), mode='bicubic', align_corners=True, antialias=False)147 src_proj_weight = rearrange(src_proj_weight, 'b c h w -> b (c h w)')148 targ_proj_weight.data.copy_(src_proj_weight)149 150 def embed_patches(self, x: torch.Tensor) -> torch.Tensor:151 patches = self.im_to_patches(x)152 patches = self.embedder(patches)153 return patches154 155 def apply_pos_enc(self,156 patches: torch.Tensor,157 patch_idxs: Optional[torch.Tensor] = None,158 input_size: Optional[Tuple[int, int]] = None,159 ) -> torch.Tensor:160 if not self.abs_pos:161 return patches162 163 pos_enc = self.get_pos_enc(patches.shape[0], patch_idxs, input_size)164 165 if self.training and self.pos_dropout > 0:166 keeps = torch.rand(patches.shape[0], 1, 1, dtype=pos_enc.dtype, device=pos_enc.device) > self.pos_dropout167 pos_enc_drop = torch.where(keeps, pos_enc, 0)168 else:169 pos_enc_drop = pos_enc170 171 return patches + pos_enc_drop, pos_enc172 173 def get_pos_enc(self,174 batch_size: int,175 patch_idxs: Optional[torch.Tensor] = None,176 input_size: Optional[Tuple[int, int]] = None,177 ) -> torch.Tensor:178 if input_size is None:179 input_dims = self.input_dims180 else:181 input_dims = tuple(d // self.patch_size for d in input_size)182 183 pos_embed = self._get_pos_embeddings(batch_size, input_dims)184 185 if patch_idxs is None:186 return pos_embed187 188 exp_patch_idxs = patch_idxs.unsqueeze(-1).expand(-1, -1, pos_embed.shape[-1])189 190 pos_embed = torch.gather(pos_embed.expand(patch_idxs.shape[0], -1, -1), dim=1, index=exp_patch_idxs)191 return pos_embed192 193 194 def _get_pos_embeddings(self, batch_size: int, input_dims: Tuple[int, int]):195 if (self.num_rows, self.num_cols) == input_dims:196 return self.pos_embed197 198 pos_embed = self.pos_embed.reshape(1, self.num_rows, self.num_cols, -1).permute(0, 3, 1, 2)199 200 def window_select(pos_embed):201 if input_dims[0] < pos_embed.shape[-2]:202 pos_embed = pos_embed[..., :input_dims[0], :]203 if input_dims[1] < pos_embed.shape[-1]:204 pos_embed = pos_embed[..., :, :input_dims[1]]205 return pos_embed206 207 if self.cpe_mode:208 if self.training:209 if self.num_video_frames is not None:210 if batch_size % self.num_video_frames != 0:211 raise ValueError(f'Batch size {batch_size} must be divisible by num_video_frames {self.num_video_frames} for CPE mode.')212 213 batch_size //= self.num_video_frames214 215 min_scale = math.sqrt(0.1)216 scale = torch.rand(batch_size, 1, 1, device=pos_embed.device) * (1 - min_scale) + min_scale217 aspect_min = math.log(3 / 4)218 aspect_max = -aspect_min219 aspect = torch.exp(torch.rand(batch_size, 1, 1, device=pos_embed.device) * (aspect_max - aspect_min) + aspect_min)220 221 scale_x = scale * aspect222 scale_y = scale * (1 / aspect)223 scale_xy = torch.stack([scale_x, scale_y], dim=-1).clamp_(0, 1)224 225 pos_xy = torch.rand(batch_size, 1, 1, 2, device=pos_embed.device) * (1 - scale_xy)226 227 lin_x = torch.linspace(0, 1, steps=input_dims[1], device=pos_embed.device)[None, None].expand(batch_size, input_dims[0], -1)228 lin_y = torch.linspace(0, 1, steps=input_dims[0], device=pos_embed.device)[None, :, None].expand(batch_size, -1, input_dims[1])229 230 lin_xy = torch.stack([lin_x, lin_y], dim=-1)231 232 grid_xy = lin_xy * scale_xy + pos_xy233 234 # Convert to [-1, 1] range235 grid_xy.mul_(2).sub_(1)236 237 pos_embed = F.grid_sample(238 pos_embed.float().expand(batch_size, -1, -1, -1),239 grid=grid_xy,240 mode='bilinear',241 padding_mode='zeros',242 align_corners=True,243 ).to(pos_embed.dtype)244 245 if self.num_video_frames is not None:246 pos_embed = torch.repeat_interleave(pos_embed, self.num_video_frames, dim=0)247 else:248 # i_rows, i_cols = input_dims249 # p_rows, p_cols = pos_embed.shape[2:]250 # if i_rows <= p_rows and i_cols <= p_cols:251 # left = (p_cols - i_cols) // 2252 # top = (p_rows - i_rows) // 2253 # pos_embed = pos_embed[..., top:top+i_rows, left:left+i_cols]254 # else:255 max_dim = max(input_dims)256 pos_embed = F.interpolate(pos_embed.float(), size=(max_dim, max_dim), align_corners=False, mode='bilinear').to(pos_embed.dtype)257 258 pos_embed = window_select(pos_embed)259 else:260 pos_embed = window_select(pos_embed)261 262 if pos_embed.shape[-2:] != input_dims:263 pos_embed = F.interpolate(pos_embed.float(), size=input_dims, align_corners=False, mode='bilinear').to(pos_embed.dtype)264 265 pos_embed = pos_embed.flatten(2).permute(0, 2, 1)266 267 return pos_embed268 269 270class Im2Patches(nn.Module):271 def __init__(self, patch_size: int):272 super().__init__()273 self.patch_size = patch_size274 275 def forward(self, x: torch.Tensor) -> torch.Tensor:276 if self.patch_size == 1:277 patches = x.flatten(2)278 patches = patches.permute(0, 2, 1)279 return patches280 281 py = x.shape[-2] // self.patch_size282 px = x.shape[-1] // self.patch_size283 patches = rearrange(x, 'b c (py yy) (px xx) -> b (py px) (c yy xx)',284 py=py, yy=self.patch_size,285 px=px, xx=self.patch_size,286 )287 return patches288 289 290class ViTPatchLinear(nn.Linear):291 def __init__(self, patch_size: int, embed_dim: int, bias: bool = False, **factory):292 super().__init__(293 3 * (patch_size ** 2),294 embed_dim,295 bias=bias,296 **factory297 )298 self.patch_size = patch_size299 