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