recursionpharma/OpenPhenom
221.2k
1# © Recursion Pharmaceuticals 20242from functools import partial3from typing import Tuple, Union4 5import torch6import torch.nn as nn7from timm.models.helpers import checkpoint_seq8from timm.models.vision_transformer import Block, Mlp, VisionTransformer9 10from .masking import transformer_random_masking11from .vit import channel_agnostic_vit12 13# If interested in training new MAEs, combine an encoder and decoder into a new module, and you should14# leverage the flattening and unflattening utilities as needed from mae_utils.py.15# Be sure to use an encoder-decoder Linear projection layer to match encoder dims with decoder dimensions.16# As described in the paper, images are self-standardized at the start.17 18 19class SelfStandardize(nn.Module):20 def __init__(self) -> None:21 super().__init__()22 self.self_standardize = nn.LazyInstanceNorm2d(23 affine=False, track_running_stats=False24 )25 26 def forward(self, pixels: torch.Tensor) -> torch.Tensor:27 x = pixels.float() / 255.028 return self.self_standardize(x)29 30 31class MAEEncoder(nn.Module):32 def __init__(33 self,34 vit_backbone: VisionTransformer,35 max_in_chans: int = 6,36 channel_agnostic: bool = False,37 ) -> None:38 super().__init__()39 if channel_agnostic:40 self.vit_backbone = channel_agnostic_vit(41 vit_backbone, max_in_chans=max_in_chans42 )43 else:44 self.vit_backbone = vit_backbone45 self.max_in_chans = max_in_chans46 self.channel_agnostic = channel_agnostic47 48 @property49 def embed_dim(self) -> int:50 return int(self.vit_backbone.embed_dim)51 52 def forward(self, x: torch.Tensor) -> torch.Tensor:53 x = self.vit_backbone.forward_features(x)54 x = self.vit_backbone.forward_head(x)55 return x # type: ignore[no-any-return]56 57 def forward_masked(58 self,59 x: torch.Tensor,60 mask_ratio: float,61 constant_noise: Union[torch.Tensor, None] = None,62 ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:63 x = self.vit_backbone.patch_embed(x)64 x = self.vit_backbone._pos_embed(x) # adds class token65 x_ = x[:, 1:, :] # no class token66 x_, mask, ind_restore = transformer_random_masking(67 x_, mask_ratio, constant_noise68 )69 x = torch.cat([x[:, :1, :], x_], dim=1) # add class token70 x = self.vit_backbone.norm_pre(x)71 72 if self.vit_backbone.grad_checkpointing and not torch.jit.is_scripting():73 x = checkpoint_seq(self.vit_backbone.blocks, x)74 else:75 x = self.vit_backbone.blocks(x)76 x = self.vit_backbone.norm(x)77 return x, mask, ind_restore78 79 80class MAEDecoder(nn.Module):81 def __init__(82 self,83 embed_dim: int = 512,84 depth: int = 8,85 num_heads: int = 16,86 mlp_ratio: float = 4,87 qkv_bias: bool = True,88 norm_layer: nn.Module = partial(nn.LayerNorm, eps=1e-6), # type: ignore[assignment]89 ) -> None:90 super().__init__()91 self.embed_dim = embed_dim92 self.pos_embeddings = None # to be overwritten by MAE class93 self.mask_token = nn.Parameter(torch.zeros(1, 1, embed_dim))94 self.blocks = nn.Sequential(95 *[96 Block(97 embed_dim,98 num_heads,99 mlp_ratio,100 qkv_bias=qkv_bias,101 norm_layer=norm_layer,102 )103 for i in range(depth)104 ]105 )106 self.norm = norm_layer(embed_dim)107 108 def forward(self, x: torch.Tensor) -> torch.Tensor:109 x = x + self.pos_embeddings110 x = self.blocks(x)111 x = self.norm(x)112 return x # type: ignore[no-any-return]113 114 def forward_masked(115 self, x: torch.Tensor, ind_restore: torch.Tensor116 ) -> torch.Tensor:117 mask_tokens = self.mask_token.repeat(118 x.shape[0], ind_restore.shape[1] + 1 - x.shape[1], 1119 )120 x_ = torch.cat([x[:, 1:, :], mask_tokens], dim=1) # remove class token121 x_ = torch.gather(122 x_, dim=1, index=ind_restore.unsqueeze(-1).repeat(1, 1, x.shape[2])123 ) # unshuffle124 x = torch.cat([x[:, :1, :], x_], dim=1) # add class token125 126 x = x + self.pos_embeddings127 x = self.blocks(x)128 x = self.norm(x)129 return x # type: ignore[no-any-return]130 131 132class CrossAttention(nn.Module):133 def __init__(134 self, embed_dim, num_heads=8, qkv_bias=False, attn_drop=0.0, proj_drop=0.0135 ):136 super().__init__()137 self.num_heads = num_heads138 head_dim = embed_dim // num_heads139 self.scale = head_dim**-0.5140 141 self.q = nn.Linear(embed_dim, embed_dim, bias=qkv_bias)142 self.kv = nn.Linear(embed_dim, embed_dim * 2, bias=qkv_bias)143 144 self.attn_drop = nn.Dropout(attn_drop)145 self.proj = nn.Linear(embed_dim, embed_dim)146 self.proj_drop = nn.Dropout(proj_drop)147 148 def forward(self, x, context):149 B, N, C = x.shape150 _, M, _ = context.shape151 152 q = (153 self.q(x)154 .reshape(B, N, self.num_heads, C // self.num_heads)155 .permute(0, 2, 1, 3)156 )157 kv = (158 self.kv(context)159 .reshape(B, M, 2, self.num_heads, C // self.num_heads)160 .permute(2, 0, 3, 1, 4)161 )162 k, v = kv[0], kv[1]163 164 attn = (q @ k.transpose(-2, -1)) * self.scale165 attn = attn.softmax(dim=-1)166 attn = self.attn_drop(attn)167 168 x = (attn @ v).transpose(1, 2).reshape(B, N, -1)169 x = self.proj(x)170 x = self.proj_drop(x)171 return x172 173 174class CAMAEDecoder(nn.Module):175 def __init__(176 self,177 num_modalities: int = 6,178 tokens_per_modality: int = 256,179 embed_dim: int = 256,180 depth: int = 2,181 num_heads: int = 16,182 mlp_ratio: float = 4,183 qkv_bias: bool = True,184 norm_layer: nn.Module = partial(nn.LayerNorm, eps=1e-6), # type: ignore[assignment]185 ) -> None:186 super().__init__()187 self.num_modalities = num_modalities188 self.tokens_per_modality = tokens_per_modality189 self.embed_dim = embed_dim190 self.pos_embeddings = None # to be overwritten by MAE class191 self.mask_token = nn.Parameter(torch.zeros(1, 1, embed_dim))192 self.placeholder = nn.Parameter(193 torch.zeros(1, 1, embed_dim), requires_grad=False194 )195 self.modality_tokens = nn.ParameterList(196 [197 nn.Parameter(torch.zeros(1, 1, self.embed_dim))198 for modality in range(self.num_modalities)199 ]200 )201 202 self.cross_attention = CrossAttention(embed_dim=self.embed_dim)203 self.mlp = Mlp(self.embed_dim, hidden_features=int(self.embed_dim * mlp_ratio))204 205 self.decoders = nn.ModuleList(206 [207 nn.Sequential(208 *[209 Block(210 embed_dim,211 num_heads,212 mlp_ratio,213 qkv_bias=qkv_bias,214 norm_layer=norm_layer,215 )216 for i in range(depth)217 ]218 )219 for modality in range(self.num_modalities)220 ]221 )222 # self.norm = norm_layer(embed_dim) # we decided to drop the last layer norm223 self.context_norm = norm_layer(embed_dim)224 self.query_norm = norm_layer(embed_dim)225 self.out_norm = norm_layer(embed_dim)226 227 def forward(self, x: torch.Tensor) -> torch.Tensor:228 x_m_s = []229 230 modality_tokens_concat = torch.cat(231 [232 self.placeholder,233 ] # placeholder for class token234 + [235 m_t.repeat(1, self.tokens_per_modality, 1)236 for m_t in self.modality_tokens237 ],238 dim=1,239 )240 241 x = (242 x + self.pos_embeddings + modality_tokens_concat243 ) # add pos and tiled modality tokens244 x_ = x[:, 1:, :] # no class token245 for m, decoder in enumerate(246 self.decoders247 ): # iterate through modalities and decoders248 x_m = x_[249 :, m * self.tokens_per_modality : (m + 1) * self.tokens_per_modality, :250 ]251 x_m = self.cross_attention(self.query_norm(x_m), self.context_norm(x_))252 x_m = x_m + self.mlp(self.out_norm(x_m))253 x_m = decoder(x_m)254 x_m_s.append(x_m)255 x_m_s = torch.cat(x_m_s, dim=1) # concat all tokens256 # x_m_s = self.norm(x_m_s) # we decided to drop the last layer norm257 x_m_s = torch.cat([x[:, :1, :], x_m_s], dim=1) # add back class token258 259 return x_m_s260 261 def forward_masked(262 self, x: torch.Tensor, ind_restore: torch.Tensor263 ) -> torch.Tensor:264 mask_tokens = self.mask_token.repeat(265 x.shape[0], ind_restore.shape[1] + 1 - x.shape[1], 1266 )267 x_ = torch.cat([x[:, 1:, :], mask_tokens], dim=1) # remove class token268 x_ = torch.gather(269 x_, dim=1, index=ind_restore.unsqueeze(-1).repeat(1, 1, x.shape[2])270 ) # unshuffle271 x = torch.cat([x[:, :1, :], x_], dim=1) # add class token272 x = self.forward(x)273 return x274 