CoolFace
Modelpublic

nvidia/C-RADIO

sourceHugging Faceotherupdated 2y agoView on Hugging Face
30likes12kdownloads
vit_patch_generator.py300 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.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: int = 0,40                 device=None, dtype=None,41    ):42        super().__init__()43 44        if isinstance(input_dims, int):45            input_dims = (input_dims, input_dims)46 47        if max_input_dims is None:48            max_input_dims = input_dims49        if isinstance(max_input_dims, int):50            max_input_dims = (max_input_dims, max_input_dims)51 52        max_input_dims = tuple(53            int(math.ceil(d / patch_size) * patch_size)54            for d in max_input_dims55        )56 57        self.cpe_mode = max_input_dims != input_dims58        self.pos_dropout = pos_dropout59        self.return_pos_enc = return_pos_enc60 61        factory = dict(device=device, dtype=dtype)62 63        self.patch_size = patch_size64        self.abs_pos = abs_pos65        self.embed_dim = embed_dim66 67        self.num_rows = max_input_dims[0] // patch_size68        self.num_cols = max_input_dims[1] // patch_size69        self.input_dims = tuple(d // patch_size for d in input_dims)70        self.num_patches = self.num_rows * self.num_cols71        self.max_input_dims = max_input_dims72 73        self.im_to_patches = Im2Patches(patch_size)74        self.embedder = ViTPatchLinear(patch_size, embed_dim, **factory)75 76        if abs_pos:77            scale = embed_dim ** -0.578            self.pos_embed = nn.Parameter(torch.randn(1, self.num_patches, embed_dim, **factory) * scale)79 80        self.cls_token = ClsToken(81            embed_dim,82            num_tokens=num_cls_tokens,83            enabled=cls_token,84            register_multiple=register_multiple,85        )86 87        self.patch_normalizer = nn.LayerNorm(embed_dim) if normalize_patches else nn.Identity()88 89    def forward(self, x: torch.Tensor) -> torch.Tensor:90        patches = self.embed_patches(x)91        patches, pos_enc = self.apply_pos_enc(patches, input_size=x.shape[2:])92        patches = self.cls_token(patches)93        patches = self.patch_normalizer(patches)94        if self.return_pos_enc:95            return patches, pos_enc96        return patches97 98    @property99    def apply_cls_token(self):100        return self.cls_token.enabled101 102    @property103    def num_cls_tokens(self):104        return self.cls_token.num_tokens105 106    @property107    def num_registers(self):108        return self.cls_token.num_registers109 110    @property111    def num_skip(self):112        return self.num_cls_tokens + self.num_registers113 114    def no_weight_decay(self):115        return [116            'pos_embed',117        ]118 119    def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs):120        if self.abs_pos:121            self._load_embed(state_dict[f'{prefix}pos_embed'], self.pos_embed)122 123    def _load_embed(self, src_embed: torch.Tensor, targ_embed: nn.Parameter):124        if src_embed.shape != targ_embed.shape:125            src_size = int(math.sqrt(src_embed.shape[1]))126 127            assert src_size ** 2 == src_embed.shape[1], 'Unable to interpolate non-square embedding'128 129            src_embed = rearrange(src_embed, 'b (h w) c -> b c h w', h=src_size, w=src_size)130            src_embed = F.interpolate(src_embed, size=(self.num_rows, self.num_cols), mode='bicubic', align_corners=True, antialias=False)131            src_embed = rearrange(src_embed, 'b c h w -> b (h w) c')132        targ_embed.data.copy_(src_embed)133 134    def _load_projection(self, src_proj_weight: torch.Tensor, targ_proj_weight: torch.Tensor):135        if src_proj_weight.shape != targ_proj_weight.shape:136            src_patch_size = int(math.sqrt(src_proj_weight.shape[1] // 3))137 138            assert (src_patch_size ** 2) * 3 == src_proj_weight.shape[1], 'Unable to interpolate non-square patch size'139 140            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)141            src_proj_weight = F.interpolate(src_proj_weight, size=(self.patch_size, self.patch_size), mode='bicubic', align_corners=True, antialias=False)142            src_proj_weight = rearrange(src_proj_weight, 'b c h w -> b (c h w)')143        targ_proj_weight.data.copy_(src_proj_weight)144 145    def embed_patches(self, x: torch.Tensor) -> torch.Tensor:146        patches = self.im_to_patches(x)147        patches = self.embedder(patches)148        return patches149 150    def apply_pos_enc(self,151                      patches: torch.Tensor,152                      patch_idxs: Optional[torch.Tensor] = None,153                      input_size: Optional[Tuple[int, int]] = None,154    ) -> torch.Tensor:155        if not self.abs_pos:156            return patches157 158        pos_enc = self.get_pos_enc(patches.shape[0], patch_idxs, input_size)159 160        if self.training and self.pos_dropout > 0:161            keeps = torch.rand(patches.shape[0], 1, 1, dtype=pos_enc.dtype, device=pos_enc.device) > self.pos_dropout162            pos_enc_drop = torch.where(keeps, pos_enc, 0)163        else:164            pos_enc_drop = pos_enc165 166        return patches + pos_enc_drop, pos_enc167 168    def get_pos_enc(self,169                    batch_size: int,170                    patch_idxs: Optional[torch.Tensor] = None,171                    input_size: Optional[Tuple[int, int]] = None,172    ) -> torch.Tensor:173        if input_size is None:174            input_dims = self.input_dims175        else:176            input_dims = tuple(d // self.patch_size for d in input_size)177 178        pos_embed = self._get_pos_embeddings(batch_size, input_dims)179 180        if patch_idxs is None:181            return pos_embed182 183        exp_patch_idxs = patch_idxs.unsqueeze(-1).expand(-1, -1, pos_embed.shape[-1])184 185        pos_embed = torch.gather(pos_embed.expand(patch_idxs.shape[0], -1, -1), dim=1, index=exp_patch_idxs)186        return pos_embed187 188 189    def _get_pos_embeddings(self, batch_size: int, input_dims: Tuple[int, int]):190        if (self.num_rows, self.num_cols) == input_dims:191            return self.pos_embed192 193        pos_embed = self.pos_embed.reshape(1, self.num_rows, self.num_cols, -1).permute(0, 3, 1, 2)194 195        def window_select(pos_embed):196            if input_dims[0] < pos_embed.shape[-2]:197                pos_embed = pos_embed[..., :input_dims[0], :]198            if input_dims[1] < pos_embed.shape[-1]:199                pos_embed = pos_embed[..., :, :input_dims[1]]200            return pos_embed201 202        if self.cpe_mode:203            if self.training:204                min_scale = math.sqrt(0.1)205                scale = torch.rand(batch_size, 1, 1, device=pos_embed.device) * (1 - min_scale) + min_scale206                aspect_min = math.log(3 / 4)207                aspect_max = -aspect_min208                aspect = torch.exp(torch.rand(batch_size, 1, 1, device=pos_embed.device) * (aspect_max - aspect_min) + aspect_min)209 210                scale_x = scale * aspect211                scale_y = scale * (1 / aspect)212                scale_xy = torch.stack([scale_x, scale_y], dim=-1).clamp_(0, 1)213 214                pos_xy = torch.rand(batch_size, 1, 1, 2, device=pos_embed.device) * (1 - scale_xy)215 216                lin_x = torch.linspace(0, 1, steps=input_dims[1], device=pos_embed.device)[None, None].expand(batch_size, input_dims[0], -1)217                lin_y = torch.linspace(0, 1, steps=input_dims[0], device=pos_embed.device)[None, :, None].expand(batch_size, -1, input_dims[1])218 219                lin_xy = torch.stack([lin_x, lin_y], dim=-1)220 221                grid_xy = lin_xy * scale_xy + pos_xy222 223                # Convert to [-1, 1] range224                grid_xy.mul_(2).sub_(1)225 226                pos_embed = F.grid_sample(227                    pos_embed.float().expand(batch_size, -1, -1, -1),228                    grid=grid_xy,229                    mode='bilinear',230                    padding_mode='zeros',231                    align_corners=True,232                ).to(pos_embed.dtype)233            else:234                # i_rows, i_cols = input_dims235                # p_rows, p_cols = pos_embed.shape[2:]236                # if i_rows <= p_rows and i_cols <= p_cols:237                #     left = (p_cols - i_cols) // 2238                #     top = (p_rows - i_rows) // 2239                #     pos_embed = pos_embed[..., top:top+i_rows, left:left+i_cols]240                # else:241                max_dim = max(input_dims)242                pos_embed = F.interpolate(pos_embed.float(), size=(max_dim, max_dim), align_corners=True, mode='bilinear').to(pos_embed.dtype)243 244                pos_embed = window_select(pos_embed)245        else:246            pos_embed = window_select(pos_embed)247 248        if pos_embed.shape[-2:] != input_dims:249            pos_embed = F.interpolate(pos_embed.float(), size=input_dims, align_corners=True, mode='bilinear').to(pos_embed.dtype)250 251        pos_embed = pos_embed.flatten(2).permute(0, 2, 1)252 253        return pos_embed254 255 256class Im2Patches(nn.Module):257    def __init__(self, patch_size: int):258        super().__init__()259        self.patch_size = patch_size260 261    def forward(self, x: torch.Tensor) -> torch.Tensor:262        if self.patch_size == 1:263            patches = x.flatten(2)264            patches = patches.permute(0, 2, 1)265            return patches266 267        py = x.shape[-2] // self.patch_size268        px = x.shape[-1] // self.patch_size269        patches = rearrange(x, 'b c (py yy) (px xx) -> b (py px) (c yy xx)',270                            py=py, yy=self.patch_size,271                            px=px, xx=self.patch_size,272        )273        return patches274 275 276class ViTPatchLinear(nn.Linear):277    def __init__(self, patch_size: int, embed_dim: int, **factory):278        super().__init__(279            3 * (patch_size ** 2),280            embed_dim,281            bias=False,282            **factory283        )284        self.patch_size = patch_size285 286    def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs):287        if self.bias is not None:288            self.bias.data.copy_(state_dict[f'{prefix}bias'])289 290        chk_weight = state_dict[f'{prefix}weight']291        if chk_weight.shape != self.weight.shape:292            src_patch_size = int(math.sqrt(chk_weight.shape[1] // 3))293 294            assert (src_patch_size ** 2) * 3 == chk_weight.shape[1], 'Unable to interpolate non-square patch size'295 296            chk_weight = rearrange(chk_weight, 'b (c h w) -> b c h w', c=3, h=src_patch_size, w=src_patch_size)297            chk_weight = F.interpolate(chk_weight, size=(self.patch_size, self.patch_size), mode='bicubic', align_corners=True, antialias=False)298            chk_weight = rearrange(chk_weight, 'b c h w -> b (c h w)')299        self.weight.data.copy_(chk_weight)300