Skip to content

Sampler: la última milla de logits a token

源码版本v0.25.1

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_samplertorch.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_temperature cambia las posiciones con temp < eps a 1.0 (sampler.py:233-237), y al final un único torch.where(temp < eps, greedy, random) combina todo, evitando ramas.
  • Prioridad in-place: apply_temperature usa logits.div_ y MinP.apply usa masked_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 es processed_logits / processed_logprobs se fuerza el paso por forward_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_logprobs usa torch._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_processors concatena los spec_token_ids al final de output_token_ids (sampler.py:384-393) para que las penalties vean los tokens del draft.

Archivos clave

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:

python
# 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_logprobs

all_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_greedy y all_random son mutuamente excluyentes: al inicio de sample hay assert 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_logprobs se calcula antes de apply_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_ids va por un gather dedicado: gather_specific_token_logprobs (sampler.py:151-225) se usa en el API generative_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_logitsprocs detecta speculative_config y se pasó un processor custom, lanza ValueError (__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.

Véase la documentación oficial: vLLM 文档 · README.