Team Ai
Modelpublic

recursionpharma/OpenPhenom

sourceHugging Faceupdated 7mo agoView on Hugging Face
22likes1.2kdownloads
vit.py310 linesDownload Raw Back to root
1# © Recursion Pharmaceuticals 20242import timm.models.vision_transformer as vit3import torch4 5 6def generate_2d_sincos_pos_embeddings(7    embedding_dim: int,8    length: int,9    scale: float = 10000.0,10    use_class_token: bool = True,11    num_modality: int = 1,12) -> torch.nn.Parameter:13    """14    Generate 2Dimensional sin/cosine positional embeddings15 16    Parameters17    ----------18    embedding_dim : int19        embedding dimension used in vit20    length : int21        number of tokens along height or width of image after patching (assuming square)22    scale : float23        scale for sin/cos functions24    use_class_token : bool25        True - add zero vector to be added to class_token, False - no vector added26    num_modality: number of modalities. If 0, a single modality is assumed.27        Otherwise one-hot modality encoding is added and sincos encoding size is appropriately reduced.28 29    Returns30    -------31    positional_encoding : torch.Tensor32        positional encoding to add to vit patch encodings33        [num_modality*length*length, embedding_dim] or [1+num_modality*length*length, embedding_dim]34        (w/ or w/o cls_token)35    """36 37    linear_positions = torch.arange(length, dtype=torch.float32)38    height_mesh, width_mesh = torch.meshgrid(39        linear_positions, linear_positions, indexing="ij"40    )41    positional_dim = embedding_dim // 4  # accomodate h and w x cos and sin embeddings42    positional_weights = (43        torch.arange(positional_dim, dtype=torch.float32) / positional_dim44    )45    positional_weights = 1.0 / (scale**positional_weights)46 47    height_weights = torch.outer(height_mesh.flatten(), positional_weights)48    width_weights = torch.outer(width_mesh.flatten(), positional_weights)49 50    positional_encoding = torch.cat(51        [52            torch.sin(height_weights),53            torch.cos(height_weights),54            torch.sin(width_weights),55            torch.cos(width_weights),56        ],57        dim=1,58    )[None, :, :]59 60    # repeat positional encoding for multiple channel modalities61    positional_encoding = positional_encoding.repeat(1, num_modality, 1)62 63    if use_class_token:64        class_token = torch.zeros([1, 1, embedding_dim], dtype=torch.float32)65        positional_encoding = torch.cat([class_token, positional_encoding], dim=1)66 67    positional_encoding = torch.nn.Parameter(positional_encoding, requires_grad=False)68 69    return positional_encoding70 71 72class ChannelAgnosticPatchEmbed(vit.PatchEmbed):  # type: ignore[misc]73    def __init__(74        self,75        img_size: int,76        patch_size: int,77        embed_dim: int,78        bias: bool = True,79    ) -> None:80        super().__init__(81            img_size=img_size,82            patch_size=patch_size,83            in_chans=1,  # in_chans is used by self.proj, which we override anyway84            embed_dim=embed_dim,85            norm_layer=None,86            flatten=False,87            bias=bias,88        )89        # channel-agnostic MAE has a single projection for all chans90        self.proj = torch.nn.Conv2d(91            1, embed_dim, kernel_size=patch_size, stride=patch_size, bias=bias92        )93 94    def forward(self, x: torch.Tensor) -> torch.Tensor:95        in_chans = x.shape[1]96        x = torch.stack(97            [self.proj(x[:, i : i + 1]) for i in range(in_chans)], dim=298        )  # single project for all chans99        x = x.flatten(2).transpose(1, 2)  # BCMHW -> BNC100        return x101 102 103class ChannelAgnosticViT(vit.VisionTransformer):  # type: ignore[misc]104    def _pos_embed(self, x: torch.Tensor) -> torch.Tensor:105        # rewrite https://github.com/huggingface/pytorch-image-models/blob/main/timm/models/vision_transformer.py#L586106        to_cat = []107        if self.cls_token is not None:108            to_cat.append(self.cls_token.expand(x.shape[0], -1, -1))109 110        # TODO: upgrade timm to get access to register tokens111        # if self.vit_backbone.reg_token is not None:112        #     to_cat.append(self.reg_token.expand(x.shape[0], -1, -1))113 114        # MAIN DIFFERENCE with Timm - we DYNAMICALLY ADDING POS EMBEDDINGS based on shape of inputs115        # this supports having CA-MAEs actually be channel-agnostic at inference time116        if self.no_embed_class:117            x = x + self.pos_embed[:, : x.shape[1]]118            if to_cat:119                x = torch.cat(to_cat + [x], dim=1)120        else:121            if to_cat:122                x = torch.cat(to_cat + [x], dim=1)123            x = x + self.pos_embed[:, : x.shape[1]]124        return self.pos_drop(x)  # type: ignore[no-any-return]125 126 127def channel_agnostic_vit(128    vit_backbone: vit.VisionTransformer, max_in_chans: int129) -> vit.VisionTransformer:130    # replace patch embedding with channel-agnostic version131    vit_backbone.patch_embed = ChannelAgnosticPatchEmbed(132        img_size=vit_backbone.patch_embed.img_size[0],133        patch_size=vit_backbone.patch_embed.patch_size[0],134        embed_dim=vit_backbone.embed_dim,135    )136 137    # replace positional embedding with channel-agnostic version138    vit_backbone.pos_embed = generate_2d_sincos_pos_embeddings(139        embedding_dim=vit_backbone.embed_dim,140        length=vit_backbone.patch_embed.grid_size[0],141        use_class_token=vit_backbone.cls_token is not None,142        num_modality=max_in_chans,143    )144 145    # change the class to be ChannelAgnostic so that it actually uses the new _pos_embed146    vit_backbone.__class__ = ChannelAgnosticViT147    return vit_backbone148 149 150def sincos_positional_encoding_vit(151    vit_backbone: vit.VisionTransformer, scale: float = 10000.0152) -> vit.VisionTransformer:153    """Attaches no-grad sin-cos positional embeddings to a pre-constructed ViT backbone model.154 155    Parameters156    ----------157    vit_backbone : timm.models.vision_transformer.VisionTransformer158        the constructed vision transformer from timm159    scale : float (default 10000.0)160        hyperparameter for sincos positional embeddings, recommend keeping at 10,000161 162    Returns163    -------164    timm.models.vision_transformer.VisionTransformer165        the same ViT but with fixed no-grad positional encodings to add to vit patch encodings166    """167    # length: number of tokens along height or width of image after patching (assuming square)168    length = (169        vit_backbone.patch_embed.img_size[0] // vit_backbone.patch_embed.patch_size[0]170    )171    pos_embeddings = generate_2d_sincos_pos_embeddings(172        vit_backbone.embed_dim,173        length=length,174        scale=scale,175        use_class_token=vit_backbone.cls_token is not None,176    )177    # note, if the model had weight_init == 'skip', this might get overwritten178    vit_backbone.pos_embed = pos_embeddings179    return vit_backbone180 181 182def vit_small_patch16_256(**kwargs):183    default_kwargs = dict(184        img_size=256,185        in_chans=6,186        num_classes=0,187        fc_norm=None,188        class_token=True,189        drop_path_rate=0.1,190        init_values=0.0001,191        block_fn=vit.ParallelScalingBlock,192        qkv_bias=False,193        qk_norm=True,194    )195    for k, v in kwargs.items():196        default_kwargs[k] = v197    return vit.vit_small_patch16_224(**default_kwargs)198 199 200def vit_small_patch32_512(**kwargs):201    default_kwargs = dict(202        img_size=512,203        in_chans=6,204        num_classes=0,205        fc_norm=None,206        class_token=True,207        drop_path_rate=0.1,208        init_values=0.0001,209        block_fn=vit.ParallelScalingBlock,210        qkv_bias=False,211        qk_norm=True,212    )213    for k, v in kwargs.items():214        default_kwargs[k] = v215    return vit.vit_small_patch32_384(**default_kwargs)216 217 218def vit_base_patch8_256(**kwargs):219    default_kwargs = dict(220        img_size=256,221        in_chans=6,222        num_classes=0,223        fc_norm=None,224        class_token=True,225        drop_path_rate=0.1,226        init_values=0.0001,227        block_fn=vit.ParallelScalingBlock,228        qkv_bias=False,229        qk_norm=True,230    )231    for k, v in kwargs.items():232        default_kwargs[k] = v233    return vit.vit_base_patch8_224(**default_kwargs)234 235 236def vit_base_patch16_256(**kwargs):237    default_kwargs = dict(238        img_size=256,239        in_chans=6,240        num_classes=0,241        fc_norm=None,242        class_token=True,243        drop_path_rate=0.1,244        init_values=0.0001,245        block_fn=vit.ParallelScalingBlock,246        qkv_bias=False,247        qk_norm=True,248    )249    for k, v in kwargs.items():250        default_kwargs[k] = v251    return vit.vit_base_patch16_224(**default_kwargs)252 253 254def vit_base_patch32_512(**kwargs):255    default_kwargs = dict(256        img_size=512,257        in_chans=6,258        num_classes=0,259        fc_norm=None,260        class_token=True,261        drop_path_rate=0.1,262        init_values=0.0001,263        block_fn=vit.ParallelScalingBlock,264        qkv_bias=False,265        qk_norm=True,266    )267    for k, v in kwargs.items():268        default_kwargs[k] = v269    return vit.vit_base_patch32_384(**default_kwargs)270 271 272def vit_large_patch8_256(**kwargs):273    default_kwargs = dict(274        img_size=256,275        in_chans=6,276        num_classes=0,277        fc_norm=None,278        class_token=True,279        patch_size=8,280        embed_dim=1024,281        depth=24,282        num_heads=16,283        drop_path_rate=0.3,284        init_values=0.0001,285        block_fn=vit.ParallelScalingBlock,286        qkv_bias=False,287        qk_norm=True,288    )289    for k, v in kwargs.items():290        default_kwargs[k] = v291    return vit.VisionTransformer(**default_kwargs)292 293 294def vit_large_patch16_256(**kwargs):295    default_kwargs = dict(296        img_size=256,297        in_chans=6,298        num_classes=0,299        fc_norm=None,300        class_token=True,301        drop_path_rate=0.3,302        init_values=0.0001,303        block_fn=vit.ParallelScalingBlock,304        qkv_bias=False,305        qk_norm=True,306    )307    for k, v in kwargs.items():308        default_kwargs[k] = v309    return vit.vit_large_patch16_384(**default_kwargs)310