From 107a03ba63e005ff03424fed9c4e6cf551b98bb2 Mon Sep 17 00:00:00 2001 From: Karthik Kothuri Date: Mon, 13 Jul 2026 01:43:34 -0700 Subject: [PATCH] =?UTF-8?q?[Core]=20Support=20fp32=20lm=5Fhead=20for=20gen?= =?UTF-8?q?eration=20models=20via=20head=5Fdtype=20(RFC=20#48305=20=C2=A73?= =?UTF-8?q?.6)=20(#48390)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Karthik Kothuri Signed-off-by: wang.yuqi Co-authored-by: Claude Co-authored-by: wang.yuqi --- .../language/generation_ppl_test/ppl_utils.py | 1 + .../language/generation_ppl_test/test_gpt.py | 12 +- tests/v1/sample/test_head_dtype.py | 172 ++++++++++++++++++ vllm/config/model.py | 16 +- vllm/lora/layers/logits_processor.py | 9 + .../model_executor/layers/logits_processor.py | 54 +++++- 6 files changed, 249 insertions(+), 15 deletions(-) create mode 100644 tests/v1/sample/test_head_dtype.py diff --git a/tests/models/language/generation_ppl_test/ppl_utils.py b/tests/models/language/generation_ppl_test/ppl_utils.py index 2b5449bddcb..45686659cd9 100644 --- a/tests/models/language/generation_ppl_test/ppl_utils.py +++ b/tests/models/language/generation_ppl_test/ppl_utils.py @@ -126,3 +126,4 @@ def wikitext_ppl_test( # We are not concerned that the vllm PPL is less than Transformers, # so we only perform one-sided testing. assert differ < atol + return vllm_ppl diff --git a/tests/models/language/generation_ppl_test/test_gpt.py b/tests/models/language/generation_ppl_test/test_gpt.py index 784f3e85a13..af0e30cc4c7 100644 --- a/tests/models/language/generation_ppl_test/test_gpt.py +++ b/tests/models/language/generation_ppl_test/test_gpt.py @@ -11,4 +11,14 @@ MODELS = [GenerateModelInfo("openai-community/gpt2-large", hf_ppl=19.45705604553 @pytest.mark.parametrize("model_info", MODELS) def test_ppl(hf_runner, vllm_runner, model_info: GenerateModelInfo): - wikitext_ppl_test(hf_runner, vllm_runner, model_info) + bf16_ppl = wikitext_ppl_test(hf_runner, vllm_runner, model_info) + fp32_ppl = wikitext_ppl_test( + hf_runner, + vllm_runner, + model_info, + vllm_extra_kwargs={"hf_overrides": {"head_dtype": "float32"}}, + ) + + differ = ((fp32_ppl - bf16_ppl) / bf16_ppl) * 100 + print("fp32 head difference (%):", differ) + assert differ < 0 diff --git a/tests/v1/sample/test_head_dtype.py b/tests/v1/sample/test_head_dtype.py new file mode 100644 index 00000000000..a64aa290aa8 --- /dev/null +++ b/tests/v1/sample/test_head_dtype.py @@ -0,0 +1,172 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Tests for running the generation lm_head in fp32 via ``head_dtype``. + +An fp32 head lets rollout logits match a trainer that computes the lm_head in +fp32, which is required for RL training-inference consistency. +""" + +import math + +import pytest +import torch + +from vllm import LLM, SamplingParams +from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.layers.vocab_parallel_embedding import ( + UnquantizedEmbeddingMethod, +) + + +class _FakeLmHead: + def __init__( + self, + weight: torch.Tensor, + quantized: bool = False, + shard_indices: object | None = None, + ): + self.weight = weight + self.quant_method = object() if quantized else UnquantizedEmbeddingMethod() + self.shard_indices = shard_indices + + +def _build_processor(vocab_size: int) -> LogitsProcessor: + lp = LogitsProcessor(vocab_size) + # The TP gather is orthogonal to the dtype behavior under test. + lp._gather_logits = lambda logits: logits + return lp + + +def test_fp32_head_runs_projection_in_fp32(default_vllm_config): + vocab_size, hidden_size, num_tokens = 64, 16, 4 + lp = _build_processor(vocab_size) + lp.head_dtype = torch.float32 + + hidden_states = torch.randn(num_tokens, hidden_size, dtype=torch.bfloat16) + weight = torch.randn(vocab_size, hidden_size, dtype=torch.bfloat16) + + logits = lp._get_logits(hidden_states, _FakeLmHead(weight), None) + + assert logits.dtype == torch.float32 + assert torch.isfinite(logits).all() + expected = torch.nn.functional.linear(hidden_states.float(), weight.float()) + torch.testing.assert_close(logits, expected) + + +def test_non_fp32_head_dtype_uses_cast_path(default_vllm_config): + # head_dtype != fp32 must not hit the CUDA out_dtype-mm fast path + # (torch.mm only supports fp32 out for fp16/bf16 inputs); the cast path + # handles any dtype. + vocab_size, hidden_size = 64, 16 + lp = _build_processor(vocab_size) + lp.head_dtype = torch.float16 + + hidden_states = torch.randn(4, hidden_size, dtype=torch.bfloat16) + weight = torch.randn(vocab_size, hidden_size, dtype=torch.bfloat16) + + logits = lp._get_logits(hidden_states, _FakeLmHead(weight), None) + + assert logits.dtype == torch.float16 + expected = torch.nn.functional.linear(hidden_states.half(), weight.half()) + torch.testing.assert_close(logits, expected) + + +def test_head_dtype_equal_to_model_dtype_uses_quant_method(default_vllm_config): + vocab_size, hidden_size = 64, 16 + lp = _build_processor(vocab_size) + lp.head_dtype = torch.bfloat16 + + hidden_states = torch.randn(4, hidden_size, dtype=torch.bfloat16) + weight = torch.randn(vocab_size, hidden_size, dtype=torch.bfloat16) + + logits = lp._get_logits(hidden_states, _FakeLmHead(weight), None) + assert logits.dtype == torch.bfloat16 + + +def test_fp32_head_rejects_quantized_lm_head(default_vllm_config): + lp = _build_processor(64) + lp.head_dtype = torch.float32 + lm_head = _FakeLmHead(torch.randn(64, 16, dtype=torch.bfloat16), quantized=True) + + with pytest.raises(ValueError, match="unquantized"): + lp._get_logits(torch.randn(4, 16, dtype=torch.bfloat16), lm_head, None) + + +def test_get_top_tokens_honors_head_dtype(default_vllm_config): + # The spec-decode local-argmax path (get_top_tokens) must run the lm_head + # in head_dtype too, not just _get_logits. + import types + from unittest import mock + + vocab_size, hidden_size = 64, 16 + lp = _build_processor(vocab_size) + lp.head_dtype = torch.float32 + + hidden_states = torch.randn(4, hidden_size, dtype=torch.bfloat16) + weight = torch.randn(vocab_size, hidden_size, dtype=torch.bfloat16) + lm_head = _FakeLmHead( + weight, + shard_indices=types.SimpleNamespace( + num_org_vocab_padding=0, org_vocab_start_index=0 + ), + ) + + with mock.patch( + "vllm.model_executor.layers.logits_processor." + "get_tensor_model_parallel_world_size", + return_value=1, + ): + top = lp.get_top_tokens(lm_head, hidden_states, None) + + expected = torch.nn.functional.linear(hidden_states.float(), weight.float()).argmax( + dim=-1 + ) + assert torch.equal(top, expected) + + +def test_fp32_head_rejected_with_lora(default_vllm_config): + from vllm.lora.layers.logits_processor import LogitsProcessorWithLoRA + + base = _build_processor(64) + base.head_dtype = torch.float32 + + with pytest.raises(ValueError, match="not yet supported with LoRA"): + LogitsProcessorWithLoRA( + base, + hidden_size=16, + dtype=torch.bfloat16, + device=torch.device("cpu"), + sharded_to_full_mapping=None, + ) + + +@pytest.mark.core_model +def test_fp32_head_e2e_no_nan(): + """An fp32 head produces finite logprobs end-to-end. + + Runs on the default (v2) model runner and exercises the + processed_logprobs path, which forces the native sampler and is where a + non-contiguous fp32 logits row previously produced NaN. + """ + llm = LLM( + model="facebook/opt-125m", + hf_overrides={"head_dtype": "float32"}, + logprobs_mode="processed_logprobs", + enforce_eager=True, + gpu_memory_utilization=0.5, + max_model_len=256, + ) + sampling_params = SamplingParams( + temperature=1.0, top_p=0.95, top_k=50, max_tokens=32, logprobs=5, seed=0 + ) + outputs = llm.generate( + ["The capital of France is", "Once upon a time,"], sampling_params + ) + + for output in outputs: + for completion in output.outputs: + for token_id, position in zip(completion.token_ids, completion.logprobs): + # The sampled token survived filtering, so its logprob is finite. + assert math.isfinite(position[token_id].logprob) + # No returned logprob is NaN. + assert not any(math.isnan(lp.logprob) for lp in position.values()) diff --git a/vllm/config/model.py b/vllm/config/model.py index b5f4e69e031..1bd625626d8 100644 --- a/vllm/config/model.py +++ b/vllm/config/model.py @@ -1709,24 +1709,18 @@ class ModelConfig: such as the lm_head in a generation model, or the score or classifier in a classification model. - `head_dtype` currently only supports pooling models. - - - The pooling model defaults to using fp32 head, you can use + - Pooling models default to an fp32 head; use --hf-overrides '{"head_dtype": "model"}' to disable it. + - Generation models default to the model dtype; set + --hf-overrides '{"head_dtype": "float32"}' to run the lm_head in + fp32, which is required for RL training-inference consistency + (the trainer computes logits in fp32). """ head_dtype = _get_head_dtype( config=self.hf_config, dtype=self.dtype, runner_type=self.runner_type ) - if self.runner_type != "pooling" and head_dtype != self.dtype: - logger.warning_once( - "`head_dtype` currently only supports pooling models, " - "fallback to model dtype [%s].", - self.dtype, - ) - return self.dtype - if head_dtype not in current_platform.supported_dtypes: logger.warning_once( "The current platform does not support [%s] head dtype, " diff --git a/vllm/lora/layers/logits_processor.py b/vllm/lora/layers/logits_processor.py index 237a61eace1..31fa79698fb 100644 --- a/vllm/lora/layers/logits_processor.py +++ b/vllm/lora/layers/logits_processor.py @@ -45,6 +45,15 @@ class LogitsProcessorWithLoRA(BaseLayerWithLoRA): self.hidden_size = hidden_size self.dtype = dtype self.device = device + # The fp32 lm_head path lives in the base LogitsProcessor._get_logits, + # which this wrapper bypasses. Rather than silently emit model-dtype + # logits, reject the combination until the LoRA path supports it. + head_dtype = getattr(base_layer, "head_dtype", None) + if head_dtype is not None and head_dtype != dtype: + raise ValueError( + "A head_dtype different from the model dtype (e.g. an fp32 " + "lm_head) is not yet supported with LoRA." + ) self.tp_size = get_tensor_model_parallel_world_size() self.tp_rank = get_tensor_model_parallel_rank() self.sharded_to_full_mapping = sharded_to_full_mapping diff --git a/vllm/model_executor/layers/logits_processor.py b/vllm/model_executor/layers/logits_processor.py index 3541b970668..eb5b1082ce7 100644 --- a/vllm/model_executor/layers/logits_processor.py +++ b/vllm/model_executor/layers/logits_processor.py @@ -3,14 +3,19 @@ """A layer that compute logits from hidden_stats.""" import torch +import torch.nn.functional as F +from vllm.config import get_current_vllm_config from vllm.distributed import ( get_tensor_model_parallel_world_size, tensor_model_parallel_all_gather, tensor_model_parallel_gather, ) from vllm.model_executor.custom_op import PluggableLayer -from vllm.model_executor.layers.vocab_parallel_embedding import VocabParallelEmbedding +from vllm.model_executor.layers.vocab_parallel_embedding import ( + UnquantizedEmbeddingMethod, + VocabParallelEmbedding, +) from vllm.platforms import current_platform @@ -50,6 +55,11 @@ class LogitsProcessor(PluggableLayer): self.soft_cap = soft_cap # Whether to use gather or all-gather to gather the logits. self.use_all_gather = current_platform.use_all_gather() + # Dtype of the lm_head projection. Defaults to the model dtype; an + # fp32 head (via `--hf-overrides '{"head_dtype": "float32"}'`) is + # required for RL training-inference consistency. + model_config = get_current_vllm_config().model_config + self.head_dtype = model_config.head_dtype if model_config is not None else None def forward( self, @@ -86,6 +96,44 @@ class LogitsProcessor(PluggableLayer): logits = tensor_model_parallel_gather(logits) return logits + def _apply_head( + self, + lm_head: VocabParallelEmbedding, + hidden_states: torch.Tensor, + embedding_bias: torch.Tensor | None, + ) -> torch.Tensor: + """Project hidden states through the lm_head, honoring head_dtype.""" + if self.head_dtype is None or self.head_dtype == hidden_states.dtype: + return lm_head.quant_method.apply( + lm_head, hidden_states, bias=embedding_bias + ) + + if not isinstance(lm_head.quant_method, UnquantizedEmbeddingMethod): + raise ValueError( + "A head_dtype different from the model dtype is only " + "supported for an unquantized lm_head." + ) + if ( + self.head_dtype == torch.float32 + and current_platform.is_cuda() + and hidden_states.is_cuda + ): + # Accumulate the projection directly into fp32. This avoids + # materializing an fp32 copy of the lm_head weight on every step, + # unlike casting both operands. `torch.mm(out_dtype=...)` is + # CUDA-only and only supports fp32 output for fp16/bf16 inputs, so + # other cases fall back to the cast path below. + flat = hidden_states.reshape(-1, hidden_states.shape[-1]) + logits = torch.mm(flat, lm_head.weight.t(), out_dtype=self.head_dtype) + if embedding_bias is not None: + logits = logits + embedding_bias.to(self.head_dtype) + return logits.reshape(*hidden_states.shape[:-1], -1) + return F.linear( + hidden_states.to(self.head_dtype), + lm_head.weight.to(self.head_dtype), + embedding_bias.to(self.head_dtype) if embedding_bias is not None else None, + ) + def _get_logits( self, hidden_states: torch.Tensor, @@ -93,7 +141,7 @@ class LogitsProcessor(PluggableLayer): embedding_bias: torch.Tensor | None, ) -> torch.Tensor | None: # Get the logits for the next tokens. - logits = lm_head.quant_method.apply(lm_head, hidden_states, bias=embedding_bias) + logits = self._apply_head(lm_head, hidden_states, embedding_bias) # Gather logits for TP logits = self._gather_logits(logits) @@ -122,7 +170,7 @@ class LogitsProcessor(PluggableLayer): ) tp_size = get_tensor_model_parallel_world_size() - logits = lm_head.quant_method.apply(lm_head, hidden_states, bias=embedding_bias) + logits = self._apply_head(lm_head, hidden_states, embedding_bias) if self.soft_cap is not None: logits = torch.tanh(logits / self.soft_cap) * self.soft_cap if self.scale != 1.0: