evo2 / scripts /predict.py
OneScience's picture
Upload folder using huggingface_hub
cfbc06c verified
Raw
History Blame Contribute Delete
20.1 kB
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-FileCopyrightText: Copyright (c) 2024 Arc Institute. All rights reserved.
# SPDX-FileCopyrightText: Copyright (c) 2024 Michael Poli. All rights reserved.
# SPDX-FileCopyrightText: Copyright (c) 2024 Stanford University. All rights reserved
# SPDX-License-Identifier: LicenseRef-Apache2
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
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 tempfile
from pathlib import Path
from typing import Literal, Optional
import nemo.lightning as nl
import torch
from lightning.pytorch import LightningDataModule
from megatron.core import parallel_state
from megatron.core.tensor_parallel.mappings import _gather_along_last_dim
from megatron.core.utils import get_batch_on_this_cp_rank
from nemo.collections.llm.gpt.model.base import get_packed_seq_params
from nemo.collections.llm.gpt.model.hyena import HYENA_MODEL_OPTIONS, HyenaModel
from nemo.collections.nlp.modules.common.tokenizer_utils import get_nmt_tokenizer
from nemo.lightning import NeMoLogger
from nemo.lightning.data import WrappedDataLoader
from torch import Tensor
from evo2.data.fasta_dataset import SimpleFastaDataset
from evo2.lightning import LightningPassthroughPredictionMixin
from evo2.utils.callbacks import PredictionWriter
CheckpointFormats = Literal["torch_dist", "zarr"]
DEFAULT_CKPT_DIR = Path(
os.environ.get(
"EVO2_CKPT_DIR",
Path(_PROJECT_ROOT) / "checkpoints" / "evo2_nemo_7b",
)
)
DEFAULT_FASTA_PATH = Path(
os.environ.get(
"EVO2_FASTA_PATH",
Path(_PROJECT_ROOT) / "data" / "predict_example.fa",
)
)
DEFAULT_OUTPUT_DIR = Path(
os.environ.get("EVO2_PREDICT_OUTPUT_DIR", Path(_PROJECT_ROOT) / "outputs" / "predict")
)
def parse_args():
"""Parse arguments for Evo2 inference."""
ap = argparse.ArgumentParser()
ap.add_argument(
"--fasta",
type=Path,
default=DEFAULT_FASTA_PATH,
help="Fasta path from which to generate logit predictions.",
)
ap.add_argument(
"--ckpt-dir",
type=Path,
default=DEFAULT_CKPT_DIR,
help="NeMo2 checkpoint directory for inference.",
)
ap.add_argument(
"--prepend-bos",
action="store_true",
help="Prepend BOS token to sequences. Defaults to False.",
)
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(
"--no-sequence-parallel",
action="store_true",
help="When using TP, skip sequence parallelism. Otherwise sequence parallelism is used whenever tensor "
"parallelism is used. sequence parallelism should save a small amount of GPU memory so it's on"
" by default.",
)
ap.add_argument(
"--batch-size",
type=int,
default=1,
help="Batch size for prediction. Defaults to 1.",
)
ap.add_argument(
"--model-size",
type=str,
default="7b",
choices=sorted(HYENA_MODEL_OPTIONS.keys()),
help="Model size to use. Defaults to '7b'.",
)
# output args:
ap.add_argument(
"--output-dir",
type=Path,
default=DEFAULT_OUTPUT_DIR,
help="Output dir that will contain the generated text produced by the Evo2 model. If not provided, the output will be logged.",
)
ap.add_argument(
"--full-fp8",
action="store_true",
help="Use full FP8 precision (faster but less accurate) rather than vortex style which "
"only applies FP8 to the projection layer of the hyena mixer, when using FP8.",
)
ap.add_argument(
"--fp8", action="store_true", help="Use FP8 precision. Defaults to BF16."
)
# extra:
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(
"--output-log-prob-seqs",
action="store_true",
help="Output log probability of sequences. Defaults to False.",
)
ap.add_argument(
"--log-prob-collapse-option",
choices=["sum", "mean"],
default="mean",
help="How to collapse the log probabilities across the sequence dimension.",
)
ap.add_argument(
"--hybrid-override-pattern",
type=str,
help="Override the hybrid override pattern in the config (specifies hyena layer ordering and type).",
)
ap.add_argument(
"--num-layers",
type=int,
help="If set, override the number of layers specified in the requested config.",
)
return ap.parse_args()
def _gather_along_cp_dim(input_, seq_dim: int = 1):
"""Gather tensors and concatenate along the last dimension."""
world_size = parallel_state.get_context_parallel_world_size()
# Bypass the function if we are using only 1 GPU.
if world_size == 1:
return input_
dim_size = list(input_.size())
dim_size[0] = dim_size[0] * world_size
output = torch.empty(
dim_size, dtype=input_.dtype, device=torch.cuda.current_device()
)
torch.distributed.all_gather_into_tensor(
output,
input_.contiguous(),
group=parallel_state.get_tensor_model_parallel_group(),
)
tensor_list = output.chunk(world_size, dim=0)
output = torch.cat(tensor_list, dim=seq_dim).contiguous()
return output
class HyenaPredictor(LightningPassthroughPredictionMixin, HyenaModel):
"""A predictor for the Hyena model. This adds in the predict step and the passthrough method."""
def __init__(
self,
*args,
output_log_prob_seqs: bool = False,
log_prob_collapse_option: Literal["sum", "mean"] = "mean",
**kwargs,
):
"""Initialize the predictor with our needs around computing log probabilities."""
super().__init__(*args, **kwargs)
self.output_log_prob_seqs = output_log_prob_seqs
self.log_prob_collapse_option = log_prob_collapse_option
def predict_step(self, batch, batch_idx: Optional[int] = None) -> Tensor:
"""Alias for forward_step, also log the pad mask since sequences may not all have the same length."""
if len(batch) == 0:
return
forward_out = self.forward_step(batch)
if not isinstance(forward_out, Tensor):
return forward_out
# Reminder: the model's predictions for input i land at output i+1. To get everything to align, we prepend the
# EOS token to the input sequences and take the outputs for all but the first token.
forward_out_tp_gathered = _gather_along_last_dim(
forward_out, group=parallel_state.get_tensor_model_parallel_group()
)
# else:
# forward_out_tp_gathered = _collect_into_dim(forward_out, dim=-1)
forward_out_gathered = _gather_along_cp_dim(forward_out_tp_gathered)
assert self.tokenizer.vocab_size == forward_out_gathered.shape[-1]
if self.output_log_prob_seqs:
softmax_logprobs = torch.log_softmax(forward_out_gathered, dim=-1)
softmax_logprobs = softmax_logprobs[:, :-1]
input_ids = batch["tokens"][:, 1:]
assert softmax_logprobs.shape[1] == input_ids.shape[1]
logprobs = torch.gather(
softmax_logprobs, # Gather likelihoods...
2, # along the vocab dimension...
input_ids.unsqueeze(-1), # using the token ids to index.
).squeeze(-1)
log_prob_seqs = torch.sum(
logprobs * batch["loss_mask"][:, 1:].float(), dim=-1
)
if self.log_prob_collapse_option == "mean":
log_prob_seqs = log_prob_seqs / (
batch["loss_mask"][:, 1:].float().sum(dim=-1) + 1e-8
)
return {
"log_probs_seqs": log_prob_seqs.cpu(),
"seq_idx": batch["seq_idx"].cpu(),
}
else:
# If the user wants to match back to logits, then they will need to do the offsetting logic themselves.
return {
"token_logits": forward_out_gathered.cpu(),
"pad_mask": batch["loss_mask"].cpu(),
"seq_idx": batch["seq_idx"].cpu(),
}
def hyena_predict_forward_step(model, batch) -> torch.Tensor:
"""Performs a forward step for the Hyena model.
Args:
model: The Hyena model
batch: Dictionary containing input batch data with keys:
- tokens: Input token IDs
- position_ids: Position IDs
- labels: Labels for loss computation
- loss_mask: Mask for loss computation
Returns:
torch.Tensor: Output from the model forward pass
"""
forward_args = {
"input_ids": batch["tokens"],
"position_ids": batch["position_ids"],
# "labels": batch["labels"],
# "loss_mask": batch["loss_mask"],
}
forward_args["attention_mask"] = None
if "cu_seqlens" in batch:
forward_args["packed_seq_params"] = get_packed_seq_params(batch)
return model(**forward_args)
def hyena_predict_data_step(dataloader_iter) -> dict[str, torch.Tensor]:
"""Data step for the Hyena model prediction. Modified from the original gpt data step to include the seq_idx."""
from megatron.core import parallel_state
# Based on: https://github.com/NVIDIA/Megatron-LM/blob/main/pretrain_gpt.py#L87
# https://github.com/NVIDIA/NeMo/blob/main/nemo/collections/nlp/models/language_modeling/megatron_gpt_model.py#L828-L842
batch = next(dataloader_iter)
_batch: dict
if isinstance(batch, tuple) and len(batch) == 3:
_batch = batch[0]
else:
_batch = batch
required_device_keys = set()
required_host_keys = set()
required_device_keys.add("attention_mask")
if "cu_seqlens" in _batch:
required_device_keys.add("cu_seqlens")
required_host_keys.add("cu_seqlens_argmin")
required_host_keys.add("max_seqlen")
if parallel_state.is_pipeline_first_stage():
required_device_keys.update(("tokens", "position_ids"))
if parallel_state.is_pipeline_last_stage():
required_device_keys.update(("labels", "loss_mask", "seq_idx"))
_batch_required_keys = {}
for key, val in _batch.items():
if key in required_device_keys:
_batch_required_keys[key] = val.cuda(non_blocking=True)
elif key in required_host_keys:
_batch_required_keys[key] = val.cpu()
else:
_batch_required_keys[key] = None
# slice batch along sequence dimension for context parallelism
output = get_batch_on_this_cp_rank(_batch_required_keys)
return output
class PredictDataModule(LightningDataModule):
"""Create a dataloader for prediction."""
def __init__(self, dataset: torch.utils.data.Dataset, batch_size: int = 1):
"""Create a dataloader for prediction."""
super().__init__()
self.dataset = dataset
self.batch_size = batch_size
def setup(self, stage: Optional[str] = None) -> None:
"""Set up the dataloader."""
def predict_dataloader(self):
"""Create a dataloader for prediction."""
# need to use this to communicate that we are in predict mode and safe to not drop last batch
return WrappedDataLoader(
mode="predict",
dataset=self.dataset,
batch_size=self.batch_size,
num_workers=8,
shuffle=False,
drop_last=False,
)
def predict(
fasta_path: Path,
ckpt_dir: str,
output_dir: Path,
tensor_parallel_size: int,
pipeline_model_parallel_size: int,
context_parallel_size: int,
model_size: str = "7b",
ckpt_format: CheckpointFormats = "torch_dist",
fp8: bool = False,
full_fp8: bool = False,
work_dir: Path | None = None,
batch_size: int = 1,
output_log_prob_seqs: bool = False,
log_prob_collapse_option: Literal["sum", "mean"] = "mean",
prepend_bos: bool = False,
no_sequence_parallel: bool = False,
hybrid_override_pattern: str | None = None,
num_layers: int | None = None,
):
"""Inference workflow for Evo2.
Returns:
None
"""
if work_dir is None:
work_dir = Path(tempfile.mkdtemp())
sequence_parallel = tensor_parallel_size > 1 and not no_sequence_parallel
output_dir.mkdir(
parents=True, exist_ok=True
) # Make sure the output directory exists, files will be written here.
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()}"
)
# Create PTL trainer.
trainer = nl.Trainer(
accelerator="gpu",
devices=model_parallel_size,
strategy=nl.MegatronStrategy(
drop_last_batch=False,
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, # Needs to be false for a normal model checkpoint.
ckpt_save_optimizer=False,
ckpt_async_save=False,
sequence_parallel=tensor_parallel_size > 1 and sequence_parallel,
save_ckpt_format=ckpt_format,
ckpt_load_strictness="log_all",
data_sampler=nl.MegatronDataSampler(
micro_batch_size=batch_size,
global_batch_size=batch_size,
seq_len=8192,
output_log=False, # this is needed for predict step to work
),
),
log_every_n_steps=1,
limit_val_batches=10,
num_sanity_val_steps=0,
callbacks=[
PredictionWriter(
output_dir=output_dir,
write_interval="epoch",
batch_dim_key_defaults={"token_logits": 0},
seq_dim_key_defaults={"token_logits": 1},
)
],
plugins=nl.MegatronMixedPrecision(
precision="bf16-mixed",
params_dtype=torch.bfloat16,
# Only use FP8 in this plugin when using full FP8 precision and FP8.
# Otherwise use vortex_style_fp8 in the model config.
fp8="hybrid" if fp8 and full_fp8 else None,
fp8_amax_history_len=16 if fp8 and full_fp8 else 1,
fp8_amax_compute_algo="max" if fp8 and full_fp8 else "most_recent",
),
)
# The following two config options are really only used for testing, but may also be useful for getting output from
# specific layers of the model.
config_modifiers_init = {}
if hybrid_override_pattern is not None:
config_modifiers_init["hybrid_override_pattern"] = hybrid_override_pattern
if num_layers is not None:
config_modifiers_init["num_layers"] = num_layers
config = HYENA_MODEL_OPTIONS[model_size](
forward_step_fn=hyena_predict_forward_step,
data_step_fn=hyena_predict_data_step, # , attention_backend=AttnBackend.fused,
distribute_saved_activations=(
False if sequence_parallel and tensor_parallel_size > 1 else True
),
# Only use vortex style FP8 in the model config if using FP8 and not full FP8. This will only apply FP8 to
# the projection layer of the hyena mixer.
vortex_style_fp8=fp8 and not full_fp8,
**config_modifiers_init,
)
trainer.strategy._setup_optimizers = False
nemo_logger = NeMoLogger(log_dir=work_dir)
nemo_logger.setup(trainer, resume_if_exists=True)
resume = nl.AutoResume(
resume_if_exists=True,
resume_ignore_no_checkpoint=False,
resume_past_end=False,
restore_config=nl.RestoreConfig(
path=str(ckpt_dir), # NeMo expects a string path.
load_model_state=True,
load_optim_state=False,
load_artifacts=False,
),
)
tokenizer = get_nmt_tokenizer("byte-level")
model = HyenaPredictor(
config,
tokenizer=tokenizer,
output_log_prob_seqs=output_log_prob_seqs,
log_prob_collapse_option=log_prob_collapse_option,
)
resume.setup(trainer, model) # this pulls weights from the starting checkpoint.
dataset = SimpleFastaDataset(fasta_path, tokenizer, prepend_bos=prepend_bos)
datamodule = PredictDataModule(dataset, batch_size=batch_size)
trainer.predict(model, datamodule=datamodule)
dataset.write_idx_map(
output_dir
) # Finally write out the index map so we can match the predictions to the original sequences.
print(f"Prediction completed. Result files are in: {output_dir}")
def main():
"""Entrypoint for Evo2 prediction (single inference step, no new tokens)."""
args = parse_args()
if args.ckpt_dir is None or str(args.ckpt_dir) in ("", "."):
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 args.ckpt_dir.is_dir():
raise SystemExit(f"ERROR: Evo2 checkpoint directory does not exist: {args.ckpt_dir}")
predict(
fasta_path=args.fasta,
ckpt_dir=args.ckpt_dir,
tensor_parallel_size=args.tensor_parallel_size,
pipeline_model_parallel_size=args.pipeline_model_parallel_size,
context_parallel_size=args.context_parallel_size,
output_dir=args.output_dir,
model_size=args.model_size,
ckpt_format=args.ckpt_format,
fp8=args.fp8,
full_fp8=args.full_fp8,
batch_size=args.batch_size,
output_log_prob_seqs=args.output_log_prob_seqs,
log_prob_collapse_option=args.log_prob_collapse_option,
prepend_bos=args.prepend_bos,
no_sequence_parallel=args.no_sequence_parallel,
hybrid_override_pattern=args.hybrid_override_pattern,
num_layers=args.num_layers,
)
if __name__ == "__main__":
main()