Trace Replay Offline¶
Source https://github.com/vllm-project/vllm/blob/main/examples/generate/trace_replay_offline.py.
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
Trace-replay with vLLM offline inference.
Trace-replay lets you supply a known sequence of decode token IDs alongside
the prompt. Instead of sampling from the model distribution, the engine
injects each decode token deterministically, step by step. All other
outputs — logprobs, token ranks, text decoding — are computed faithfully
from the real logit distribution for that token.
How it works:
1. Start the engine with ``--enable-trace-replay`` (or
``LLM(..., enable_trace_replay=True)``).
2. Set ``SamplingParams.trace_decode_token_ids`` to the list of decode
token IDs you want to force.
3. The engine will output exactly those tokens and stop. ``max_tokens``
is overwritten with the trace length, and the trace is truncated if it
does not fit within ``max_model_len``. EOS tokens inside the trace
sequence do **not** halt generation early.
Requires model runner V2. Requests are rejected with ``ValueError`` when
combined with any of:
* n > 1 * Speculative decoding
* prompt_logprobs * Structured outputs
* repetition_detection * thinking_token_budget
* bad_words
Typical use-cases:
* Reproduce exact outputs from a previous run for benchmarking.
* Compute logprobs for an already-known output (e.g. reference answers).
* Dataset annotation: given (prompt, response) pairs, obtain per-token
logprob scores without altering the response.
Usage:
python examples/generate/trace_replay_offline.py
python examples/generate/trace_replay_offline.py --model facebook/opt-125m
"""
import argparse
from vllm import LLM, SamplingParams
DEFAULT_PROMPT = "Hello, my name is"
def build_llm(args: argparse.Namespace) -> LLM:
"""Construct an LLM from common CLI args."""
llm_kwargs: dict = {
"model": args.model,
"trust_remote_code": args.trust_remote_code,
"tensor_parallel_size": args.tensor_parallel_size,
"enforce_eager": args.enforce_eager,
"gpu_memory_utilization": args.gpu_memory_utilization,
"max_num_seqs": args.max_num_seqs,
"enable_trace_replay": True,
}
if args.max_model_len is not None:
llm_kwargs["max_model_len"] = args.max_model_len
return LLM(**llm_kwargs)
def run_normal_generation(llm: LLM, prompt: str, max_tokens: int) -> list[int]:
"""Run a standard greedy generation and return the output token IDs."""
sampling_params = SamplingParams(
temperature=0.0,
max_tokens=max_tokens,
# Request logprobs for the greedy token at each step so we can
# compare them against the trace-replay logprobs below.
logprobs=1,
)
outputs = llm.generate([prompt], sampling_params=sampling_params)
result = outputs[0].outputs[0]
output_token_ids = list(result.token_ids)
print("[Normal generation]")
print(f" Prompt : {prompt!r}")
print(f" Output token IDs : {output_token_ids}")
print(f" Output text : {result.text!r}")
return output_token_ids
def run_trace_replay(llm: LLM, prompt: str, decode_token_ids: list[int]) -> None:
"""Replay a known decode sequence and print per-token logprobs."""
sampling_params = SamplingParams(
# Provide the decode tokens to replay.
trace_decode_token_ids=decode_token_ids,
# Request top-5 logprobs so we can inspect the distribution.
logprobs=5,
)
outputs = llm.generate([prompt], sampling_params=sampling_params)
result = outputs[0].outputs[0]
replayed_ids = list(result.token_ids)
print("\n[Trace-replay]")
print(f" Requested decode token IDs : {decode_token_ids}")
print(f" Replayed output token IDs : {replayed_ids}")
print(f" Replayed output text : {result.text!r}")
# Verify the replayed tokens match exactly.
assert replayed_ids == decode_token_ids, (
f"Mismatch!\n expected: {decode_token_ids}\n got: {replayed_ids}"
)
print(" Replayed tokens match the requested trace exactly.")
# Show per-token logprobs (computed from the real distribution).
if result.logprobs:
print("\n Per-token logprobs (trace token):")
for step, (token_id, logprob_dict) in enumerate(
zip(replayed_ids, result.logprobs)
):
sampled_lp = logprob_dict.get(token_id)
lp_value = f"{sampled_lp.logprob:.4f}" if sampled_lp is not None else "n/a"
rank = sampled_lp.rank if sampled_lp is not None else "n/a"
print(
f" step {step:2d}: token_id={token_id:6d} "
f"logprob={lp_value} rank={rank}"
)
def run_demo(args: argparse.Namespace) -> None:
print(f"Loading model: {args.model}\n")
llm = build_llm(args)
prompt = args.prompt
print("=" * 60)
print("Step 1 — Normal greedy generation (captures decode tokens)")
print("=" * 60)
decode_token_ids = run_normal_generation(llm, prompt, max_tokens=8)
print("\n" + "=" * 60)
print("Step 2 — Trace-replay with the captured tokens")
print("=" * 60)
run_trace_replay(llm, prompt, decode_token_ids)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Trace-replay with vLLM offline inference"
)
parser.add_argument(
"--prompt",
type=str,
default=DEFAULT_PROMPT,
help="Text prompt to use for generation (default: %(default)r)",
)
parser.add_argument(
"--model",
type=str,
default="facebook/opt-125m",
help="Name or path of the HuggingFace model to use",
)
parser.add_argument(
"--trust-remote-code",
action="store_true",
help="Trust remote code from HuggingFace",
)
parser.add_argument(
"--tensor-parallel-size",
type=int,
default=1,
help="Number of tensor parallel replicas",
)
parser.add_argument(
"--enforce-eager",
action="store_true",
help="Always use eager-mode PyTorch (disable CUDA graph)",
)
parser.add_argument(
"--gpu-memory-utilization",
type=float,
default=0.9,
help="Fraction of GPU memory to use",
)
parser.add_argument(
"--max-model-len",
type=int,
default=None,
help="Model context length",
)
parser.add_argument(
"--max-num-seqs",
type=int,
default=256,
help="Maximum number of sequences per iteration",
)
return parser.parse_args()
def main() -> None:
args = parse_args()
run_demo(args)
if __name__ == "__main__":
main()