Use lm-eval 0.4.12's built-in trtllm backend, deprecate lm_eval_tensorrt_llm.py (#2066)

### What does this PR do?

Type of change: documentation / example update (with a behaviour fix)

lm-evaluation-harness **0.4.12** is the first release that ships a
TensorRT-LLM backend
(`lm_eval.models.trtllm_causallms`, registered as `trtllm`) — it is
absent in 0.4.10 and
0.4.11. This example no longer maintains its own, so:

- Pin `lm_eval[api,ifeval]>=0.4.12,<0.5` (the 0.5.0.dev line drops the
file) and bump
  `lm_eval_hf.py`'s version guard to match.
- **Delete** `examples/llm_eval/lm_eval_tensorrt_llm.py` (the `trt-llm`
model). Replace
`python lm_eval_tensorrt_llm.py --model trt-llm --model_args
tokenizer=<tok>,checkpoint_dir=<ckpt>`
with `python lm_eval_trtllm.py --model trtllm --model_args
model=<ckpt>,tokenizer=<tok>`.
- Add `examples/llm_eval/lm_eval_trtllm.py`, whose entire content is one
corrected
`_parse_logprobs` plus `cli_evaluate()` (see below). `lm_eval_hf.py`
stays HF-only.
- `examples/hf_ptq/scripts/huggingface_example.sh` and the docs use the
upstream backend.
`parser.sh` gains `--input` (`BUILD_MAX_INPUT_LEN`, default 4096) — it
already *echoed*
that variable but never parsed or defaulted it, so it printed empty on
every run.

#### Why `lm_eval_trtllm.py` exists: an upstream off-by-one

TensorRT-LLM aligns `prompt_logprobs` to the *next* token.
`executor/base_worker.py`:

```python
# Pass prompt_token_ids with an offset of 1 for correct mapping to the context logits
prompt_token_ids = generation_result._generation_request.prompt_token_ids[1:] + first_generation_token
```

So entry `i` is the distribution that predicted `tokens[i + 1]`, and
`_topk_logprobs`
appends that token's id when it is not in the top-k. lm-eval's
`_parse_logprobs` instead
reads `prompt_logprobs[i][tokens[i]]` and applies its own shift on top,
which raises
`KeyError` on the **first request of every loglikelihood task**
(hellaswag, mmlu, arc, ...):

```
File ".../lm_eval/models/trtllm_causallms.py", line 324, in _parse_logprobs
    current_token_logprob = prompt_logprob[tokens[i]]
KeyError: 6503
```

Probed against TRT-LLM 1.3.0rc23 with a 14-token prompt for
`prompt_logprobs` 0, 1 and 2:
`tokens[i]` is missing at **every** position, `tokens[i+1]` is present
at every position.
Only `generate_until` tasks work unpatched. **This wants an upstream
issue against
EleutherAI/lm-evaluation-harness.**

The override also fails loudly rather than quietly: it checks
`prompt_logprobs` covers
every prompt token and raises on a missing token, instead of skipping
the term and
silently inflating the reported accuracy.

#### Defaults that must be set explicitly

`TRTLLM.__init__` accepts `**kwargs` but forwards only a fixed set to
the **TensorRT-LLM
`LLM` API**, so extra `--model_args` aimed at the engine are silently
dropped. (lm-eval's
own named parameters — `max_gen_toks`, `batch_size`, `truncation_side`,
... — are honored
normally.) Two engine defaults are unsafe for few-shot eval:

- `tensor_parallel_size` defaults to **1** (the deleted wrapper used
every visible GPU).
- `max_input_len` defaults to **2048**, and longer prompts are silently
left-truncated —
  5-shot MMLU/gsm8k prompts exceed that.

### Usage

```bash
python lm_eval_trtllm.py --model trtllm \
    --model_args model=<quantized checkpoint dir>,tokenizer=<HF model folder>,tensor_parallel_size=<tp>,max_batch_size=<bs>,max_input_len=4096,max_output_len=512 \
    --tasks hellaswag,gsm8k \
    --batch_size <bs>
```

Flat arguments (no `run` subcommand) are what 0.4.12's
`HarnessCLI.parse_args` inserts
`run` for automatically (`_cli/harness.py:48-51`); this is the exact
command form used for
the results below.

### Testing

