recursionpharma/OpenPhenom
221.2k
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 