Team Ai
Modelpublic

recursionpharma/OpenPhenom

sourceHugging Faceupdated 7mo agoView on Hugging Face
22likes1.2kdownloads
mae_modules.py274 linesDownload Raw Back to root
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