**Unit** — `tests/examples/llm_eval/test_lm_eval_trtllm.py`, no GPU and
no `tensorrt_llm`
install: stubs the response object and pins the `i-1` alignment, the
`rank != 1` →
`is_greedy` rule, the `ctxlen=0` edge, and both `RuntimeError` paths.
Mutation-checked —
dropping the `-1` shift is caught by 5/5 cases, ignoring `ctxlen` by
4/5. A sixth test is a
**tripwire**: it asserts lm-eval's own implementation is still
misaligned, so a future
0.4.x that fixes the bug fails the test and says to delete this file
rather than being
silently re-broken by the override.

**End to end** — `nvidia/Qwen3.5-122B-A10B-NVFP4` (NVFP4 MoE, 256
experts) on **4x B300**,
TRT-LLM 1.3.0rc23, lm-eval 0.4.12, `--limit 32`:

| run | hellaswag acc | hellaswag acc_norm | gsm8k flexible | gsm8k
strict |
|---|---|---|---|---|
| deleted impl (`trt-llm`), tp=4 | 0.7188 | 0.7812 | 0.8438 | 0.7812 |
| `lm_eval_trtllm.py`, tp=1 | 0.7188 | 0.7812 | 0.8438 | 0.8125 |
| `lm_eval_trtllm.py`, tp=2 | 0.7188 | 0.7812 | 0.9062 | 0.8125 |
| `lm_eval_trtllm.py`, tp=4 | 0.7188 | 0.7812 | 0.8750 | 0.8438 |

- hellaswag (the loglikelihood path this PR fixes) is **identical at
every tp and identical
to the deleted implementation** — the alignment fix is exact, not
approximate.
- gsm8k varies by 1–2 samples out of 32 (generation path: upstream uses
native `stop=`
sequences and per-request `SamplingParams`; the old wrapper used
beam-search-of-1 with
  post-hoc string truncation).
- Without the override, every hellaswag run above dies with the
`KeyError`.
- Re-verified at tp=4 after the code moved out of `lm_eval_hf.py` into
`lm_eval_trtllm.py`.

Note: NVFP4 fused-MoE has no CUTLASS tactic on Hopper (`No supported MoE
GEMM tactic
remains after replacing unsupported NO_SMEM epilogues.`), so this had to
be validated on
Blackwell.

### Feature parity notes

Gained from upstream: `loglikelihood_rolling` (was
`NotImplementedError`), pipeline
parallelism, `add_bos_token` auto-detection, prompt truncation,
per-request sampling params,
`prompt_logprobs` instead of full-vocab context logits (much lower
memory), thinking-tag
handling, `batch_size=auto`.

Not reachable through the upstream backend (were set by
`modelopt.deploy.llm.LLM`):
`enable_attention_dp` for MoE, `CudaGraphConfig`,
`enable_chunked_prefill`,
`moe_expert_parallel_size=1`, and `free_gpu_memory_fraction=0.7` with a
capped
`kv_cache.max_tokens` — upstream uses the TRT-LLM default 0.9 (observed
allocating 218 GiB
of paged KV cache on B300), so OOM risk is higher on smaller GPUs. This
is documented in
`examples/llm_eval/README.md`, and `huggingface_example.sh` honours a
preset `LM_EVAL_TP`
so users can lower the tensor-parallel size without editing the script.

### Before your PR is "*Ready for review*"

- Is this change backward compatible?: ❌ — `lm_eval_tensorrt_llm.py` is
removed and the CLI changes (`--model trt-llm` → `trtllm`,
`checkpoint_dir=` → `model=`). Migration command is in the README and
CHANGELOG.
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ — no new
dependency; existing `lm_eval` pin tightened.
- Did you write any new necessary tests?: ✅ —
`tests/examples/llm_eval/test_lm_eval_trtllm.py` (6 cases, no GPU).
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ — under 0.47 *Deprecations*.
- Did you get Claude approval on this PR?: ✅ — reviewed, feedback
addressed in `dcedd37b4` and `622b97c26`.

