Team Ai
Modelpublic

OneScience-Group/ESM

sourceHugging Facemitupdated 2mo agoView on Hugging Face
0likes24downloads
torchscript.py216 linesDownload Raw Back to openfold
1# Copyright 2021 AlQuraishi Laboratory2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7#      http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14 15from typing import Optional, Sequence, Tuple16 17import torch18import torch.nn as nn19 20from model.openfold.dropout import (21    DropoutRowwise,22    DropoutColumnwise,23)24from model.openfold.evoformer import (25    EvoformerBlock,26    EvoformerStack,27)28from model.openfold.outer_product_mean import OuterProductMean29from model.openfold.msa import (30    MSARowAttentionWithPairBias, 31    MSAColumnAttention,32    MSAColumnGlobalAttention,33)34from model.openfold.pair_transition import PairTransition35from model.openfold.primitives import Attention, GlobalAttention36from model.openfold.structure_module import (37    InvariantPointAttention,38    BackboneUpdate,39)40from model.openfold.template import TemplatePairStackBlock41from model.openfold.triangular_attention import (42    TriangleAttentionStartingNode,43    TriangleAttentionEndingNode,44)45from model.openfold.triangular_multiplicative_update import (46    TriangleMultiplicationOutgoing,47    TriangleMultiplicationIncoming,48)49 50 51def script_preset_(model: torch.nn.Module):52    """53    TorchScript a handful of low-level but frequently used submodule types 54    that are known to be scriptable.55 56    Args:57        model: 58            A torch.nn.Module. It should contain at least some modules from 59            this repository, or this function won't do anything.60    """61    script_submodules_(62        model, 63        [64            nn.Dropout,65            Attention,66            GlobalAttention,67            EvoformerBlock,68            #TemplatePairStackBlock,69        ], 70        attempt_trace=False,71        batch_dims=None,72    ) 73 74    75def _get_module_device(module: torch.nn.Module) -> torch.device:76    """77    Fetches the device of a module, assuming that all of the module's78    parameters reside on a single device79 80    Args:81        module: A torch.nn.Module82    Returns:83        The module's device84    """85    return next(module.parameters()).device86 87 88def _trace_module(module, batch_dims=None):89    if(batch_dims is None):90        batch_dims = ()91 92    # Stand-in values93    n_seq = 1094    n_res = 1095 96    device = _get_module_device(module)97 98    def msa(channel_dim):99        return torch.rand(100            (*batch_dims, n_seq, n_res, channel_dim),101            device=device,102        )103 104    def pair(channel_dim):105        return torch.rand(106            (*batch_dims, n_res, n_res, channel_dim),107            device=device,108        )109 110    if(isinstance(module, MSARowAttentionWithPairBias)):111        inputs = {112            "forward": (113                msa(module.c_in), # m114                pair(module.c_z), # z115                torch.randint(116                    0, 2, 117                    (*batch_dims, n_seq, n_res)118                ), # mask119            ),120        }121    elif(isinstance(module, MSAColumnAttention)):122        inputs = {123            "forward": (124                msa(module.c_in), # m125                torch.randint(126                    0, 2, 127                    (*batch_dims, n_seq, n_res)128                ), # mask129            ),130        }131    elif(isinstance(module, OuterProductMean)):132        inputs = {133            "forward": (134                msa(module.c_m),135                torch.randint(136                    0, 2,137                    (*batch_dims, n_seq, n_res)138                )139            )140        }141    else:142        raise TypeError(143            f"tracing is not supported for modules of type {type(module)}"144        )145 146    return torch.jit.trace_module(module, inputs)147 148 149def _script_submodules_helper_(150    model,151    types,152    attempt_trace,153    to_trace,154):155    for name, child in model.named_children():156        if(types is None or any(isinstance(child, t) for t in types)):157            try:158                scripted = torch.jit.script(child)159                setattr(model, name, scripted)160                continue161            except (RuntimeError, torch.jit.frontend.NotSupportedError) as e:162                if(attempt_trace):163                    to_trace.add(type(child))164                else:165                    raise e166        167        _script_submodules_helper_(child, types, attempt_trace, to_trace)168 169 170def _trace_submodules_(171    model,172    types,173    batch_dims=None,174):175    for name, child in model.named_children():176        if(any(isinstance(child, t) for t in types)):177            traced = _trace_module(child, batch_dims=batch_dims)178            setattr(model, name, traced)179        else:180            _trace_submodules_(child, types, batch_dims=batch_dims)181 182 183def script_submodules_(184    model: nn.Module,185    types: Optional[Sequence[type]] = None,186    attempt_trace: Optional[bool] = True,187    batch_dims: Optional[Tuple[int]] = None,188):189    """190    Convert all submodules whose types match one of those in the input 191    list to recursively scripted equivalents in place. To script the entire192    model, just call torch.jit.script on it directly.193 194    When types is None, all submodules are scripted.195 196    Args:197        model: 198            A torch.nn.Module199        types: 200            A list of types of submodules to script201        attempt_trace: 202            Whether to attempt to trace specified modules if scripting 203            fails. Recall that tracing eliminates all conditional 204            logic---with great tracing comes the mild responsibility of 205            having to remember to ensure that the modules in question 206            perform the same computations no matter what.207    """208    to_trace = set()209 210    # Aggressively script as much as possible first...211    _script_submodules_helper_(model, types, attempt_trace, to_trace)212  213    # ... and then trace stragglers.214    if(attempt_trace and len(to_trace) > 0):215        _trace_submodules_(model, to_trace, batch_dims=batch_dims)216