Team Ai
Modelpublic

OneScience-Group/ESM

sourceHugging Facemitupdated 2mo agoView on Hugging Face
0likes24downloads
triangular_multiplicative_update.py648 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.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