Skip to content

Commit 862bb93

Browse files
committed
Wrap vllm inputs to compatible with VLLM>=0.10.2
1 parent b1d45e3 commit 862bb93

2 files changed

Lines changed: 6 additions & 3 deletions

File tree

pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -98,7 +98,7 @@ nanotron = [
9898
"tensorboardX"
9999
]
100100
tensorboardX = ["tensorboardX"]
101-
vllm = ["vllm>=0.10.0,<0.10.2", "ray", "more_itertools"]
101+
vllm = ["vllm>=0.10.0", "ray", "more_itertools"]
102102
sglang = ["sglang"]
103103
quality = ["ruff>=v0.11.0","pre-commit"]
104104
tests = ["pytest>=7.4.0","deepdiff","pip>=25.2"]

src/lighteval/models/vllm/vllm_model.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@
3030
import torch
3131
from pydantic import NonNegativeFloat, NonNegativeInt, PositiveInt
3232
from tqdm import tqdm
33+
from vllm.inputs.data import TokensPrompt
3334

3435
from lighteval.data import GenerativeTaskDataset, LoglikelihoodDataset
3536
from lighteval.models.abstract_model import LightevalModel, ModelConfig
@@ -415,6 +416,8 @@ def _generate(
415416
generate: bool = True,
416417
) -> list:
417418
"""Contains the actual logic of the generation."""
419+
# Wrap inputs with TokensPrompt to make compatible with VLLM >= 0.10.2
420+
inputs = [TokensPrompt(prompt_token_ids=token_ids) for token_ids in inputs]
418421
sampling_params = SamplingParams(**self.config.generation_parameters.to_vllm_dict())
419422

420423
if generate:
@@ -437,7 +440,7 @@ def _generate(
437440
@ray.remote(num_gpus=self.tensor_parallel_size)
438441
def run_inference_one_model(model_args: dict, sampling_params: SamplingParams, requests):
439442
llm = LLM(**model_args)
440-
return llm.generate(prompt_token_ids=requests, sampling_params=sampling_params)
443+
return llm.generate(requests, sampling_params=sampling_params)
441444

442445
# dispatch requests to all self.data_parallel_size workers, in interleaved fashion
443446
# interleaved important to balance context lengths across workers
@@ -455,7 +458,7 @@ def run_inference_one_model(model_args: dict, sampling_params: SamplingParams, r
455458
]
456459
else:
457460
outputs = self.model.generate(
458-
prompt_token_ids=inputs,
461+
inputs,
459462
sampling_params=sampling_params,
460463
use_tqdm=True,
461464
)

0 commit comments

Comments
 (0)