| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
|
|
| import argparse |
| import os |
| import sys |
|
|
| _PROJECT_ROOT = os.path.abspath(os.path.dirname(__file__)) |
| while _PROJECT_ROOT and not os.path.isdir(os.path.join(_PROJECT_ROOT, "model")): |
| _PARENT = os.path.dirname(_PROJECT_ROOT) |
| if _PARENT == _PROJECT_ROOT: |
| break |
| _PROJECT_ROOT = _PARENT |
| _MODEL_ROOT = os.path.join(_PROJECT_ROOT, "model") |
| _ONESCIENCE_ROOT = os.environ.get("ONESCIENCE_ROOT") |
| for _path in (_MODEL_ROOT, _PROJECT_ROOT): |
| if os.path.exists(_path) and _path not in sys.path: |
| sys.path.insert(0, _path) |
| if _ONESCIENCE_ROOT: |
| _ONESCIENCE_SRC = os.path.join(_ONESCIENCE_ROOT, "src") |
| for _path in (_ONESCIENCE_SRC, _ONESCIENCE_ROOT): |
| if os.path.exists(_path) and _path not in sys.path: |
| sys.path.insert(0, _path) |
| import sys |
| import time |
| from typing import Literal, Optional |
|
|
| import nemo.lightning as nl |
| import torch |
| from megatron.core.inference.common_inference_params import CommonInferenceParams |
| from megatron.core.inference.inference_request import InferenceRequest |
| from nemo.collections.llm import inference |
| from nemo.utils import logging |
|
|
|
|
| CheckpointFormats = Literal["torch_dist", "zarr"] |
| DEFAULT_CKPT_DIR = os.environ.get( |
| "EVO2_CKPT_DIR", |
| os.path.join(_PROJECT_ROOT, "checkpoints", "evo2_nemo_7b"), |
| ) |
|
|
|
|
| def parse_args(): |
| """Parse arguments for Evo2 inference.""" |
| ap = argparse.ArgumentParser() |
|
|
| |
| default_prompt = ( |
| "|d__Bacteria;" |
| + "p__Pseudomonadota;" |
| + "c__Gammaproteobacteria;" |
| + "o__Enterobacterales;" |
| + "f__Enterobacteriaceae;" |
| + "g__Escherichia;" |
| + "s__Escherichia|" |
| ) |
| ap.add_argument( |
| "--prompt", |
| type=str, |
| default=default_prompt, |
| help="Prompt to generate text from Evo2. Defaults to a phylogenetic lineage tag for E coli.", |
| ) |
| ap.add_argument( |
| "--ckpt-dir", |
| type=str, |
| default=DEFAULT_CKPT_DIR, |
| help="Path to checkpoint directory containing pre-trained Evo2 model.", |
| ) |
| ap.add_argument("--temperature", type=float, default=1.0, help="Temperature during sampling for generation.") |
| ap.add_argument("--top-k", type=int, default=0, help="Top K during sampling for generation.") |
| ap.add_argument("--top-p", type=float, default=0.0, help="Top P during sampling for generation.") |
| ap.add_argument("--max-new-tokens", type=int, default=1024, help="Maximum number of tokens to generate.") |
| ap.add_argument("--seed", type=int, default=None, help="Random seed for generation.") |
| |
| ap.add_argument("--tensor-parallel-size", type=int, default=1, help="Order of tensor parallelism. Defaults to 1.") |
| ap.add_argument( |
| "--pipeline-model-parallel-size", type=int, default=1, help="Order of pipeline parallelism. Defaults to 1." |
| ) |
| ap.add_argument( |
| "--context-parallel-size", type=int, default=1, help="Order of context parallelism. Defaults to 1." |
| ) |
| |
| ap.add_argument( |
| "--output-file", |
| type=str, |
| default=None, |
| help="Output file containing the generated text produced by the Evo2 model. If not provided, the output will be logged.", |
| ) |
| |
| ap.add_argument( |
| "--ckpt-format", |
| type=str, |
| choices=["torch_dist", "zarr"], |
| default="torch_dist", |
| help="Specify checkpoint format to use. Defaults to 'torch_dist', as 'zarr' is deprecated.", |
| ) |
| ap.add_argument( |
| "--fp8", |
| action="store_true", |
| default=False, |
| help="Whether to use vortex style FP8. Defaults to False.", |
| ) |
| ap.add_argument( |
| "--flash-decode", |
| action="store_true", |
| default=False, |
| help="Whether to use flash decode. Defaults to True.", |
| ) |
| return ap.parse_args() |
|
|
|
|
| def infer( |
| prompt: str, |
| ckpt_dir: str, |
| temperature: float, |
| top_k: int, |
| top_p: float, |
| max_new_tokens: int, |
| tensor_parallel_size: int, |
| pipeline_model_parallel_size: int, |
| context_parallel_size: int, |
| output_file: Optional[str] = None, |
| ckpt_format: CheckpointFormats = "torch_dist", |
| seed: Optional[int] = None, |
| vortex_style_fp8: bool = False, |
| flash_decode: bool = False, |
| return_log_probs: bool = False, |
| ) -> list[InferenceRequest]: |
| """Inference workflow for Evo2. |
| |
| Args: |
| prompt (str): Prompt to generate text from Evo2. |
| ckpt_dir (str): Path to checkpoint directory containing pre-trained Evo2 model. |
| temperature (float): Temperature during sampling for generation. |
| top_k (int): Top K during sampling for generation. |
| top_p (float): Top P during sampling for generation. |
| max_new_tokens (int): Maximum number of tokens to generate. |
| tensor_parallel_size (int): Order of tensor parallelism. |
| pipeline_model_parallel_size (int): Order of pipeline parallelism. |
| context_parallel_size (int): Order of context parallelism. |
| output_file (str): Output file containing the generated text produced by the Evo2 model. |
| ckpt_format (CheckpointFormats): Checkpoint format to use. |
| seed (int): Random seed for generation. |
| vortex_style_fp8 (bool): Whether to use vortex style FP8. |
| flash_decode (bool): Whether to use flash decode. |
| return_log_probs (bool): Whether to return log probabilities. |
| |
| Returns: |
| None |
| """ |
| model_parallel_size = tensor_parallel_size * pipeline_model_parallel_size * context_parallel_size |
| if model_parallel_size > torch.cuda.device_count(): |
| raise ValueError( |
| f"Requested model parallel size {model_parallel_size} is greater than the " |
| f"number of available CUDA devices {torch.cuda.device_count()}" |
| ) |
| |
| trainer = nl.Trainer( |
| accelerator="gpu", |
| devices=model_parallel_size, |
| strategy=nl.MegatronStrategy( |
| tensor_model_parallel_size=tensor_parallel_size, |
| pipeline_model_parallel_size=pipeline_model_parallel_size, |
| context_parallel_size=context_parallel_size, |
| pipeline_dtype=torch.bfloat16, |
| ckpt_load_optimizer=False, |
| ckpt_save_optimizer=False, |
| ckpt_async_save=False, |
| save_ckpt_format=ckpt_format, |
| ckpt_load_strictness="log_all", |
| ), |
| log_every_n_steps=1, |
| limit_val_batches=10, |
| num_sanity_val_steps=0, |
| plugins=nl.MegatronMixedPrecision( |
| precision="bf16-mixed", |
| params_dtype=torch.bfloat16, |
| ), |
| ) |
| inference_wrapped_model, mcore_tokenizer = inference.setup_model_and_tokenizer( |
| path=ckpt_dir, |
| trainer=trainer, |
| params_dtype=torch.bfloat16, |
| inference_batch_times_seqlen_threshold=8192, |
| inference_max_seq_length=8192, |
| recompute_granularity=None, |
| recompute_num_layers=None, |
| recompute_method=None, |
| vortex_style_fp8=vortex_style_fp8, |
| flash_decode=flash_decode, |
| enable_flash_decode=flash_decode, |
| ) |
| t0 = time.perf_counter_ns() |
| |
| results: list[InferenceRequest] = inference.generate( |
| model=inference_wrapped_model, |
| max_batch_size=1, |
| tokenizer=mcore_tokenizer, |
| prompts=[prompt], |
| random_seed=seed, |
| inference_params=CommonInferenceParams( |
| temperature=temperature, |
| top_k=top_k, |
| top_p=top_p, |
| return_log_probs=return_log_probs, |
| num_tokens_to_generate=max_new_tokens, |
| ), |
| ) |
| dt = (time.perf_counter_ns() - t0) / 1e9 |
| tokens_per_sec = (len(results[0].generated_text) + 1) / dt |
|
|
| print(f"Inference time: {dt} seconds, {tokens_per_sec} tokens/sec", file=sys.stderr) |
| if torch.distributed.get_rank() == 0: |
| if output_file is None: |
| logging.info(results) |
| else: |
| with open(output_file, "w") as f: |
| f.write(f"{results[0]}\n") |
|
|
| return results |
|
|
|
|
| def main(): |
| """Main function for Evo2 inference.""" |
| |
| args = parse_args() |
| if not args.ckpt_dir: |
| raise SystemExit( |
| "ERROR: Evo2 checkpoint path is required. Put the checkpoint under " |
| "checkpoints/evo2_nemo_7b, set EVO2_CKPT_DIR, or pass --ckpt-dir." |
| ) |
| if not os.path.isdir(args.ckpt_dir): |
| raise SystemExit(f"ERROR: Evo2 checkpoint directory does not exist: {args.ckpt_dir}") |
| infer( |
| prompt=args.prompt, |
| ckpt_dir=args.ckpt_dir, |
| temperature=args.temperature, |
| top_k=args.top_k, |
| top_p=args.top_p, |
| max_new_tokens=args.max_new_tokens, |
| tensor_parallel_size=args.tensor_parallel_size, |
| pipeline_model_parallel_size=args.pipeline_model_parallel_size, |
| context_parallel_size=args.context_parallel_size, |
| output_file=args.output_file, |
| ckpt_format=args.ckpt_format, |
| seed=args.seed, |
| vortex_style_fp8=args.fp8, |
| flash_decode=args.flash_decode, |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|