OneScience-Group/ESM
024
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 )