Team Ai
Modelpublic

OneScience-Group/ESM

sourceHugging Facemitupdated 2mo agoView on Hugging Face
0likes24downloads
structure_module.py1253 linesDownload Raw Back to openfold
1# Copyright 2021 AlQuraishi Laboratory2# Copyright 2021 DeepMind Technologies Limited3#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8#      http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15from functools import reduce16import importlib17import math18import sys19from operator import mul20 21import torch22import torch.nn as nn23from typing import Optional, Tuple, Sequence, Union24 25from model.openfold.primitives import Linear, LayerNorm, ipa_point_weights_init_26from onescience.utils.openfold.np.residue_constants import (27    restype_rigid_group_default_frame,28    restype_atom14_to_rigid_group,29    restype_atom14_mask,30    restype_atom14_rigid_group_positions,31)32from onescience.utils.openfold.geometry.quat_rigid import QuatRigid33from onescience.utils.openfold.geometry.rigid_matrix_vector import Rigid3Array34from onescience.utils.openfold.geometry.vector import Vec3Array, square_euclidean_distance35from onescience.utils.openfold.feats import (36    frames_and_literature_positions_to_atom14_pos,37    torsion_angles_to_frames,38)39from onescience.utils.openfold.precision_utils import is_fp16_enabled40from onescience.utils.openfold.rigid_utils import Rotation, Rigid41from onescience.utils.openfold.tensor_utils import (42    dict_multimap,43    permute_final_dims,44    flatten_final_dims,45)46 47attn_core_inplace_cuda = importlib.import_module("attn_core_inplace_cuda")48 49 50class AngleResnetBlock(nn.Module):51    def __init__(self, c_hidden):52        """53        Args:54            c_hidden:55                Hidden channel dimension56        """57        super(AngleResnetBlock, self).__init__()58 59        self.c_hidden = c_hidden60 61        self.linear_1 = Linear(self.c_hidden, self.c_hidden, init="relu")62        self.linear_2 = Linear(self.c_hidden, self.c_hidden, init="final")63 64        self.relu = nn.ReLU()65 66    def forward(self, a: torch.Tensor) -> torch.Tensor:67 68        s_initial = a69 70        a = self.relu(a)71        a = self.linear_1(a)72        a = self.relu(a)73        a = self.linear_2(a)74 75        return a + s_initial76 77 78class AngleResnet(nn.Module):79    """80    Implements Algorithm 20, lines 11-1481    """82 83    def __init__(self, c_in, c_hidden, no_blocks, no_angles, epsilon):84        """85        Args:86            c_in:87                Input channel dimension88            c_hidden:89                Hidden channel dimension90            no_blocks:91                Number of resnet blocks92            no_angles:93                Number of torsion angles to generate94            epsilon:95                Small constant for normalization96        """97        super(AngleResnet, self).__init__()98 99        self.c_in = c_in100        self.c_hidden = c_hidden101        self.no_blocks = no_blocks102        self.no_angles = no_angles103        self.eps = epsilon104 105        self.linear_in = Linear(self.c_in, self.c_hidden)106        self.linear_initial = Linear(self.c_in, self.c_hidden)107 108        self.layers = nn.ModuleList()109        for _ in range(self.no_blocks):110            layer = AngleResnetBlock(c_hidden=self.c_hidden)111            self.layers.append(layer)112 113        self.linear_out = Linear(self.c_hidden, self.no_angles * 2)114 115        self.relu = nn.ReLU()116 117    def forward(118        self, s: torch.Tensor, s_initial: torch.Tensor119    ) -> Tuple[torch.Tensor, torch.Tensor]:120        """121        Args:122            s:123                [*, C_hidden] single embedding124            s_initial:125                [*, C_hidden] single embedding as of the start of the126                StructureModule127        Returns:128            [*, no_angles, 2] predicted angles129        """130        # NOTE: The ReLU's applied to the inputs are absent from the supplement131        # pseudocode but present in the source. For maximal compatibility with132        # the pretrained weights, I'm going with the source.133 134        # [*, C_hidden]135        s_initial = self.relu(s_initial)136        s_initial = self.linear_initial(s_initial)137        s = self.relu(s)138        s = self.linear_in(s)139        s = s + s_initial140 141        for l in self.layers:142            s = l(s)143 144        s = self.relu(s)145 146        # [*, no_angles * 2]147        s = self.linear_out(s)148 149        # [*, no_angles, 2]150        s = s.view(s.shape[:-1] + (-1, 2))151 152        unnormalized_s = s153        norm_denom = torch.sqrt(154            torch.clamp(155                torch.sum(s ** 2, dim=-1, keepdim=True),156                min=self.eps,157            )158        )159        s = s / norm_denom160 161        return unnormalized_s, s162 163 164class PointProjection(nn.Module):165    def __init__(self,166        c_hidden: int,167        num_points: int,168        no_heads: int,169        is_multimer: bool,170        return_local_points: bool = False,171    ):172        super().__init__()173        self.return_local_points = return_local_points174        self.no_heads = no_heads175        self.num_points = num_points176        self.is_multimer = is_multimer177 178        # Multimer requires this to be run with fp32 precision during training179        precision = torch.float32 if self.is_multimer else None180        self.linear = Linear(c_hidden, no_heads * 3 * num_points, precision=precision)181 182    def forward(self, 183        activations: torch.Tensor, 184        rigids: Union[Rigid, Rigid3Array],185    ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:186        # TODO: Needs to run in high precision during training187        points_local = self.linear(activations)188        out_shape = points_local.shape[:-1] + (self.no_heads, self.num_points, 3)189 190        if self.is_multimer:191            points_local = points_local.view(192                points_local.shape[:-1] + (self.no_heads, -1)193            )194 195        points_local = torch.split(196            points_local, points_local.shape[-1] // 3, dim=-1197        )198 199        points_local = torch.stack(points_local, dim=-1).view(out_shape)200 201        points_global = rigids[..., None, None].apply(points_local)202 203        if(self.return_local_points):204            return points_global, points_local205 206        return points_global207 208 209class InvariantPointAttention(nn.Module):210    """211    Implements Algorithm 22.212    """213    def __init__(214        self,215        c_s: int,216        c_z: int,217        c_hidden: int,218        no_heads: int,219        no_qk_points: int,220        no_v_points: int,221        inf: float = 1e5,222        eps: float = 1e-8,223        is_multimer: bool = False,224    ):225        """226        Args:227            c_s:228                Single representation channel dimension229            c_z:230                Pair representation channel dimension231            c_hidden:232                Hidden channel dimension233            no_heads:234                Number of attention heads235            no_qk_points:236                Number of query/key points to generate237            no_v_points:238                Number of value points to generate239        """240        super(InvariantPointAttention, self).__init__()241 242        self.c_s = c_s243        self.c_z = c_z244        self.c_hidden = c_hidden245        self.no_heads = no_heads246        self.no_qk_points = no_qk_points247        self.no_v_points = no_v_points248        self.inf = inf249        self.eps = eps250        self.is_multimer = is_multimer251 252        # These linear layers differ from their specifications in the253        # supplement. There, they lack bias and use Glorot initialization.254        # Here as in the official source, they have bias and use the default255        # Lecun initialization.256        hc = self.c_hidden * self.no_heads257        self.linear_q = Linear(self.c_s, hc, bias=(not is_multimer))258 259        self.linear_q_points = PointProjection(260            self.c_s,261            self.no_qk_points,262            self.no_heads,263            self.is_multimer264        )265 266        if(is_multimer):267            self.linear_k = Linear(self.c_s, hc, bias=False)268            self.linear_v = Linear(self.c_s, hc, bias=False)269            self.linear_k_points = PointProjection(270                self.c_s,271                self.no_qk_points,272                self.no_heads,273                self.is_multimer274            )275 276            self.linear_v_points = PointProjection(277                self.c_s,278                self.no_v_points,279                self.no_heads,280                self.is_multimer281            )282        else:283            self.linear_kv = Linear(self.c_s, 2 * hc)284            self.linear_kv_points = PointProjection(285                self.c_s,286                self.no_qk_points + self.no_v_points,287                self.no_heads,288                self.is_multimer289            )290 291        self.linear_b = Linear(self.c_z, self.no_heads)292 293        self.head_weights = nn.Parameter(torch.zeros((no_heads)))294        ipa_point_weights_init_(self.head_weights)295 296        concat_out_dim = self.no_heads * (297            self.c_z + self.c_hidden + self.no_v_points * 4298        )299        self.linear_out = Linear(concat_out_dim, self.c_s, init="final")300 301        self.softmax = nn.Softmax(dim=-1)302        self.softplus = nn.Softplus()303 304    def forward(305        self,306        s: torch.Tensor,307        z: torch.Tensor,308        r: Union[Rigid, Rigid3Array],309        mask: torch.Tensor,310        inplace_safe: bool = False,311        _offload_inference: bool = False,312        _z_reference_list: Optional[Sequence[torch.Tensor]] = None,313    ) -> torch.Tensor:314        """315        Args:316            s:317                [*, N_res, C_s] single representation318            z:319                [*, N_res, N_res, C_z] pair representation320            r:321                [*, N_res] transformation object322            mask:323                [*, N_res] mask324        Returns:325            [*, N_res, C_s] single representation update326        """327        if (_offload_inference and inplace_safe):328            z = _z_reference_list329        else:330            z = [z]331 332        #######################################333        # Generate scalar and point activations334        #######################################335        # [*, N_res, H * C_hidden]336        q = self.linear_q(s)337 338        # [*, N_res, H, C_hidden]339        q = q.view(q.shape[:-1] + (self.no_heads, -1))340 341        # [*, N_res, H, P_qk]342        q_pts = self.linear_q_points(s, r)343 344        # The following two blocks are equivalent345        # They're separated only to preserve compatibility with old AF weights346        if(self.is_multimer):347            # [*, N_res, H * C_hidden]348            k = self.linear_k(s)349            v = self.linear_v(s)350 351            # [*, N_res, H, C_hidden]352            k = k.view(k.shape[:-1] + (self.no_heads, -1))353            v = v.view(v.shape[:-1] + (self.no_heads, -1))354 355            # [*, N_res, H, P_qk, 3]356            k_pts = self.linear_k_points(s, r)357 358            # [*, N_res, H, P_v, 3]359            v_pts = self.linear_v_points(s, r)360        else:361            # [*, N_res, H * 2 * C_hidden]362            kv = self.linear_kv(s)363 364            # [*, N_res, H, 2 * C_hidden]365            kv = kv.view(kv.shape[:-1] + (self.no_heads, -1))366 367            # [*, N_res, H, C_hidden]368            k, v = torch.split(kv, self.c_hidden, dim=-1)369 370            kv_pts = self.linear_kv_points(s, r)371 372            # [*, N_res, H, P_q/P_v, 3]373            k_pts, v_pts = torch.split(374                kv_pts, [self.no_qk_points, self.no_v_points], dim=-2375            )376 377        ##########################378        # Compute attention scores379        ##########################380        # [*, N_res, N_res, H]381        b = self.linear_b(z[0])382 383        if (_offload_inference):384            assert (sys.getrefcount(z[0]) == 2)385            z[0] = z[0].cpu()386 387        # [*, H, N_res, N_res]388        if (is_fp16_enabled()):389            with torch.cuda.amp.autocast(enabled=False):390                a = torch.matmul(391                    permute_final_dims(q.float(), (1, 0, 2)),  # [*, H, N_res, C_hidden]392                    permute_final_dims(k.float(), (1, 2, 0)),  # [*, H, C_hidden, N_res]393                )394        else:395            a = torch.matmul(396                permute_final_dims(q, (1, 0, 2)),  # [*, H, N_res, C_hidden]397                permute_final_dims(k, (1, 2, 0)),  # [*, H, C_hidden, N_res]398            )399 400        a *= math.sqrt(1.0 / (3 * self.c_hidden))401        a += (math.sqrt(1.0 / 3) * permute_final_dims(b, (2, 0, 1)))402 403        # [*, N_res, N_res, H, P_q, 3]404        pt_att = q_pts.unsqueeze(-4) - k_pts.unsqueeze(-5)405 406        if (inplace_safe):407            pt_att *= pt_att408        else:409            pt_att = pt_att ** 2410 411        pt_att = sum(torch.unbind(pt_att, dim=-1))412 413        head_weights = self.softplus(self.head_weights).view(414            *((1,) * len(pt_att.shape[:-2]) + (-1, 1))415        )416        head_weights = head_weights * math.sqrt(417            1.0 / (3 * (self.no_qk_points * 9.0 / 2))418        )419 420        if (inplace_safe):421            pt_att *= head_weights422        else:423            pt_att = pt_att * head_weights424 425        # [*, N_res, N_res, H]426        pt_att = torch.sum(pt_att, dim=-1) * (-0.5)427 428        # [*, N_res, N_res]429        square_mask = mask.unsqueeze(-1) * mask.unsqueeze(-2)430        square_mask = self.inf * (square_mask - 1)431 432        # [*, H, N_res, N_res]433        pt_att = permute_final_dims(pt_att, (2, 0, 1))434 435        if (inplace_safe):436            a += pt_att437            del pt_att438            a += square_mask.unsqueeze(-3)439            # in-place softmax440            attn_core_inplace_cuda.forward_(441                a,442                reduce(mul, a.shape[:-1]),443                a.shape[-1],444            )445        else:446            a = a + pt_att447            a = a + square_mask.unsqueeze(-3)448            a = self.softmax(a)449 450        ################451        # Compute output452        ################453        # [*, N_res, H, C_hidden]454        o = torch.matmul(455            a, v.transpose(-2, -3).to(dtype=a.dtype)456        ).transpose(-2, -3)457 458        # [*, N_res, H * C_hidden]459        o = flatten_final_dims(o, 2)460 461        # [*, H, 3, N_res, P_v]462        if (inplace_safe):463            v_pts = permute_final_dims(v_pts, (1, 3, 0, 2))464            o_pt = [465                torch.matmul(a, v.to(a.dtype))466                for v in torch.unbind(v_pts, dim=-3)467            ]468            o_pt = torch.stack(o_pt, dim=-3)469        else:470            o_pt = torch.sum(471                (472                        a[..., None, :, :, None]473                        * permute_final_dims(v_pts, (1, 3, 0, 2))[..., None, :, :]474                ),475                dim=-2,476            )477 478        # [*, N_res, H, P_v, 3]479        o_pt = permute_final_dims(o_pt, (2, 0, 3, 1))480        o_pt = r[..., None, None].invert_apply(o_pt)481 482        # [*, N_res, H * P_v]483        o_pt_norm = flatten_final_dims(484            torch.sqrt(torch.sum(o_pt ** 2, dim=-1) + self.eps), 2485        )486 487        # [*, N_res, H * P_v, 3]488        o_pt = o_pt.reshape(*o_pt.shape[:-3], -1, 3)489        o_pt = torch.unbind(o_pt, dim=-1)490 491        if (_offload_inference):492            z[0] = z[0].to(o_pt.device)493 494        # [*, N_res, H, C_z]495        o_pair = torch.matmul(a.transpose(-2, -3), z[0].to(dtype=a.dtype))496 497        # [*, N_res, H * C_z]498        o_pair = flatten_final_dims(o_pair, 2)499 500        # [*, N_res, C_s]501        s = self.linear_out(502            torch.cat(503                (o, *o_pt, o_pt_norm, o_pair), dim=-1504            ).to(dtype=z[0].dtype)505        )506 507        return s508 509 510#TODO: This module follows the refactoring done in IPA for multimer. Running the regular IPA above511# in multimer mode should be equivalent, but tests do not pass unless using this version. Determine512# whether or not the increase in test error matters in practice.513class InvariantPointAttentionMultimer(nn.Module):514    """515    Implements Algorithm 22.516    """517    def __init__(518        self,519        c_s: int,520        c_z: int,521        c_hidden: int,522        no_heads: int,523        no_qk_points: int,524        no_v_points: int,525        inf: float = 1e5,526        eps: float = 1e-8,527        is_multimer: bool = True,528    ):529        """530        Args:531            c_s:532                Single representation channel dimension533            c_z:534                Pair representation channel dimension535            c_hidden:536                Hidden channel dimension537            no_heads:538                Number of attention heads539            no_qk_points:540                Number of query/key points to generate541            no_v_points:542                Number of value points to generate543        """544        super(InvariantPointAttentionMultimer, self).__init__()545 546        self.c_s = c_s547        self.c_z = c_z548        self.c_hidden = c_hidden549        self.no_heads = no_heads550        self.no_qk_points = no_qk_points551        self.no_v_points = no_v_points552        self.inf = inf553        self.eps = eps554 555        # These linear layers differ from their specifications in the556        # supplement. There, they lack bias and use Glorot initialization.557        # Here as in the official source, they have bias and use the default558        # Lecun initialization.559        hc = self.c_hidden * self.no_heads560        self.linear_q = Linear(self.c_s, hc, bias=False)561 562        self.linear_q_points = PointProjection(563            self.c_s,564            self.no_qk_points,565            self.no_heads,566            is_multimer=True567        )568 569        self.linear_k = Linear(self.c_s, hc, bias=False)570        self.linear_v = Linear(self.c_s, hc, bias=False)571        self.linear_k_points = PointProjection(572            self.c_s,573            self.no_qk_points,574            self.no_heads,575            is_multimer=True576        )577 578        self.linear_v_points = PointProjection(579            self.c_s,580            self.no_v_points,581            self.no_heads,582            is_multimer=True583        )584 585        self.linear_b = Linear(self.c_z, self.no_heads)586 587        self.head_weights = nn.Parameter(torch.zeros((no_heads)))588        ipa_point_weights_init_(self.head_weights)589 590        concat_out_dim = self.no_heads * (591            self.c_z + self.c_hidden + self.no_v_points * 4592        )593        self.linear_out = Linear(concat_out_dim, self.c_s, init="final")594 595        self.softmax = nn.Softmax(dim=-2)596 597    def forward(598        self,599        s: torch.Tensor,600        z: Optional[torch.Tensor],601        r: Union[Rigid, Rigid3Array],602        mask: torch.Tensor,603        inplace_safe: bool = False,604        _offload_inference: bool = False,605        _z_reference_list: Optional[Sequence[torch.Tensor]] = None,606    ) -> torch.Tensor:607        """608        Args:609            s:610                [*, N_res, C_s] single representation611            z:612                [*, N_res, N_res, C_z] pair representation613            r:614                [*, N_res] transformation object615            mask:616                [*, N_res] mask617        Returns:618            [*, N_res, C_s] single representation update619        """620        if(_offload_inference and inplace_safe):621            z = _z_reference_list622        else:623            z = [z]624 625        a = 0.626 627        point_variance = (max(self.no_qk_points, 1) * 9.0 / 2)628        point_weights = math.sqrt(1.0 / point_variance)629 630        softplus = lambda x: torch.logaddexp(x, torch.zeros_like(x))631 632        head_weights = softplus(self.head_weights)633        point_weights = point_weights * head_weights634 635        #######################################636        # Generate scalar and point activations637        #######################################638 639        # [*, N_res, H, P_qk]640        q_pts = Vec3Array.from_array(self.linear_q_points(s, r))641 642        # [*, N_res, H, P_qk, 3]643        k_pts = Vec3Array.from_array(self.linear_k_points(s, r))644 645        pt_att = square_euclidean_distance(q_pts.unsqueeze(-3), k_pts.unsqueeze(-4), epsilon=0.)646        pt_att = torch.sum(pt_att * point_weights[..., None], dim=-1) * (-0.5)647        pt_att = pt_att.to(dtype=s.dtype)648        a = a + pt_att649 650        scalar_variance = max(self.c_hidden, 1) * 1.651        scalar_weights = math.sqrt(1.0 / scalar_variance)652 653        # [*, N_res, H * C_hidden]654        q = self.linear_q(s)655        k = self.linear_k(s)656 657        # [*, N_res, H, C_hidden]658        q = q.view(q.shape[:-1] + (self.no_heads, -1))659        k = k.view(k.shape[:-1] + (self.no_heads, -1))660 661        q = q * scalar_weights662        a = a + torch.einsum('...qhc,...khc->...qkh', q, k)663 664        ##########################665        # Compute attention scores666        ##########################667        # [*, N_res, N_res, H]668        b = self.linear_b(z[0])669 670        if (_offload_inference):671            assert (sys.getrefcount(z[0]) == 2)672            z[0] = z[0].cpu()673 674        a = a + b675 676        # [*, N_res, N_res]677        square_mask = mask.unsqueeze(-1) * mask.unsqueeze(-2)678        square_mask = self.inf * (square_mask - 1)679 680        a = a + square_mask.unsqueeze(-1)681        a = a * math.sqrt(1. / 3)  # Normalize by number of logit terms (3)682        a = self.softmax(a)683 684        # [*, N_res, H * C_hidden]685        v = self.linear_v(s)686 687        # [*, N_res, H, C_hidden]688        v = v.view(v.shape[:-1] + (self.no_heads, -1))689 690        o = torch.einsum('...qkh, ...khc->...qhc', a, v)691 692        # [*, N_res, H * C_hidden]693        o = flatten_final_dims(o, 2)694 695        # [*, N_res, H, P_v, 3]696        v_pts = Vec3Array.from_array(self.linear_v_points(s, r))697 698        # [*, N_res, H, P_v]699        o_pt = v_pts[..., None, :, :, :] * a.unsqueeze(-1)700        o_pt = o_pt.sum(dim=-3)701        # o_pt = Vec3Array(702        #     torch.sum(a.unsqueeze(-1) * v_pts[..., None, :, :, :].x, dim=-3),703        #     torch.sum(a.unsqueeze(-1) * v_pts[..., None, :, :, :].y, dim=-3),704        #     torch.sum(a.unsqueeze(-1) * v_pts[..., None, :, :, :].z, dim=-3),705        # )706 707        # [*, N_res, H * P_v, 3]708        o_pt = o_pt.reshape(o_pt.shape[:-2] + (-1,))709 710        # [*, N_res, H, P_v]711        o_pt = r[..., None].apply_inverse_to_point(o_pt)712        o_pt_flat = [o_pt.x, o_pt.y, o_pt.z]713        o_pt_flat = [x.to(dtype=a.dtype) for x in o_pt_flat]714 715        # [*, N_res, H * P_v]716        o_pt_norm = o_pt.norm(epsilon=1e-8)717 718        if (_offload_inference):719            z[0] = z[0].to(o_pt.x.device)720 721        o_pair = torch.einsum('...ijh, ...ijc->...ihc', a, z[0].to(dtype=a.dtype))722 723        # [*, N_res, H * C_z]724        o_pair = flatten_final_dims(o_pair, 2)725 726        # [*, N_res, C_s]727        s = self.linear_out(728            torch.cat(729                (o, *o_pt_flat, o_pt_norm, o_pair), dim=-1730            ).to(dtype=z[0].dtype)731        )732 733        return s734 735 736class BackboneUpdate(nn.Module):737    """738    Implements part of Algorithm 23.739    """740 741    def __init__(self, c_s):742        """743        Args:744            c_s:745                Single representation channel dimension746        """747        super(BackboneUpdate, self).__init__()748 749        self.c_s = c_s750 751        self.linear = Linear(self.c_s, 6, init="final")752 753    def forward(self, s: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:754        """755        Args:756            [*, N_res, C_s] single representation757        Returns:758            [*, N_res, 6] update vector 759        """760        # [*, 6]761        update = self.linear(s)762 763        return update 764 765 766class StructureModuleTransitionLayer(nn.Module):767    def __init__(self, c):768        super(StructureModuleTransitionLayer, self).__init__()769 770        self.c = c771 772        self.linear_1 = Linear(self.c, self.c, init="relu")773        self.linear_2 = Linear(self.c, self.c, init="relu")774        self.linear_3 = Linear(self.c, self.c, init="final")775 776        self.relu = nn.ReLU()777 778    def forward(self, s):779        s_initial = s780        s = self.linear_1(s)781        s = self.relu(s)782        s = self.linear_2(s)783        s = self.relu(s)784        s = self.linear_3(s)785 786        s = s + s_initial787 788        return s789 790 791class StructureModuleTransition(nn.Module):792    def __init__(self, c, num_layers, dropout_rate):793        super(StructureModuleTransition, self).__init__()794 795        self.c = c796        self.num_layers = num_layers797        self.dropout_rate = dropout_rate798 799        self.layers = nn.ModuleList()800        for _ in range(self.num_layers):801            l = StructureModuleTransitionLayer(self.c)802            self.layers.append(l)803 804        self.dropout = nn.Dropout(self.dropout_rate)805        self.layer_norm = LayerNorm(self.c)806 807    def forward(self, s):808        for l in self.layers:809            s = l(s)810 811        s = self.dropout(s)812        s = self.layer_norm(s)813 814        return s815 816 817class StructureModule(nn.Module):818    def __init__(819        self,820        c_s,821        c_z,822        c_ipa,823        c_resnet,824        no_heads_ipa,825        no_qk_points,826        no_v_points,827        dropout_rate,828        no_blocks,829        no_transition_layers,830        no_resnet_blocks,831        no_angles,832        trans_scale_factor,833        epsilon,834        inf,835        is_multimer=False,836        **kwargs,837    ):838        """839        Args:840            c_s:841                Single representation channel dimension842            c_z:843                Pair representation channel dimension844            c_ipa:845                IPA hidden channel dimension846            c_resnet:847                Angle resnet (Alg. 23 lines 11-14) hidden channel dimension848            no_heads_ipa:849                Number of IPA heads850            no_qk_points:851                Number of query/key points to generate during IPA852            no_v_points:853                Number of value points to generate during IPA854            dropout_rate:855                Dropout rate used throughout the layer856            no_blocks:857                Number of structure module blocks858            no_transition_layers:859                Number of layers in the single representation transition860                (Alg. 23 lines 8-9)861            no_resnet_blocks:862                Number of blocks in the angle resnet863            no_angles:864                Number of angles to generate in the angle resnet865            trans_scale_factor:866                Scale of single representation transition hidden dimension867            epsilon:868                Small number used in angle resnet normalization869            inf:870                Large number used for attention masking871        """872        super(StructureModule, self).__init__()873 874        self.c_s = c_s875        self.c_z = c_z876        self.c_ipa = c_ipa877        self.c_resnet = c_resnet878        self.no_heads_ipa = no_heads_ipa879        self.no_qk_points = no_qk_points880        self.no_v_points = no_v_points881        self.dropout_rate = dropout_rate882        self.no_blocks = no_blocks883        self.no_transition_layers = no_transition_layers884        self.no_resnet_blocks = no_resnet_blocks885        self.no_angles = no_angles886        self.trans_scale_factor = trans_scale_factor887        self.epsilon = epsilon888        self.inf = inf889        self.is_multimer = is_multimer890 891        # Buffers to be lazily initialized later892        # self.default_frames893        # self.group_idx894        # self.atom_mask895        # self.lit_positions896 897        self.layer_norm_s = LayerNorm(self.c_s)898        self.layer_norm_z = LayerNorm(self.c_z)899 900        self.linear_in = Linear(self.c_s, self.c_s)901 902        ipa = InvariantPointAttention if not self.is_multimer else InvariantPointAttentionMultimer903        self.ipa = ipa(904            self.c_s,905            self.c_z,906            self.c_ipa,907            self.no_heads_ipa,908            self.no_qk_points,909            self.no_v_points,910            inf=self.inf,911            eps=self.epsilon,912            is_multimer=self.is_multimer,913        )914 915        self.ipa_dropout = nn.Dropout(self.dropout_rate)916        self.layer_norm_ipa = LayerNorm(self.c_s)917 918        self.transition = StructureModuleTransition(919            self.c_s,920            self.no_transition_layers,921            self.dropout_rate,922        )923 924        if self.is_multimer:925            self.bb_update = QuatRigid(self.c_s, full_quat=False)926        else:927            self.bb_update = BackboneUpdate(self.c_s)928 929        self.angle_resnet = AngleResnet(930            self.c_s,931            self.c_resnet,932            self.no_resnet_blocks,933            self.no_angles,934            self.epsilon,935        )936 937    def _forward_monomer(938        self,939        evoformer_output_dict,940        aatype,941        mask=None,942        inplace_safe=False,943        _offload_inference=False,944    ):945        """946        Args:947            evoformer_output_dict:948                Dictionary containing:949                    "single":950                        [*, N_res, C_s] single representation951                    "pair":952                        [*, N_res, N_res, C_z] pair representation953            aatype:954                [*, N_res] amino acid indices955            mask:956                Optional [*, N_res] sequence mask957        Returns:958            A dictionary of outputs959        """960        s = evoformer_output_dict["single"]961 962        if mask is None:963            # [*, N]964            mask = s.new_ones(s.shape[:-1])965 966        # [*, N, C_s]967        s = self.layer_norm_s(s)968 969        # [*, N, N, C_z]970        z = self.layer_norm_z(evoformer_output_dict["pair"])971 972        z_reference_list = None973        if (_offload_inference):974            assert (sys.getrefcount(evoformer_output_dict["pair"]) == 2)975            evoformer_output_dict["pair"] = evoformer_output_dict["pair"].cpu()976            z_reference_list = [z]977            z = None978 979        # [*, N, C_s]980        s_initial = s981        s = self.linear_in(s)982 983        # [*, N]984        rigids = Rigid.identity(985            s.shape[:-1], 986            s.dtype, 987            s.device, 988            self.training,989            fmt="quat",990        )991        outputs = []992        for i in range(self.no_blocks):993            # [*, N, C_s]994            s = s + self.ipa(995                s, 996                z, 997                rigids, 998                mask, 999                inplace_safe=inplace_safe,1000                _offload_inference=_offload_inference, 1001                _z_reference_list=z_reference_list1002            )1003            s = self.ipa_dropout(s)1004            s = self.layer_norm_ipa(s)1005            s = self.transition(s)1006           1007            # [*, N]1008            rigids = rigids.compose_q_update_vec(self.bb_update(s))1009 1010            # To hew as closely as possible to AlphaFold, we convert our1011            # quaternion-based transformations to rotation-matrix ones1012            # here1013            backb_to_global = Rigid(1014                Rotation(1015                    rot_mats=rigids.get_rots().get_rot_mats(), 1016                    quats=None1017                ),1018                rigids.get_trans(),1019            )1020 1021            backb_to_global = backb_to_global.scale_translation(1022                self.trans_scale_factor1023            )1024 1025            # [*, N, 7, 2]1026            unnormalized_angles, angles = self.angle_resnet(s, s_initial)1027 1028            all_frames_to_global = self.torsion_angles_to_frames(1029                backb_to_global,1030                angles,1031                aatype,1032            )1033 1034            pred_xyz = self.frames_and_literature_positions_to_atom14_pos(1035                all_frames_to_global,1036                aatype,1037            )1038 1039            scaled_rigids = rigids.scale_translation(self.trans_scale_factor)1040            1041            preds = {1042                "frames": scaled_rigids.to_tensor_7(),1043                "sidechain_frames": all_frames_to_global.to_tensor_4x4(),1044                "unnormalized_angles": unnormalized_angles,1045                "angles": angles,1046                "positions": pred_xyz,1047                "states": s,1048            }1049 1050            outputs.append(preds)1051 1052            rigids = rigids.stop_rot_gradient()1053 1054        del z, z_reference_list1055 1056        if (_offload_inference):1057            evoformer_output_dict["pair"] = (1058                evoformer_output_dict["pair"].to(s.device)1059            )1060 1061        outputs = dict_multimap(torch.stack, outputs)1062        outputs["single"] = s1063 1064        return outputs1065 1066    def _forward_multimer(1067            self,1068            evoformer_output_dict,1069            aatype,1070            mask=None,1071            inplace_safe=False,1072            _offload_inference=False,1073    ):1074        s = evoformer_output_dict["single"]1075 1076        if mask is None:1077            # [*, N]1078            mask = s.new_ones(s.shape[:-1])1079 1080        # [*, N, C_s]1081        s = self.layer_norm_s(s)1082 1083        # [*, N, N, C_z]1084        z = self.layer_norm_z(evoformer_output_dict["pair"])1085 1086        z_reference_list = None1087        if (_offload_inference):1088            assert (sys.getrefcount(evoformer_output_dict["pair"]) == 2)1089            evoformer_output_dict["pair"] = evoformer_output_dict["pair"].cpu()1090            z_reference_list = [z]1091            z = None1092 1093        # [*, N, C_s]1094        s_initial = s1095        s = self.linear_in(s)1096 1097        # [*, N]1098        rigids = Rigid3Array.identity(1099            s.shape[:-1], 1100            s.device, 1101        )1102        outputs = []1103        for i in range(self.no_blocks):1104            # [*, N, C_s]1105            s = s + self.ipa(1106                s,1107                z,1108                rigids,1109                mask,1110                inplace_safe=inplace_safe,1111                _offload_inference=_offload_inference,1112                _z_reference_list=z_reference_list1113            )1114            s = self.ipa_dropout(s)1115            s = self.layer_norm_ipa(s)1116            s = self.transition(s)1117 1118            # [*, N]1119            rigids = rigids @ self.bb_update(s)1120 1121            # [*, N, 7, 2]1122            unnormalized_angles, angles = self.angle_resnet(s, s_initial)1123 1124            all_frames_to_global = self.torsion_angles_to_frames(1125                rigids.scale_translation(self.trans_scale_factor),1126                angles,1127                aatype,1128            )1129 1130            pred_xyz = self.frames_and_literature_positions_to_atom14_pos(1131                all_frames_to_global,1132                aatype,1133            )1134            1135            preds = {1136                "frames": rigids.scale_translation(self.trans_scale_factor).to_tensor(),1137                "sidechain_frames": all_frames_to_global.to_tensor_4x4(),1138                "unnormalized_angles": unnormalized_angles,1139                "angles": angles,1140                "positions": pred_xyz,1141            }1142 1143            preds = {k: v.to(dtype=s.dtype) for k, v in preds.items()}1144 1145            outputs.append(preds)1146 1147            rigids = rigids.stop_rot_gradient()1148 1149        del z, z_reference_list1150 1151        if (_offload_inference):1152            evoformer_output_dict["pair"] = (1153                evoformer_output_dict["pair"].to(s.device)1154            )1155 1156        outputs = dict_multimap(torch.stack, outputs)1157        outputs["single"] = s1158 1159        return outputs1160 1161    def forward(1162        self,1163        evoformer_output_dict,1164        aatype,1165        mask=None,1166        inplace_safe=False,1167        _offload_inference=False,1168    ):1169        """1170        Args:1171            s:1172                [*, N_res, C_s] single representation1173            z:1174                [*, N_res, N_res, C_z] pair representation1175            aatype:1176                [*, N_res] amino acid indices1177            mask:1178                Optional [*, N_res] sequence mask1179        Returns:1180            A dictionary of outputs1181        """1182        if(self.is_multimer):1183            outputs = self._forward_multimer(evoformer_output_dict, aatype, mask, inplace_safe, _offload_inference)1184        else:1185            outputs = self._forward_monomer(evoformer_output_dict, aatype, mask, inplace_safe, _offload_inference)1186 1187        return outputs1188 1189    def _init_residue_constants(self, float_dtype, device):1190        if not hasattr(self, "default_frames"):1191            self.register_buffer(1192                "default_frames",1193                torch.tensor(1194                    restype_rigid_group_default_frame,1195                    dtype=float_dtype,1196                    device=device,1197                    requires_grad=False,1198                ),1199                persistent=False,1200            )

Showing the first 1,200 of 1253 lines. Download the file for the rest.