Team Ai
Modelpublic

OneScience-Group/ESM

sourceHugging Facemitupdated 2mo agoView on Hugging Face
0likes24downloads
fold.py215 linesDownload Raw Back to scripts
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