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.15 16from functools import partialmethod17from typing import Optional18from abc import ABC, abstractmethod19 20import torch21import torch.nn as nn22 23from model.openfold.primitives import Linear, LayerNorm24from onescience.utils.openfold.chunk_utils import chunk_layer25from onescience.utils.openfold.precision_utils import is_fp16_enabled26from onescience.utils.openfold.tensor_utils import add, permute_final_dims27 28 29class BaseTriangleMultiplicativeUpdate(nn.Module, ABC):30 """31 Implements Algorithms 11 and 12.32 """33 @abstractmethod34 def __init__(self, c_z, c_hidden, _outgoing, bias=True):35 """36 Args:37 c_z:38 Input channel dimension39 c:40 Hidden channel dimension41 """42 super(BaseTriangleMultiplicativeUpdate, self).__init__()43 self.c_z = c_z44 self.c_hidden = c_hidden45 self._outgoing = _outgoing46 self.bias = bias47 48 self.linear_g = Linear(self.c_z, self.c_z, bias=bias, init="gating")49 self.linear_z = Linear(self.c_hidden, self.c_z, bias=bias, init="final")50 51 self.layer_norm_in = LayerNorm(self.c_z)52 self.layer_norm_out = LayerNorm(self.c_hidden)53 54 self.sigmoid = nn.Sigmoid()55 56 def _combine_projections(self,57 a: torch.Tensor,58 b: torch.Tensor,59 _inplace_chunk_size: Optional[int] = None60 ) -> torch.Tensor:61 if(self._outgoing):62 a = permute_final_dims(a, (2, 0, 1))63 b = permute_final_dims(b, (2, 1, 0))64 else:65 a = permute_final_dims(a, (2, 1, 0))66 b = permute_final_dims(b, (2, 0, 1))67 68 if(_inplace_chunk_size is not None):69 # To be replaced by torch vmap70 for i in range(0, a.shape[-3], _inplace_chunk_size):71 a_chunk = a[..., i: i + _inplace_chunk_size, :, :]72 b_chunk = b[..., i: i + _inplace_chunk_size, :, :]73 a[..., i: i + _inplace_chunk_size, :, :] = (74 torch.matmul(75 a_chunk,76 b_chunk,77 )78 )79 80 p = a81 else:82 p = torch.matmul(a, b)83 84 return permute_final_dims(p, (1, 2, 0))85 86 @abstractmethod87 def forward(self,88 z: torch.Tensor,89 mask: Optional[torch.Tensor] = None,90 inplace_safe: bool = False,91 _add_with_inplace: bool = False92 ) -> torch.Tensor:93 """94 Args:95 x:96 [*, N_res, N_res, C_z] input tensor97 mask:98 [*, N_res, N_res] input mask99 Returns:100 [*, N_res, N_res, C_z] output tensor101 """102 pass103 104 105class TriangleMultiplicativeUpdate(BaseTriangleMultiplicativeUpdate):106 """107 Implements Algorithms 11 and 12.108 """109 def __init__(self, c_z, c_hidden, _outgoing=True, bias: bool=True):110 """111 Args:112 c_z:113 Input channel dimension114 c:115 Hidden channel dimension116 """117 super(TriangleMultiplicativeUpdate, self).__init__(c_z=c_z,118 c_hidden=c_hidden,119 _outgoing=_outgoing,120 bias=bias121 )122 123 self.linear_a_p = Linear(self.c_z, self.c_hidden, bias=bias)124 self.linear_a_g = Linear(self.c_z, self.c_hidden, bias=bias, init="gating")125 self.linear_b_p = Linear(self.c_z, self.c_hidden, bias=bias,)126 self.linear_b_g = Linear(self.c_z, self.c_hidden, bias=bias, init="gating")127 128 def _inference_forward(self,129 z: torch.Tensor,130 mask: Optional[torch.Tensor] = None,131 inplace_chunk_size: Optional[int] = None,132 with_add: bool = True,133 ):134 """135 Args:136 z:137 A [*, N, N, C_z] pair representation138 mask:139 A [*, N, N] pair mask140 inplace_chunk_size:141 Size of chunks used in the main computation. Increase to trade142 memory for speed.143 with_add:144 If True, z is overwritten with (z + update). Otherwise, it is145 overwritten with (update).146 Returns:147 A reference to the overwritten z148 149 More memory-efficient, inference-only version of the forward function.150 Uses in-place operations, fusion of the addition that happens after151 this module in the Evoformer, a smidge of recomputation, and 152 a cache of overwritten values to lower peak memory consumption of this153 module from 5x the size of the input tensor z to 2.5x its size. Useful154 for inference on extremely long sequences. 155 156 It works as follows. We will make reference to variables used in the157 default forward implementation below. Naively, triangle multiplication158 attention requires the manifestation of 5 tensors the size of z:159 1) z, the "square" input tensor, 2) a, the first projection of z, 160 3) b, the second projection of b, 4) g, a z-sized mask, and 5) a 161 z-sized tensor for intermediate computations. For large N, this is 162 prohibitively expensive; for N=4000, for example, z is more than 8GB 163 alone. To avoid this problem, we compute b, g, and all intermediate164 tensors in small chunks, noting that the chunks required to compute a165 chunk of the output depend only on the tensor a and corresponding 166 vertical and horizontal chunks of z. This suggests an algorithm that 167 loops over pairs of chunks of z: hereafter "columns" and "rows" of168 z, even though each "column" and "row" in fact contains169 inplace_chunk_size contiguous true columns and rows of z. Writing 170 output chunks to a new tensor would bring total memory consumption171 down to 3x the size of z. However, more memory can be saved by writing172 output chunks directly to z in-place. WLOG, we choose to write output173 chunks vertically, overwriting the ith "column" of z at the end of174 the ith iteration of the main loop. Despite this overwriting, the 175 ith column is always one column ahead of previously overwritten columns 176 and can be recovered directly from z. After the first iteration,177 however, the ith row of z is always at least partially overwritten. For178 this reason, we introduce the z-cache, a tensor one-half the size of 179 z. The z-cache initially contains the left half (2nd and 3rd quadrants)180 of z. For 0 < i < N/2, the missing left part of the ith row of z is181 recovered from this cache at the beginning of the ith iteration. Once i 182 exceeds n/2, the cache is "reoriented" to encompass the 3rd and 4th 183 quadrants of z instead. Though the 3rd quadrant of the original z is 184 entirely overwritten at this point, it can be recovered from the z-cache 185 itself. Thereafter, the ith row of z can be recovered in its entirety 186 from the reoriented z-cache. After the final iteration, z has been 187 completely overwritten and contains the triangular multiplicative 188 update. If with_add is True, it instead contains the sum of z and the189 triangular multiplicative update. In either case, peak memory 190 consumption is just 2.5x the size of z, disregarding memory used for 191 chunks and other small variables.192 """193 if mask is None:194 mask = z.new_ones(z.shape[:-1])195 196 mask = mask.unsqueeze(-1)197 198 def compute_projection_helper(pair, mask, a=True):199 if(a):200 linear_g = self.linear_a_g201 linear_p = self.linear_a_p202 else:203 linear_g = self.linear_b_g204 linear_p = self.linear_b_p205 206 pair = self.layer_norm_in(pair)207 p = linear_g(pair)208 p.sigmoid_()209 p *= linear_p(pair)210 p *= mask211 p = permute_final_dims(p, (2, 0, 1))212 return p213 214 def compute_projection(pair, mask, a=True, chunked=True): 215 need_transpose = self._outgoing ^ a216 if(not chunked):217 p = compute_projection_helper(pair, mask, a)218 if(need_transpose):219 p = p.transpose(-1, -2)220 else:221 # This computation is chunked so as not to exceed our 2.5x 222 # budget with a large intermediate tensor223 linear_g = self.linear_a_g if a else self.linear_b_g224 #c = linear_g.bias.shape[-1]225 if self.bias:226 c = linear_g.bias.shape[-1]227 else:228 c = linear_g.weight.shape[0]229 out_shape = pair.shape[:-3] + (c,) + pair.shape[-3:-1]230 p = pair.new_zeros(out_shape)231 for i in range(0, pair.shape[-3], inplace_chunk_size):232 pair_chunk = pair[..., i: i + inplace_chunk_size, :, :]233 mask_chunk = mask[..., i: i + inplace_chunk_size, :, :]234 pair_chunk = compute_projection_helper(235 pair[..., i: i + inplace_chunk_size, :, :],236 mask[..., i: i + inplace_chunk_size, :, :], 237 a,238 )239 if(need_transpose):240 pair_chunk = pair_chunk.transpose(-1, -2)241 p[..., i: i + inplace_chunk_size] = pair_chunk242 else:243 p[..., i: i + inplace_chunk_size, :] = pair_chunk244 245 del pair_chunk246 247 return p248 249 # We start by fully manifesting a. In addition to the input, this250 # brings total memory consumption to 2x z (disregarding size of chunks)251 # [*, N, N, c]252 a = compute_projection(z, mask, True, chunked=True)253 #if bias==True:254 # a = compute_projection(z, mask, True, True, chunked=True)255 #else:256 # a = compute_projection(z, mask, False, True, chunked=True)257 258 if(inplace_chunk_size is not None):259 n = a.shape[-1]260 half_n = n // 2 + n % 2261 row_dim = -3262 col_dim = -2263 b_chunk_dim = row_dim if self._outgoing else col_dim264 265 def empty_slicer(t):266 return [slice(None) for _ in t.shape]267 268 def slice_tensor(t, start, end, dim):269 # Slices start:end from the dim dimension of t270 s = empty_slicer(t)271 s[dim] = slice(start, end)272 return t[s]273 274 def flip_z_cache_(z_cache, z):275 # "Reorient" the z_cache (see below), filling it with quadrants276 # 3---recovered from the z_cache---and 4---recovered from z---277 # of the input tensor z. 278 quadrant_3 = slice_tensor(279 z_cache, half_n, None, row_dim280 )281 z_cache = z_cache.transpose(row_dim, col_dim)282 283 # If n is odd, we need to shrink the z_cache by one row284 z_cache = z_cache[..., :(n // 2), :, :]285 286 # Move the 3rd quadrant of z into the 287 first_half_slicer = empty_slicer(z_cache)288 first_half_slicer[col_dim] = slice(0, half_n)289 z_cache[first_half_slicer] = quadrant_3290 291 # Get the fourth quadrant of z292 quadrant_4 = slice_tensor(z, half_n, None, row_dim)293 quadrant_4 = slice_tensor(294 quadrant_4, half_n, None, col_dim295 )296 297 # Insert said quadrant into the rotated z-cache298 quadrant_3_slicer = empty_slicer(z_cache)299 quadrant_3_slicer[col_dim] = slice(half_n, None)300 301 z_cache[quadrant_3_slicer] = quadrant_4302 303 return z_cache304 305 # Initialize the z cache to the left half of z.306 z_cache_shape = list(z.shape)307 z_cache_shape[col_dim] = half_n308 z_cache = z.new_zeros(z_cache_shape)309 z_cache_slicer = empty_slicer(z_cache)310 z_cache_slicer[col_dim] = slice(0, half_n)311 z_cache.copy_(z[z_cache_slicer])312 z_cache_rotated = False313 314 # We need to reorient the z-cache at the halfway point, and we 315 # don't want a single chunk to straddle that point. We contract one316 # of the chunks in the middle to address that problem.317 i_range = list(range(0, half_n, inplace_chunk_size))318 initial_offsets = [319 i_2 - i_1 for i_1, i_2 in zip(i_range, i_range[1:] + [half_n])320 ]321 after_half = list(range(half_n, n, inplace_chunk_size))322 after_half_offsets = [inplace_chunk_size for _ in after_half]323 combined_range_with_offsets = zip(324 i_range + after_half, initial_offsets + after_half_offsets325 )326 for i, offset in combined_range_with_offsets:327 if(not z_cache_rotated and i >= half_n):328 z_cache = flip_z_cache_(z_cache, z)329 z_cache_rotated = True330 331 z_chunk_b = slice_tensor(332 z, i, i + offset, b_chunk_dim,333 )334 mask_chunk = slice_tensor(335 mask, i, i + offset, b_chunk_dim,336 )337 338 z_chunk_b = z_chunk_b.clone()339 if(b_chunk_dim == col_dim):340 z_chunk_b = slice_tensor(341 z, i, i + offset, col_dim342 )343 else: # b_chunk_dim == row_dim344 # In this case, the b-dimension (b_chunk_dim) is partially 345 # overwritten at the end of each iteration. We need to 346 # restore the missing component from the z-cache.347 if(not z_cache_rotated):348 z_chunk_slicer = empty_slicer(z_chunk_b)349 z_chunk_slicer[col_dim] = slice(0, half_n)350 z_chunk_b[z_chunk_slicer] = slice_tensor(351 z_cache, i, i + offset, row_dim,352 )353 else:354 z_cache_offset = i - half_n355 z_chunk_b = slice_tensor(356 z_cache, 357 z_cache_offset, z_cache_offset + offset, 358 row_dim359 )360 b_chunk = compute_projection(z_chunk_b, mask_chunk, a=False, chunked=False)361 #if bias==True:362 # b_chunk = compute_projection(363 # z_chunk_b, mask_chunk, True, a=False, chunked=False364 # )365 #else:366 # b_chunk = compute_projection(z_chunk_b, mask_chunk, False, a=False, chunked=False)367 del z_chunk_b368 369 x_chunk = torch.matmul(370 a,371 b_chunk,372 )373 x_chunk = permute_final_dims(x_chunk, (1, 2, 0))374 x_chunk = self.layer_norm_out(x_chunk)375 x_chunk = self.linear_z(x_chunk)376 377 # The g dimension (col_dim) is parallel to and ahead of the 378 # overwrites in z. We can extract the g chunk normally.379 z_chunk_g = slice_tensor(380 z, i, i + offset, col_dim381 )382 g_chunk = self.linear_g(self.layer_norm_in(z_chunk_g)) 383 g_chunk.sigmoid_()384 del z_chunk_g385 386 x_chunk *= g_chunk387 388 # Write the columns into z in-place389 z_slicer = empty_slicer(z)390 z_slicer[col_dim] = slice(i, i + offset)391 if(with_add):392 z[z_slicer] += x_chunk393 else:394 z[z_slicer] = x_chunk395 else:396 b = compute_projection(z, mask, False, False)397 #if bias==True:398 # b = compute_projection(z, mask, True, False, False)399 #else:400 # b = compute_projection(z, mask, False, False, False) 401 x = torch.matmul(a, b)402 x = self.layer_norm_out(x)403 x = self.linear_z(x)404 g = self.linear_g(z)405 g.sigmoid_()406 x *= g407 if(with_add):408 z += x409 else:410 z = x411 412 return z413 414 def forward(self, 415 z: torch.Tensor, 416 mask: Optional[torch.Tensor] = None,417 inplace_safe: bool = False,418 _add_with_inplace: bool = False,419 _inplace_chunk_size: Optional[int] = 256,420 ) -> torch.Tensor:421 """422 Args:423 x:424 [*, N_res, N_res, C_z] input tensor425 mask:426 [*, N_res, N_res] input mask427 Returns:428 [*, N_res, N_res, C_z] output tensor429 """430 if(inplace_safe):431 x = self._inference_forward(432 z, 433 mask, 434 inplace_chunk_size=_inplace_chunk_size,435 with_add=_add_with_inplace,436 )437 return x438 439 if mask is None:440 mask = z.new_ones(z.shape[:-1])441 442 mask = mask.unsqueeze(-1)443 444 z = self.layer_norm_in(z)445 a = mask446 a = a * self.sigmoid(self.linear_a_g(z)) 447 a = a * self.linear_a_p(z)448 b = mask449 b = b * self.sigmoid(self.linear_b_g(z))450 b = b * self.linear_b_p(z)451 452 # Prevents overflow of torch.matmul in combine projections in453 # reduced-precision modes454 a_std = a.std()455 b_std = b.std()456 if(is_fp16_enabled() and a_std != 0. and b_std != 0.):457 a = a / a.std()458 b = b / b.std()459 460 if(is_fp16_enabled()):461 with torch.cuda.amp.autocast(enabled=False):462 x = self._combine_projections(a.float(), b.float())463 else:464 x = self._combine_projections(a, b)465 466 del a, b467 x = self.layer_norm_out(x)468 x = self.linear_z(x)469 g = self.sigmoid(self.linear_g(z))470 x = x * g471 472 return x473 474 475class TriangleMultiplicationOutgoing(TriangleMultiplicativeUpdate):476 """477 Implements Algorithm 11.478 """479 __init__ = partialmethod(TriangleMultiplicativeUpdate.__init__, _outgoing=True)480 481 482class TriangleMultiplicationIncoming(TriangleMultiplicativeUpdate):483 """484 Implements Algorithm 12.485 """486 __init__ = partialmethod(TriangleMultiplicativeUpdate.__init__, _outgoing=False)487 488class ProtenixTriangleMultiplicationOutgoing(TriangleMultiplicativeUpdate):489 """490 Implements Algorithm 11.491 """492 __init__ = partialmethod(TriangleMultiplicativeUpdate.__init__, _outgoing=True, bias=False)493 494class ProtenixTriangleMultiplicationIncoming(TriangleMultiplicativeUpdate):495 """496 Implements Algorithm 12.497 """498 __init__ = partialmethod(TriangleMultiplicativeUpdate.__init__, _outgoing=False, bias=False)499 500 501class FusedTriangleMultiplicativeUpdate(BaseTriangleMultiplicativeUpdate):502 """503 Implements Algorithms 11 and 12.504 """505 506 def __init__(self, c_z, c_hidden, _outgoing=True):507 """508 Args:509 c_z:510 Input channel dimension511 c:512 Hidden channel dimension513 """514 super(FusedTriangleMultiplicativeUpdate, self).__init__(c_z=c_z,515 c_hidden=c_hidden,516 _outgoing=_outgoing)517 518 self.linear_ab_p = Linear(self.c_z, self.c_hidden * 2)519 self.linear_ab_g = Linear(self.c_z, self.c_hidden * 2, init="gating")520 521 def _inference_forward(self,522 z: torch.Tensor,523 mask: Optional[torch.Tensor] = None,524 _inplace_chunk_size: Optional[int] = None,525 with_add: bool = True,526 ):527 """528 Args:529 z:530 A [*, N, N, C_z] pair representation531 mask:532 A [*, N, N] pair mask533 with_add:534 If True, z is overwritten with (z + update). Otherwise, it is535 overwritten with (update).536 Returns:537 A reference to the overwritten z538 """539 if mask is None:540 mask = z.new_ones(z.shape[:-1])541 542 mask = mask.unsqueeze(-1)543 544 def compute_projection_helper(pair, mask):545 p = self.linear_ab_g(pair)546 p.sigmoid_()547 p *= self.linear_ab_p(pair)548 p *= mask549 550 return p551 552 def compute_projection(pair, mask):553 p = compute_projection_helper(pair, mask)554 left = p[..., :self.c_hidden]555 right = p[..., self.c_hidden:]556 557 return left, right558 559 z_norm_in = self.layer_norm_in(z)560 a, b = compute_projection(z_norm_in, mask)561 x = self._combine_projections(a, b, _inplace_chunk_size=_inplace_chunk_size)562 x = self.layer_norm_out(x)563 x = self.linear_z(x)564 g = self.linear_g(z_norm_in)565 g.sigmoid_()566 x *= g567 if (with_add):568 z += x569 else:570 z = x571 572 return z573 574 def forward(self,575 z: torch.Tensor,576 mask: Optional[torch.Tensor] = None,577 inplace_safe: bool = False,578 _add_with_inplace: bool = False,579 _inplace_chunk_size: Optional[int] = 256580 ) -> torch.Tensor:581 """582 Args:583 x:584 [*, N_res, N_res, C_z] input tensor585 mask:586 [*, N_res, N_res] input mask587 Returns:588 [*, N_res, N_res, C_z] output tensor589 """590 if (inplace_safe):591 x = self._inference_forward(592 z,593 mask,594 _inplace_chunk_size=_inplace_chunk_size,595 with_add=_add_with_inplace,596 )597 return x598 599 if mask is None:600 mask = z.new_ones(z.shape[:-1])601 602 mask = mask.unsqueeze(-1)603 604 z = self.layer_norm_in(z)605 ab = mask606 ab = ab * self.sigmoid(self.linear_ab_g(z))607 ab = ab * self.linear_ab_p(z)608 609 a = ab[..., :self.c_hidden]610 b = ab[..., self.c_hidden:]611 612 # Prevents overflow of torch.matmul in combine projections in613 # reduced-precision modes614 a_std = a.std()615 b_std = b.std()616 if (is_fp16_enabled() and a_std != 0. and b_std != 0.):617 a = a / a.std()618 b = b / b.std()619 620 if (is_fp16_enabled()):621 with torch.cuda.amp.autocast(enabled=False):622 x = self._combine_projections(a.float(), b.float())623 else:624 x = self._combine_projections(a, b)625 626 del a, b627 x = self.layer_norm_out(x)628 x = self.linear_z(x)629 g = self.sigmoid(self.linear_g(z))630 x = x * g631 632 return x633 634 635class FusedTriangleMultiplicationOutgoing(FusedTriangleMultiplicativeUpdate):636 """637 Implements Algorithm 11.638 """639 __init__ = partialmethod(FusedTriangleMultiplicativeUpdate.__init__, _outgoing=True)640 641 642class FusedTriangleMultiplicationIncoming(FusedTriangleMultiplicativeUpdate):643 """644 Implements Algorithm 12.645 """646 __init__ = partialmethod(FusedTriangleMultiplicativeUpdate.__init__, _outgoing=False)647 648 