🤖 Generated with [Claude Code](https://claude.com/claude-code)


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **New Features**
* Added TensorRT-LLM evaluation through lm-evaluation-harness’s `trtllm`
backend.
* Added configurable input/output lengths, batching, tensor parallelism,
and build input length.
* Improved prompt log-probability alignment for more accurate evaluation
results.

* **Documentation**
* Updated evaluation instructions, truncation guidance, backend
limitations, and configuration examples.

* **Deprecations**
  * Removed the legacy TensorRT-LLM evaluation script and entry point.

* **Updates**
  * lm-evaluation-harness now requires versions 0.4.12 through 0.4.x.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Chenjie Luo <chenjiel@nvidia.com>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
Chenjie Luo
2026-08-07 22:36:55 +00:00
committed by GitHub
co-authored by Claude Opus 5
parent bd3798a794
commit 9b8caf623a
11 changed files with 412 additions and 230 deletions
+11 -4
View File
@@ -297,11 +297,18 @@ if [[ $TASKS =~ "lm_eval" ]]; then
pip install -r requirements.txt
echo "Using the following config: max output $BUILD_MAX_OUTPUT_LEN max batch $BUILD_MAX_BATCH_SIZE"
# lm-eval's `trtllm` backend defaults to 1 GPU; shard over every visible one instead.
# Override LM_EVAL_TP to lower it -- TRT-LLM enables expert parallelism at higher TP,
# which fails in DeepEP kernels for MoE checkpoints on some GPUs (e.g. SM 12.0).
LM_EVAL_TP=${LM_EVAL_TP:-$(python -c "import torch; print(max(torch.cuda.device_count(), 1))")}
python lm_eval_tensorrt_llm.py \
--model trt-llm \
--model_args tokenizer=$MODEL_PATH,checkpoint_dir=$SAVE_PATH,max_gen_toks=$BUILD_MAX_OUTPUT_LEN \
echo "Using the following config: max input $BUILD_MAX_INPUT_LEN max output $BUILD_MAX_OUTPUT_LEN max batch $BUILD_MAX_BATCH_SIZE tp $LM_EVAL_TP"
# max_input_len defaults to 2048, which silently truncates 5-shot prompts, so pass it
# explicitly; the engine's max_seq_len is max_input_len + max_output_len.
python lm_eval_trtllm.py \
--model trtllm \
--model_args "model=$SAVE_PATH,tokenizer=$MODEL_ABS_PATH,tensor_parallel_size=$LM_EVAL_TP,max_batch_size=$BUILD_MAX_BATCH_SIZE,max_gen_toks=$BUILD_MAX_OUTPUT_LEN,max_input_len=$BUILD_MAX_INPUT_LEN,max_output_len=$BUILD_MAX_OUTPUT_LEN" \
--tasks $LM_EVAL_TASKS \
--batch_size $BUILD_MAX_BATCH_SIZE $lm_eval_flags | tee $LM_EVAL_RESULT
+6 -1
View File
@@ -41,7 +41,7 @@ parse_options() {
CALIB_WITH_IMAGES=false
# Parse command-line options
ARGS=$(getopt -o "" -l "model:,quant:,recipe:,kv_cache_quant:,tp:,pp:,sparsity:,awq_block_size:,calib:,calib_batch_size:,output:,batch:,tasks:,lm_eval_tasks:,lm_eval_limit:,simple_eval_tasks:,simple_eval_limit:,mmlu_limit:,trust_remote_code,use_seq_device_map,gpu_max_mem_percentage:,kv_cache_free_gpu_memory_fraction:,low_memory_mode,no-verbose,calib_dataset:,calib_seq:,auto_quantize_checkpoint:,auto_quantize_bits:,auto_quantize_method:,auto_quantize_score_size:,auto_quantize_cost_model:,auto_quantize_active_moe_expert_ratio:,moe_calib_experts_ratio:,cast_mxfp4_to_nvfp4,vlm,calib_with_images" -n "$0" -- "$@")
ARGS=$(getopt -o "" -l "model:,quant:,recipe:,kv_cache_quant:,tp:,pp:,sparsity:,awq_block_size:,calib:,calib_batch_size:,input:,output:,batch:,tasks:,lm_eval_tasks:,lm_eval_limit:,simple_eval_tasks:,simple_eval_limit:,mmlu_limit:,trust_remote_code,use_seq_device_map,gpu_max_mem_percentage:,kv_cache_free_gpu_memory_fraction:,low_memory_mode,no-verbose,calib_dataset:,calib_seq:,auto_quantize_checkpoint:,auto_quantize_bits:,auto_quantize_method:,auto_quantize_score_size:,auto_quantize_cost_model:,auto_quantize_active_moe_expert_ratio:,moe_calib_experts_ratio:,cast_mxfp4_to_nvfp4,vlm,calib_with_images" -n "$0" -- "$@")
eval set -- "$ARGS"
while true; do
@@ -56,6 +56,7 @@ parse_options() {
--awq_block_size ) AWQ_BLOCK_SIZE="$2"; shift 2;;
--calib ) CALIB_SIZE="$2"; shift 2;;
--calib_batch_size ) CALIB_BATCH_SIZE="$2"; shift 2;;
--input ) BUILD_MAX_INPUT_LEN="$2"; shift 2;;
--output ) BUILD_MAX_OUTPUT_LEN="$2"; shift 2;;
--batch ) BUILD_MAX_BATCH_SIZE="$2"; shift 2;;
--tasks ) TASKS="$2"; shift 2;;
@@ -90,6 +91,7 @@ parse_options() {
DEFAULT_CALIB_SIZE=512
DEFAULT_CALIB_SEQ=512
DEFAULT_CALIB_BATCH_SIZE=0
DEFAULT_BUILD_MAX_INPUT_LEN=4096
DEFAULT_BUILD_MAX_OUTPUT_LEN=1024
DEFAULT_BUILD_MAX_BATCH_SIZE=2
@@ -102,6 +104,9 @@ parse_options() {
if [ -z "$CALIB_BATCH_SIZE" ]; then
CALIB_BATCH_SIZE=$DEFAULT_CALIB_BATCH_SIZE
fi
if [ -z "$BUILD_MAX_INPUT_LEN" ]; then
BUILD_MAX_INPUT_LEN=$DEFAULT_BUILD_MAX_INPUT_LEN
fi
if [ -z "$BUILD_MAX_OUTPUT_LEN" ]; then
BUILD_MAX_OUTPUT_LEN=$DEFAULT_BUILD_MAX_OUTPUT_LEN
fi
+33 -1
View File
@@ -109,10 +109,42 @@ If `trust_remote_code` needs to be true, please append the command with the `--t
### TensorRT-LLM
Uses the `trtllm` backend built into lm-eval (>= 0.4.12), which loads the quantized
checkpoint directly with the TensorRT-LLM LLM API.
```sh
python lm_eval_tensorrt_llm.py --model trt-llm --model_args tokenizer=<HF model folder>,checkpoint_dir=<Quantized checkpoint dir> --tasks <comma separated tasks> --batch_size <max batch size>
python lm_eval_trtllm.py --model trtllm \
--model_args model=<Quantized checkpoint dir>,tokenizer=<HF model folder>,tensor_parallel_size=<tp>,max_batch_size=<max batch size>,max_input_len=4096,max_output_len=512 \
--tasks <comma separated tasks> \
--batch_size <max batch size>
```
> **_NOTE:_** Loglikelihood tasks (mmlu, hellaswag, arc, ...) need **TensorRT-LLM >=
> 1.3.0rc11**, which is when the engine started returning the requested token in every
> `prompt_logprobs` entry. Earlier releases return only the top-1 token per position, so a
> continuation token's logprob cannot be recovered and the run aborts with a clear error.
> Generative tasks (gsm8k, ifeval) are unaffected.
> **_NOTE:_** Set `max_input_len` and `max_output_len` explicitly. They default to 2048 and
> 512, and prompts longer than `max_input_len` are silently truncated — 5-shot MMLU or
> gsm8k prompts exceed 2048 tokens. `max_seq_len` of the engine is their sum.
> **_NOTE:_** `tensor_parallel_size` defaults to 1; set it to the number of GPUs the
> checkpoint needs. `pipeline_parallel_size` is also supported.
> **_NOTE:_** Use `lm_eval_trtllm.py` rather than the plain `lm_eval` CLI. lm-eval 0.4.12's
> `trtllm` backend misaligns TensorRT-LLM's `prompt_logprobs` by one position, so every
> loglikelihood task (hellaswag, mmlu, arc, ...) fails with a `KeyError`;
> `lm_eval_trtllm.py` overrides the alignment. It goes away once the fix lands upstream.
> **_NOTE:_** The backend forwards only a fixed set of arguments to TensorRT-LLM, so the
> tuning the old `lm_eval_tensorrt_llm.py` applied is not reachable: expert parallelism is
> left at the TensorRT-LLM default (MoE checkpoints can fail in DeepEP kernels on some
> GPUs, e.g. SM 12.0) and the KV cache uses 90% of free GPU memory rather than 70%. Lower
> `tensor_parallel_size` if you hit either.
`lm_eval_tensorrt_llm.py` (`--model trt-llm`) has been removed; use the command above.
## MMLU
[Massive Multitask Language Understanding](https://arxiv.org/abs/2009.03300). A score (0-1, higher is better) will be printed at the end of the benchmark.
+3 -2
View File
@@ -48,8 +48,9 @@ import datasets
from lm_eval import utils
from packaging.version import Version
if Version(version("lm_eval")) < Version("0.4.10"):
raise ImportError(f"lm_eval_hf.py requires lm-eval >= 0.4.10; found {version('lm_eval')}.")
if Version(version("lm_eval")) < Version("0.4.12"):
# Matches the floor pinned in requirements.txt.
raise ImportError(f"lm_eval_hf.py requires lm-eval >= 0.4.12; found {version('lm_eval')}.")
from lm_eval._cli import HarnessCLI
from lm_eval.api.model import T
-213
View File
@@ -1,213 +0,0 @@
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# 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 copy
import gc
import logging
import os
import signal
import threading
import time
from collections.abc import Iterable
from typing import Any
import torch
import torch.nn.functional as F
from lm_eval.__main__ import cli_evaluate
from lm_eval.api.registry import register_model
from lm_eval.models.api_models import TemplateAPI
from transformers import BatchEncoding
from modelopt.deploy.llm import LLM
logger = logging.getLogger(__name__)
TokenSequence = list[int] | torch.LongTensor | torch.Tensor | BatchEncoding
@register_model("trt-llm")
class TRTLLM(TemplateAPI):
def __init__(
self,
tokenizer: str,
checkpoint_dir: str,
batch_size: int = 1,
**kwargs,
):
assert isinstance(tokenizer, str)
super().__init__(
tokenizer=tokenizer,
batch_size=int(batch_size),
**kwargs,
)
if self.tokenizer.pad_token_id is None:
self.tokenizer.pad_token_id = self.tokenizer.eos_token_id
assert isinstance(checkpoint_dir, str)
max_length = kwargs.get("max_length", self._max_gen_toks + 4096)
self.llm = LLM(
checkpoint_dir=checkpoint_dir,
tokenizer=self.tokenizer,
max_batch_size=int(batch_size),
max_seq_len=max_length,
# Loglikelihood tasks request context logits. KV cache prefix reuse would return
# logits only for the recomputed suffix on shared-prefix requests (e.g. hellaswag),
# truncating context_logits and breaking parse_logprobs. Disable it.
enable_kv_cache_reuse=False,
trust_remote_code=bool(kwargs.get("trust_remote_code", False)),
)
self.max_length = max_length - 1
logger.info("Loaded TRT-LLM")
def model_call(
self,
messages: Iterable[list[int]],
*,
generate: bool = True,
gen_kwargs: dict | None = None,
**kwargs,
):
# !!! Copy: shared dict for each request, need new object !!!
gen_kwargs = copy.deepcopy(gen_kwargs)
assert isinstance(messages, Iterable), "Expect the messages to be Iterable[list[int]]"
first_element = next(iter(messages))
assert isinstance(first_element, list) and isinstance(first_element[0], int), (
"Expect the messages to be Iterable[list[int]]"
)
if not generate:
return self.llm.generate_context_logits(prompts=messages)
llm_kwargs = {}
max_new_tokens = self._max_gen_toks
stop_words = []
if gen_kwargs:
if "until" in gen_kwargs:
stop_words = gen_kwargs.pop("until")
llm_kwargs["stop_words"] = stop_words
if "temperature" in gen_kwargs:
llm_kwargs["temperature"] = gen_kwargs.pop("temperature")
if "top_p" in gen_kwargs:
llm_kwargs["top_p"] = gen_kwargs.pop("top_p")
if "max_gen_toks" in gen_kwargs:
max_new_tokens = gen_kwargs.pop("max_gen_toks")
output_texts: list[str] = self.llm.generate_text(
prompts=messages,
max_new_tokens=max_new_tokens,
**llm_kwargs,
)
# Manually filter out keyword if not supported by llm.
for i, text in enumerate(output_texts):
for word in stop_words:
word_index = text.find(word)
if word_index >= 0:
text = text[:word_index]
output_texts[i] = text
return output_texts
async def amodel_call(
self,
session,
messages: Iterable[list[int]],
*,
generate: bool = True,
cache_keys: list | None = None,
ctxlens: list[int] | None = None,
gen_kwargs: dict | None = None,
**kwargs,
):
raise NotImplementedError
def loglikelihood_rolling(self, requests):
raise NotImplementedError
def _create_payload(
self,
messages: list[list[int]] | list[dict] | list[str] | str,
*,
generate: bool = True,
gen_kwargs: dict | None = None,
seed: int = 1234,
**kwargs,
) -> dict:
"""This method is responsible for creating the json payload that will be sent to the API."""
raise NotImplementedError
@staticmethod
def parse_generations(outputs: Any | list[Any], **kwargs) -> list[str]:
"""Method used to parse the generations from the (batched) API response."""
return outputs
@staticmethod
def parse_logprobs(
outputs: Any | list[Any],
tokens: list[list[int]] | None = None,
ctxlens: list[int] | None = None,
**kwargs,
) -> list[tuple[float, bool]]:
"""Method used to parse the logprobs from the (batched) API response.
The provided tokens have two parts: The context tokens (length as ctxlens) and the continuation tokens.
The logprobs returned is computed from the continuation tokens.
We return the sum of the logprob of the continuation tokens
[assuming the continuation tokens are the golden output].
"""
res = []
for logits_single_batch, tokens_single_batch, ctxlen_single_batch in zip(
outputs,
tokens, # type: ignore[arg-type]
ctxlens, # type: ignore[arg-type]
):
logits_single_batch = logits_single_batch.to("cuda")
continuation_logprob = F.log_softmax(
logits_single_batch[(ctxlen_single_batch - 1) : -1], dim=-1
)
continuation_tokens = torch.tensor(tokens_single_batch[ctxlen_single_batch:])
top_tokens = continuation_logprob.argmax(dim=-1).cpu()
is_greedy = torch.equal(top_tokens, continuation_tokens)
logprob_sum = (
continuation_logprob[
torch.arange(continuation_logprob.size(0)), continuation_tokens
]
.sum()
.cpu()
)
res.append((logprob_sum, is_greedy))
return res
if __name__ == "__main__":
cli_evaluate()
# Force clean up the LLM instance and void hanging.
gc.collect()
# Force terminate in case gc.collect() is not enough.
def _terminate():
time.sleep(10)
os.kill(os.getpid(), signal.SIGTERM)
termination_thread = threading.Thread(target=_terminate, daemon=True)
termination_thread.start()
+136
View File
@@ -0,0 +1,136 @@
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# 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.
"""Run lm-evaluation-harness against a TensorRT-LLM checkpoint.
Entry point around lm-eval's built-in ``trtllm`` backend
(``lm_eval.models.trtllm_causallms``, new in 0.4.12). It exists only to correct that
backend's ``prompt_logprobs`` handling -- everything else is upstream. Drop this file and
call ``lm_eval`` directly once the fix lands upstream.
python lm_eval_trtllm.py --model trtllm \
--model_args model=<quantized checkpoint dir>,tokenizer=<HF model folder>,\
tensor_parallel_size=<tp>,max_batch_size=<max batch size>,max_input_len=4096 \
--tasks <comma separated tasks> --batch_size <max batch size>
"""
import sys
from importlib.metadata import version
from lm_eval.__main__ import cli_evaluate
from packaging.version import Version
if Version(version("lm_eval")) < Version("0.4.12"):
# 0.4.12 is the first release shipping lm_eval.models.trtllm_causallms.
raise ImportError(f"lm_eval_trtllm.py requires lm-eval >= 0.4.12; found {version('lm_eval')}.")
from lm_eval.models.trtllm_causallms import TRTLLM
# TensorRT-LLM only started passing the prompt token ids into `compute_logprobs` in
# 1.3.0rc11 (`executor/base_worker.py`), which is what makes the requested token always
# present in each `prompt_logprobs` entry. On 1.2.0 and earlier, `prompt_logprobs=1` keeps
# only the top-1 token, so a non-greedy continuation token is simply absent and no correct
# continuation logprob can be recovered -- by this file or by lm-eval's own version.
_MIN_TRTLLM_VERSION = "1.3.0rc11"
_trtllm_version_checked = False
def _check_trtllm_version() -> None:
"""Raise if TensorRT-LLM predates the `prompt_logprobs` layout scored below."""
try:
import tensorrt_llm
except ImportError:
# Nothing to check, and unreachable in a real run: the backend refuses to build a
# model without tensorrt_llm long before any logprob is scored.
return
if Version(tensorrt_llm.__version__) < Version(_MIN_TRTLLM_VERSION):
raise RuntimeError(
f"Loglikelihood tasks need TensorRT-LLM >= {_MIN_TRTLLM_VERSION}; found "
f"{tensorrt_llm.__version__}. Earlier releases return only the top-1 token per "
"prompt position, so the continuation token's logprob is unavailable. Use a "
"newer TensorRT-LLM container, or restrict the run to generative tasks."
)
def _parse_logprobs(tokens: list[int], outputs, ctxlen: int) -> tuple[float, bool]:
"""Sum the continuation logprobs of one request, correcting upstream's alignment.
TensorRT-LLM aligns ``prompt_logprobs`` to the *next* token: its worker computes them
from ``prompt_token_ids[1:] + first_generated_token`` (``executor/base_worker.py``), so
entry ``i`` is the distribution that predicted ``tokens[i + 1]`` and always contains
that token's id -- either in the top-k or appended by ``_topk_logprobs``.
lm-eval 0.4.12's ``TRTLLM._parse_logprobs`` instead reads
``prompt_logprobs[i][tokens[i]]`` and applies its own shift on top, which raises
``KeyError`` on the first request of every loglikelihood task (hellaswag, mmlu, arc).
"""
global _trtllm_version_checked
if not _trtllm_version_checked:
# Checked here rather than at startup so generative-only runs, which never reach
# this path, still work on older TensorRT-LLM releases.
_check_trtllm_version()
_trtllm_version_checked = True
prompt_logprobs = outputs.outputs[0].prompt_logprobs
# Scoring tokens[ctxlen:] reads entries ctxlen-1 .. len(tokens)-2; a shorter list means
# the engine saw a different prompt than we asked about, which would shift every index.
if len(prompt_logprobs) < len(tokens) - 1:
raise RuntimeError(
f"prompt_logprobs has {len(prompt_logprobs)} entries for {len(tokens)} tokens; "
"the engine scored a different prompt than was requested."
)
continuation_logprobs = 0.0
is_greedy = True
# Token 0 has no preceding distribution, so it can never be scored.
for i in range(max(ctxlen, 1), len(tokens)):
logprob = prompt_logprobs[i - 1].get(tokens[i])
if logprob is None:
# Dropping the term instead would silently inflate the reported accuracy.
raise RuntimeError(
f"tokens[{i}] is missing from prompt_logprobs[{i - 1}]; the returned "
"logprobs are misaligned with the requested tokens."
)
continuation_logprobs += logprob.logprob
if logprob.rank != 1:
is_greedy = False
return continuation_logprobs, is_greedy
if not hasattr(TRTLLM, "_parse_logprobs"):
raise RuntimeError(
"lm_eval.models.trtllm_causallms.TRTLLM has no _parse_logprobs to override; the "
f"backend changed shape in lm-eval {version('lm_eval')}. Recheck whether this file "
"is still needed."
)
# Kept so the unit tests can assert the upstream implementation is still the broken one.
# When that assertion starts failing, upstream has fixed the alignment and this whole file
# should be deleted in favour of calling `lm_eval` directly.
_UPSTREAM_PARSE_LOGPROBS = TRTLLM._parse_logprobs
TRTLLM._parse_logprobs = staticmethod(_parse_logprobs)
if __name__ == "__main__":
# Warn up front so an unusable container is obvious before the model loads, but do not
# abort: generative tasks are unaffected by the old prompt_logprobs layout.
try:
_check_trtllm_version()
except RuntimeError as e:
print(f"WARNING: {e}", file=sys.stderr)
cli_evaluate()
+1 -1
View File
@@ -1,5 +1,5 @@
fire>=0.5.0
lm_eval[api,ifeval]>=0.4.10
lm_eval[api,ifeval]>=0.4.12,<0.5
peft>=0.5.0
rwkv>=0.7.3
torchvision