OneScience-Group/ESM
024
1from pathlib import Path2import sys3 4_PROJECT_FILE = Path(__file__).resolve()5for _PROJECT_ROOT in _PROJECT_FILE.parents:6 if (_PROJECT_ROOT / "model").is_dir():7 if str(_PROJECT_ROOT) not in sys.path:8 sys.path.insert(0, str(_PROJECT_ROOT))9 break10# Copyright (c) Meta Platforms, Inc. and affiliates.11 12# This source code is licensed under the MIT license found in the13# LICENSE file in the root directory of this source tree.14 15 16from pathlib import Path17import sys,os18import argparse19import logging20import sys21import typing as T22from pathlib import Path23from timeit import default_timer as timer24 25import torch26 27import model.esm as esm28from onescience.datapipes.esm import read_fasta29 30logger = logging.getLogger()31logger.setLevel(logging.INFO)32 33formatter = logging.Formatter(34 "%(asctime)s | %(levelname)s | %(name)s | %(message)s",35 datefmt="%y/%m/%d %H:%M:%S",36)37 38console_handler = logging.StreamHandler(sys.stdout)39console_handler.setLevel(logging.INFO)40console_handler.setFormatter(formatter)41logger.addHandler(console_handler)42 43 44PathLike = T.Union[str, Path]45 46 47def enable_cpu_offloading(model):48 from torch.distributed.fsdp import CPUOffload, FullyShardedDataParallel49 from torch.distributed.fsdp.wrap import enable_wrap, wrap50 51 torch.distributed.init_process_group(52 backend="nccl", init_method="tcp://localhost:9999", world_size=1, rank=053 )54 55 wrapper_kwargs = dict(cpu_offload=CPUOffload(offload_params=True))56 57 with enable_wrap(wrapper_cls=FullyShardedDataParallel, **wrapper_kwargs):58 for layer_name, layer in model.layers.named_children():59 wrapped_layer = wrap(layer)60 setattr(model.layers, layer_name, wrapped_layer)61 model = wrap(model)62 63 return model64 65 66def init_model_on_gpu_with_cpu_offloading(model):67 model = model.eval()68 model_esm = enable_cpu_offloading(model.esm)69 del model.esm70 model.cuda()71 model.esm = model_esm72 return model73 74 75def create_batched_sequence_datasest(76 sequences: T.List[T.Tuple[str, str]], max_tokens_per_batch: int = 102477) -> T.Generator[T.Tuple[T.List[str], T.List[str]], None, None]:78 79 batch_headers, batch_sequences, num_tokens = [], [], 080 for header, seq in sequences:81 if (len(seq) + num_tokens > max_tokens_per_batch) and num_tokens > 0:82 yield batch_headers, batch_sequences83 batch_headers, batch_sequences, num_tokens = [], [], 084 batch_headers.append(header)85 batch_sequences.append(seq)86 num_tokens += len(seq)87 88 yield batch_headers, batch_sequences89 90 91def create_parser():92 parser = argparse.ArgumentParser()93 parser.add_argument(94 "-i",95 "--fasta",96 help="Path to input FASTA file",97 type=Path,98 required=True,99 )100 parser.add_argument(101 "-o", "--pdb", help="Path to output PDB directory", type=Path, required=True102 )103 parser.add_argument(104 "-m", "--model-dir", help="Parent path to Pretrained ESM data directory. ", type=Path, default=None105 )106 parser.add_argument(107 "--num-recycles",108 type=int,109 default=None,110 help="Number of recycles to run. Defaults to number used in training (4).",111 )112 parser.add_argument(113 "--max-tokens-per-batch",114 type=int,115 default=1024,116 help="Maximum number of tokens per gpu forward-pass. This will group shorter sequences together "117 "for batched prediction. Lowering this can help with out of memory issues, if these occur on "118 "short sequences.",119 )120 parser.add_argument(121 "--chunk-size",122 type=int,123 default=None,124 help="Chunks axial attention computation to reduce memory usage from O(L^2) to O(L). "125 "Equivalent to running a for loop over chunks of of each dimension. Lower values will "126 "result in lower memory usage at the cost of speed. Recommended values: 128, 64, 32. "127 "Default: None.",128 )129 parser.add_argument("--cpu-only", help="CPU only", action="store_true")130 parser.add_argument("--cpu-offload", help="Enable CPU offloading", action="store_true")131 return parser132 133 134def run(args):135 if not args.fasta.exists():136 raise FileNotFoundError(args.fasta)137 138 args.pdb.mkdir(exist_ok=True)139 140 # Read fasta and sort sequences by length141 logger.info(f"Reading sequences from {args.fasta}")142 all_sequences = sorted(read_fasta(args.fasta), key=lambda header_seq: len(header_seq[1]))143 logger.info(f"Loaded {len(all_sequences)} sequences from {args.fasta}")144 145 logger.info("Loading model")146 147 # Use pre-downloaded ESM weights from model_pth.148 if args.model_dir is not None:149 # if pretrained model path is available150 torch.hub.set_dir(args.model_dir)151 152 model = esm.pretrained.esmfold_v1()153 154 155 model = model.eval()156 model.set_chunk_size(args.chunk_size)157 158 if args.cpu_only:159 model.esm.float() # convert to fp32 as ESM-2 in fp16 is not supported on CPU160 model.cpu()161 elif args.cpu_offload:162 model = init_model_on_gpu_with_cpu_offloading(model)163 else:164 model.cuda()165 logger.info("Starting Predictions")166 batched_sequences = create_batched_sequence_datasest(all_sequences, args.max_tokens_per_batch)167 168 num_completed = 0169 num_sequences = len(all_sequences)170 for headers, sequences in batched_sequences:171 start = timer()172 try:173 output = model.infer(sequences, num_recycles=args.num_recycles)174 except RuntimeError as e:175 if e.args[0].startswith("CUDA out of memory"):176 if len(sequences) > 1:177 logger.info(178 f"Failed (CUDA out of memory) to predict batch of size {len(sequences)}. "179 "Try lowering `--max-tokens-per-batch`."180 )181 else:182 logger.info(183 f"Failed (CUDA out of memory) on sequence {headers[0]} of length {len(sequences[0])}."184 )185 186 continue187 raise188 189 output = {key: value.cpu() for key, value in output.items()}190 pdbs = model.output_to_pdb(output)191 tottime = timer() - start192 time_string = f"{tottime / len(headers):0.1f}s"193 if len(sequences) > 1:194 time_string = time_string + f" (amortized, batch size {len(sequences)})"195 for header, seq, pdb_string, mean_plddt, ptm in zip(196 headers, sequences, pdbs, output["mean_plddt"], output["ptm"]197 ):198 output_file = args.pdb / f"{header}.pdb"199 output_file.write_text(pdb_string)200 num_completed += 1201 logger.info(202 f"Predicted structure for {header} with length {len(seq)}, pLDDT {mean_plddt:0.1f}, "203 f"pTM {ptm:0.3f} in {time_string}. "204 f"{num_completed} / {num_sequences} completed."205 )206 207 208def main():209 parser = create_parser()210 args = parser.parse_args()211 run(args)212 213if __name__ == "__main__":214 main()215 