Sampler: la última milla de logits a token
Responsabilidades
Cuando el forward del modelo termina, GPUModelRunner recibe logits de [num_tokens, vocab_size]. Lo que queda es convertir esos logits en token ids: aplicar temperature, top_k, top_p, penalties, hacer muestreo greedy o multinomial, y de paso recolectar los top-k logprobs. Sampler (sampler.py:20) es el nn.Module que hace eso; cada GPUModelRunner tiene una instancia (gpu_model_runner.py:532).
Sampler.forward es un pipeline fijo de nueve pasos (sampler.py:22-59): primero calcula según haga falta raw_logprobs / clona raw_logits, convierte los logits a float32, aplica la whitelist allowed_token_ids, la blacklist bad_words, corre los logits processors non-argmax-invariant (MinTokens, LogitBias), aplica repetition / frequency / presence penalty, y luego entra en sample() para hacer muestreo greedy o random. sample (sampler.py:243) a su vez se parte: greedy va por argmax, random va por apply_temperature → processor argmax-invariant (MinP por defecto) → topk_topp_sampler → torch.where(temp < eps, greedy, random).
La implementación concreta de top-k/top-p no está en Sampler, sino en TopKTopPSampler (topk_topp_sampler.py:70) — en CUDA va por top_p_sampling_from_probs / top_k_sampling_from_probs de FlashInfer (topk_topp_sampler.py:471), en CPU/XPU/ROCm va por sus implementaciones nativas, y si no cae en ninguno hace fallback a forward_native.
Motivación de diseño
- Clasificación por invarianza a argmax:
LogitsProcessor.is_argmax_invariant(interface.py:84-92) separa los que cambian o no el resultado greedy; MinP, MinTokens, LogitBias caen a cada lado; las requests greedy saltan todo el flujo random detrás del argmax-invariant (sampler.py:257-271). - Temperatura 0 tratada como greedy:
apply_temperaturecambia las posiciones contemp < epsa 1.0 (sampler.py:233-237), y al final un únicotorch.where(temp < eps, greedy, random)combina todo, evitando ramas. - Prioridad in-place:
apply_temperatureusalogits.div_yMinP.applyusamasked_fill_(builtin.py:115) para bajar el pico de VRAM. - FlashInfer sampler rápido pero con funciones limitadas: no devuelve los logits filtrados (
topk_topp_sampler.py:86-95), por eso cuando logprobs_mode esprocessed_logits/processed_logprobsse fuerza el paso porforward_native. - Recolección de logprobs separada del muestreo:
gather_logprobs(sampler.py:308-356) corre después del muestreo y usa los logits crudos (sin escalar por temperature) para calcular log_softmax, evitando que las penalties afecten al valor del logprob devuelto al usuario. - Protección contra dinámica de dimensión de batch:
gather_logprobsusatorch._dynamo.decorators.mark_unbacked(sampler.py:345-346) para que dynamo no recompile cuando batch_size pasa de 1 a ≥2. - Compatibilidad con spec decode: en
predict_bonus_token,apply_logits_processorsconcatena losspec_token_idsal final deoutput_token_ids(sampler.py:384-393) para que las penalties vean los tokens del draft.
Archivos clave
Sampler 类:20—nn.Module, mantieneTopKTopPSampler, implementa el pipeline de 9 pasos.forward docstring:22-59— la especificación del pipeline de 9 pasos, anotando qué hace cada uno.Sampler.forward:72-149— entrada principal, encadena apply_logits_processors + sample + gather_logprobs.Sampler.sample:243— bifurcación greedy / random, application de temperature + top-k/top-p.apply_temperature:227-237—div_in-place, cambia temperatura 0 a 1.0 para evitar división por cero.greedy_sample:239-241—logits.argmax(dim=-1).view(-1).apply_logits_processors:371-420— allowed_token_ids_mask + bad_words + non_argmax_invariant processors + penalties.apply_penalties:422-439— delega a la opapply_all_penalties.SamplingMetadata:14— mantiene temperature/top_p/top_k/generators/penalties y demás metadatos por lote.TopKTopPSampler:70— despacho por plataforma: CUDA→FlashInfer, CPU→native, XPU→kernel, ROCm→aiter.flashinfer_sample:471-508— llama atop_p_renorm_probs/top_k_sampling_from_probsde FlashInfer; estadísticamente equivalente pero más rápido que rejection.GPUModelRunner 持有 Sampler:532— en el constructor haceSampler(logprobs_mode, use_fp64_gumbel).
Flujo de datos
La entrada de cada paso son logits de [num_tokens, vocab_size], donde num_tokens es la cantidad total de tokens a calcular en el lote del step actual (prefill calcula varios tokens del prompt, decode calcula 1 por request). Esta es la bifurcación greedy / random dentro del método sample:
# vllm/v1/sample/sampler.py L243-L302
def sample(
self,
logits: torch.Tensor,
sampling_metadata: SamplingMetadata,
logprobs_mode_override: LogprobsMode | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
logprobs_mode = logprobs_mode_override or self.logprobs_mode
assert not (sampling_metadata.all_greedy and sampling_metadata.all_random)
if sampling_metadata.all_random:
greedy_sampled = None
else:
greedy_sampled = self.greedy_sample(logits)
if sampling_metadata.all_greedy:
processed_logprobs = None
if (
sampling_metadata.max_num_logprobs is not None
or sampling_metadata.logprob_token_ids
):
if logprobs_mode == "processed_logits":
processed_logprobs = logits
elif logprobs_mode == "processed_logprobs":
processed_logprobs = self.compute_logprobs(logits)
return greedy_sampled, processed_logprobs
assert sampling_metadata.temperature is not None
# Apply temperature.
logits = self.apply_temperature(
logits, sampling_metadata.temperature, sampling_metadata.all_random
)
# Apply logits processors that only apply to random sampling
# (argmax invariant)
for processor in sampling_metadata.logitsprocs.argmax_invariant:
logits = processor.apply(logits)
# Apply top_k and/or top_p.
random_sampled, processed_logprobs = self.topk_topp_sampler(
logits,
sampling_metadata.generators,
sampling_metadata.top_k,
sampling_metadata.top_p,
)
if greedy_sampled is None:
return random_sampled, processed_logprobs
sampled = torch.where(
sampling_metadata.temperature < _SAMPLING_EPS,
greedy_sampled,
random_sampled,
out=greedy_sampled, # Reuse tensor
)
return sampled, processed_logprobsall_greedy / all_random son flags de lote precomputados en SamplingMetadata que permiten tomar el fast path para todo el lote. En un lote mixto se calculan ambos caminos y luego torch.where elige según la temperatura por-request — así la GPU no se frena por ramas.
Límites y fallos
all_greedyyall_randomson mutuamente excluyentes: al inicio desamplehayassert not (all_greedy and all_random)(sampler.py:256); si la configuración activa ambos, falla con AssertionError.- Los tokens con temperatura 0 no van por random:
torch.where(temp < eps, greedy, random)(sampler.py:296-301) garantiza que las requests greedy siempre tomen argmax, sin contaminarse con random. - logprobs usan logits crudos:
compute_logprobsse calcula antes deapply_temperature(sampler.py:85-93); penalties y temperature no afectan los valores de logprob devueltos al usuario — esta es una diferencia con v0. logprob_token_idsva por un gather dedicado:gather_specific_token_logprobs(sampler.py:151-225) se usa en el APIgenerative_scoring, salta el top-k y hace gather directo de los tokens especificados, con padding y valid_mask.- Spec decode desactiva logits processors: si
build_logitsprocsdetectaspeculative_configy se pasó un processor custom, lanzaValueError(__init__.py:201-209) y solo conserva MinTokens. - FlashInfer devuelve int32: al final del
forward,sampled = sampled.long()(sampler.py:109) unifica a int64, porque el sampler de FlashInfer devuelve int32 mientras que PyTorch argmax / topk devuelven int64, y las operaciones de indexado posteriores deben ser compatibles.
Resumen
Sampler es la entrada única de muestreo de vLLM: pipeline de nueve pasos + bifurcación greedy/random + aceleración con FlashInfer. Depende de la abstracción LogitsProcessor de /sampling/logits para las transformaciones parametrizadas, y él mismo solo se ocupa de "cómo convertir logits en tokens". El SamplerOutput devuelto sube a /worker/gpu-model-runner, y finalmente OutputProcessor lo recompone en RequestOutput. Para cómo el núcleo del motor programa estos steps, ver /engine/engine-core.