From 877dae9c684cd8cb2f66eed33f1c3261c6792c10 Mon Sep 17 00:00:00 2001 From: Wentao Ye <44945378+yewentao256@users.noreply.github.com> Date: Fri, 17 Jul 2026 10:57:13 -0400 Subject: [PATCH 01/51] [Refactor] Remove deepseek dead code (#48780) Signed-off-by: yewentao256 Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> --- vllm/model_executor/models/deepseek_eagle3.py | 3 - vllm/model_executor/models/deepseek_mtp.py | 3 - .../warmup/deepseek_v4_mhc_warmup.py | 18 -- vllm/renderers/deepseek_v32.py | 3 - vllm/renderers/deepseek_v4.py | 3 - vllm/tokenizers/deepseek_v32_encoding.py | 169 ---------------- vllm/tokenizers/deepseek_v4_encoding.py | 191 +----------------- 7 files changed, 3 insertions(+), 387 deletions(-) diff --git a/vllm/model_executor/models/deepseek_eagle3.py b/vllm/model_executor/models/deepseek_eagle3.py index 127734808dc..71ce971d530 100644 --- a/vllm/model_executor/models/deepseek_eagle3.py +++ b/vllm/model_executor/models/deepseek_eagle3.py @@ -12,7 +12,6 @@ from transformers import DeepseekV2Config, DeepseekV3Config from vllm.compilation.decorators import support_torch_compile from vllm.config import VllmConfig, get_current_vllm_config -from vllm.logger import init_logger from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ReplicatedLinear from vllm.model_executor.layers.logits_processor import LogitsProcessor @@ -36,8 +35,6 @@ from .utils import ( process_eagle_weight, ) -logger = init_logger(__name__) - class DeepseekV2Eagle3DecoderLayer(nn.Module): """ diff --git a/vllm/model_executor/models/deepseek_mtp.py b/vllm/model_executor/models/deepseek_mtp.py index 746f20e7461..a5dd6ae8f9f 100644 --- a/vllm/model_executor/models/deepseek_mtp.py +++ b/vllm/model_executor/models/deepseek_mtp.py @@ -11,7 +11,6 @@ from vllm._aiter_ops import rocm_aiter_ops from vllm.compilation.decorators import support_torch_compile from vllm.config import VllmConfig from vllm.distributed import tensor_model_parallel_all_gather -from vllm.logger import init_logger from vllm.model_executor.layers.fused_moe import ( fused_moe_make_expert_params_mapping, ) @@ -38,8 +37,6 @@ from .deepseek_v2 import ( ) from .utils import get_pp_missing_layer_names, maybe_prefix -logger = init_logger(__name__) - def _restore_full_token_layout_if_needed( hidden_states: torch.Tensor, diff --git a/vllm/model_executor/warmup/deepseek_v4_mhc_warmup.py b/vllm/model_executor/warmup/deepseek_v4_mhc_warmup.py index 5b4900b50ce..6ca8c94fff0 100644 --- a/vllm/model_executor/warmup/deepseek_v4_mhc_warmup.py +++ b/vllm/model_executor/warmup/deepseek_v4_mhc_warmup.py @@ -15,7 +15,6 @@ import torch from vllm.logger import init_logger from vllm.tracing import instrument -from vllm.utils.math_utils import cdiv logger = init_logger(__name__) @@ -39,23 +38,6 @@ _DEFAULT_TOKEN_SIZE_CANDIDATES = ( ) -def _compute_mhc_pre_num_split( - *, - num_tokens: int, - hidden_size: int, - hc_mult: int, - num_sms: int, -) -> int: - block_k = 64 - block_m = 64 - k = hc_mult * hidden_size - grid_size = cdiv(num_tokens, block_m) - split_k = num_sms // grid_size - num_block_k = cdiv(k, block_k) - split_k = min(split_k, num_block_k // 4) - return max(split_k, 1) - - def _normalize_token_sizes( token_sizes: Iterable[int], *, diff --git a/vllm/renderers/deepseek_v32.py b/vllm/renderers/deepseek_v32.py index 45a46b23283..ba2e0d60c1a 100644 --- a/vllm/renderers/deepseek_v32.py +++ b/vllm/renderers/deepseek_v32.py @@ -8,7 +8,6 @@ from vllm.entrypoints.chat_utils import ( parse_chat_messages, parse_chat_messages_async, ) -from vllm.logger import init_logger from vllm.tokenizers.deepseek_v32 import DeepseekV32Tokenizer from vllm.utils.async_utils import make_async @@ -17,8 +16,6 @@ from .inputs import DictPrompt from .inputs.preprocess import parse_dec_only_prompt from .params import ChatParams -logger = init_logger(__name__) - class DeepseekV32Renderer(BaseRenderer[DeepseekV32Tokenizer]): def __init__( diff --git a/vllm/renderers/deepseek_v4.py b/vllm/renderers/deepseek_v4.py index 3dc82b9622e..e93069209a2 100644 --- a/vllm/renderers/deepseek_v4.py +++ b/vllm/renderers/deepseek_v4.py @@ -8,7 +8,6 @@ from vllm.entrypoints.chat_utils import ( parse_chat_messages, parse_chat_messages_async, ) -from vllm.logger import init_logger from vllm.tokenizers.deepseek_v4 import DeepseekV4Tokenizer from vllm.utils.async_utils import make_async @@ -17,8 +16,6 @@ from .inputs import DictPrompt from .inputs.preprocess import parse_dec_only_prompt from .params import ChatParams -logger = init_logger(__name__) - class DeepseekV4Renderer(BaseRenderer[DeepseekV4Tokenizer]): def __init__( diff --git a/vllm/tokenizers/deepseek_v32_encoding.py b/vllm/tokenizers/deepseek_v32_encoding.py index 249b5326275..b02449020e7 100644 --- a/vllm/tokenizers/deepseek_v32_encoding.py +++ b/vllm/tokenizers/deepseek_v32_encoding.py @@ -7,8 +7,6 @@ import copy import json from typing import Any -import regex as re - # flake8: noqa: E501 TOOLS_SYSTEM_TEMPLATE = """## Tools You have access to a set of tools you can use to answer the user's question. @@ -79,19 +77,6 @@ def tool_calls_from_openai_format(tool_calls): ] -def tool_calls_to_openai_format(tool_calls): - return [ - { - "type": "function", - "function": { - "name": tool_call["name"], - "arguments": tool_call["arguments"], - }, - } - for tool_call in tool_calls - ] - - def encode_arguments_to_dsml(tool_call: dict[str, str]) -> str: p_dsml_template = """<{dsml_token}parameter name="{key}" string="{is_str}">{value}""" P_dsml_strs = [] @@ -113,24 +98,6 @@ def encode_arguments_to_dsml(tool_call: dict[str, str]) -> str: return "\n".join(P_dsml_strs) -def decode_dsml_to_arguments( - tool_name: str, tool_args: dict[str, tuple[str, str]] -) -> dict[str, str]: - def _decode_value(key: str, value: str, string: str): - if string == "true": - value = to_json(value) - return f"{to_json(key)}: {value}" - - tool_args_json = ( - "{" - + ", ".join( - [_decode_value(k, v, string=is_str) for k, (v, is_str) in tool_args.items()] - ) - + "}" - ) - return dict(name=tool_name, arguments=tool_args_json) - - def render_tools(tools: list[dict[str, str | dict[str, Any]]]) -> str: tools_json = [to_json(t) for t in tools] @@ -333,139 +300,3 @@ def encode_messages( ) return prompt - - -def _read_until_stop( - index: int, text: str, stop: list[str] -) -> tuple[int, str, None | str]: - min_pos = len(text) - matched_stop = None - - for s in stop: - pos = text.find(s, index) - if pos != -1 and pos < min_pos: - min_pos = pos - matched_stop = s - - if matched_stop: - content = text[index:min_pos] - return min_pos + len(matched_stop), content, matched_stop - else: - content = text[index:] - return len(text), content, None - - -def parse_tool_calls(index: int, text: str): - tool_calls: list[dict[str, Any]] = [] - stop_token = None - tool_calls_end_token = f"" - - while index < len(text): - index, _, stop_token = _read_until_stop( - index, text, [f"<{dsml_token}invoke", tool_calls_end_token] - ) - if _ != ">\n": - raise RuntimeError("Tool call format error") - - if stop_token == tool_calls_end_token: - break - - if stop_token is None: - raise RuntimeError("Missing special token") - - index, tool_name_content, stop_token = _read_until_stop( - index, text, [f"<{dsml_token}parameter", f"\n$', tool_name_content, flags=re.DOTALL - ) - if len(p_tool_name) != 1: - raise RuntimeError("Tool name format error") - tool_name = p_tool_name[0] - - tool_args: dict[str, tuple[str, str]] = {} - while stop_token == f"<{dsml_token}parameter": - index, param_content, stop_token = _read_until_stop( - index, text, [f"/{dsml_token}parameter"] - ) - - param_kv = re.findall( - r'^ name="(.*?)" string="(true|false)">(.*?)<$', - param_content, - flags=re.DOTALL, - ) - if len(param_kv) != 1: - raise RuntimeError("Parameter format error") - param_name, string, param_value = param_kv[0] - - if param_name in tool_args: - raise RuntimeError("Duplicate parameter name") - tool_args[param_name] = (param_value, string) - - index, content, stop_token = _read_until_stop( - index, text, [f"<{dsml_token}parameter", f"\n": - raise RuntimeError("Parameter format error") - - tool_call = decode_dsml_to_arguments(tool_name=tool_name, tool_args=tool_args) - tool_calls.append(tool_call) - - return index, stop_token, tool_calls - - -# NOTE: This function is designed to parse only correctly -# formatted string and will not attempt to correct malformed output -# that may be generated by the model. -def parse_message_from_completion_text(text: str, thinking_mode: str): - summary_content, reasoning, tool_calls = "", "", [] - index, stop_token = 0, None - tool_calls_start_token = f"\n\n<{dsml_token}function_calls" - - is_thinking, is_tool_calling = thinking_mode == "thinking", False - - if is_thinking: - index, content_delta, stop_token = _read_until_stop( - index, text, [thinking_end_token, tool_calls_start_token] - ) - reasoning = content_delta - if stop_token != thinking_end_token: - raise RuntimeError("Invalid thinking format") - - index, content_delta, stop_token = _read_until_stop( - index, text, [eos_token, tool_calls_start_token] - ) - summary_content = content_delta - if stop_token == tool_calls_start_token: - is_tool_calling = True - else: - if stop_token != eos_token: - raise RuntimeError("Invalid summary format") - - if is_tool_calling: - index, stop_token, tool_calls = parse_tool_calls(index, text) - - index, tool_ends_text, stop_token = _read_until_stop(index, text, [eos_token]) - if tool_ends_text: - raise RuntimeError("Unexpected content after tool calls") - - if not (len(text) == index and stop_token in [eos_token, None]): - raise RuntimeError("Unexpected content at end") - - for sp_token in [ - bos_token, - eos_token, - thinking_start_token, - thinking_end_token, - dsml_token, - ]: - if sp_token in summary_content or sp_token in reasoning: - raise RuntimeError("Unexpected special token in content") - - return { - "role": "assistant", - "content": summary_content, - "reasoning": reasoning, - "tool_calls": tool_calls_to_openai_format(tool_calls), - } diff --git a/vllm/tokenizers/deepseek_v4_encoding.py b/vllm/tokenizers/deepseek_v4_encoding.py index 6895771e2f5..16bfa1a99a1 100644 --- a/vllm/tokenizers/deepseek_v4_encoding.py +++ b/vllm/tokenizers/deepseek_v4_encoding.py @@ -6,16 +6,14 @@ """ DeepSeek-V4 Encoding -A self-contained implementation for encoding/decoding DeepSeek-V4 chat messages -with tool calling, thinking mode, and quick instruction task support. +A self-contained implementation for encoding DeepSeek-V4 chat messages with tool +calling, thinking mode, and quick instruction task support. """ -from typing import Any, Dict, List, Union, Optional, Tuple +from typing import Any, Dict, List, Union, Optional import copy import json -import regex as re - # ============================================================ # Special Tokens # ============================================================ @@ -128,20 +126,6 @@ def tool_calls_from_openai_format(tool_calls): ] -def tool_calls_to_openai_format(tool_calls): - """Convert internal tool calls to OpenAI format.""" - return [ - { - "type": "function", - "function": { - "name": tool_call["name"], - "arguments": tool_call["arguments"], - } - } - for tool_call in tool_calls - ] - - def encode_arguments_to_dsml(tool_call: Dict[str, Any]) -> str: """ Encode tool call arguments into DSML parameter format. @@ -172,26 +156,6 @@ def encode_arguments_to_dsml(tool_call: Dict[str, Any]) -> str: return "\n".join(P_dsml_strs) -def decode_dsml_to_arguments(tool_name: str, tool_args: Dict[str, Tuple[str, str]]) -> Dict[str, str]: - """ - Decode DSML parameters back to a tool call dict. - - Args: - tool_name: Name of the tool. - tool_args: Dict mapping param_name -> (value, is_string_flag). - - Returns: - Dict with "name" and "arguments" (JSON string) keys. - """ - def _decode_value(key: str, value: str, string: str): - if string == "true": - value = to_json(value) - return f"{to_json(key)}: {value}" - - tool_args_json = "{" + ", ".join([_decode_value(k, v, string=is_str) for k, (v, is_str) in tool_args.items()]) + "}" - return dict(name=tool_name, arguments=tool_args_json) - - def render_tools(tools: List[Dict[str, Union[str, Dict[str, Any]]]]) -> str: """ Render tool schemas into the system prompt format. @@ -605,153 +569,4 @@ def _drop_thinking_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, An return result -# ============================================================ -# Parsing (Decoding model output) -# ============================================================ - -def _read_until_stop(index: int, text: str, stop: List[str]) -> Tuple[int, str, Optional[str]]: - """ - Read text from index until one of the stop strings is found. - - Returns: - Tuple of (new_index, content_before_stop, matched_stop_string_or_None). - """ - min_pos = len(text) - matched_stop = None - - for s in stop: - pos = text.find(s, index) - if pos != -1 and pos < min_pos: - min_pos = pos - matched_stop = s - - if matched_stop: - content = text[index:min_pos] - return min_pos + len(matched_stop), content, matched_stop - else: - content = text[index:] - return len(text), content, None - - -def parse_tool_calls(index: int, text: str) -> Tuple[int, Optional[str], List[Dict[str, str]]]: - """ - Parse DSML tool calls from text starting at the given index. - - Args: - index: Starting position in text. - text: The full text to parse. - - Returns: - Tuple of (new_index, last_stop_token, list_of_tool_call_dicts). - Each tool call dict has "name" and "arguments" keys. - """ - tool_calls: List[Dict[str, Any]] = [] - stop_token = None - tool_calls_end_token = f"" - - while index < len(text): - index, content_before, stop_token = _read_until_stop(index, text, [f"<{dsml_token}invoke", tool_calls_end_token]) - if content_before != ">\n": - raise ValueError(f"Tool call format error: expected '>\\n' but got '{content_before}'") - - if stop_token == tool_calls_end_token: - break - - if stop_token is None: - raise ValueError("Missing special token in tool calls") - - index, tool_name_content, stop_token = _read_until_stop(index, text, [f"<{dsml_token}parameter", f"\n$', tool_name_content, flags=re.DOTALL) - if len(p_tool_name) != 1: - raise ValueError(f"Tool name format error: '{tool_name_content}'") - tool_name = p_tool_name[0] - - tool_args: Dict[str, Tuple[str, str]] = {} - while stop_token == f"<{dsml_token}parameter": - index, param_content, stop_token = _read_until_stop(index, text, [f"/{dsml_token}parameter"]) - - param_kv = re.findall(r'^ name="(.*?)" string="(true|false)">(.*?)<$', param_content, flags=re.DOTALL) - if len(param_kv) != 1: - raise ValueError(f"Parameter format error: '{param_content}'") - param_name, string, param_value = param_kv[0] - - if param_name in tool_args: - raise ValueError(f"Duplicate parameter name: '{param_name}'") - tool_args[param_name] = (param_value, string) - - index, content, stop_token = _read_until_stop(index, text, [f"<{dsml_token}parameter", f"\n": - raise ValueError(f"Parameter format error: expected '>\\n' but got '{content}'") - - tool_call = decode_dsml_to_arguments(tool_name=tool_name, tool_args=tool_args) - tool_calls.append(tool_call) - - return index, stop_token, tool_calls - - -def parse_message_from_completion_text(text: str, thinking_mode: str) -> Dict[str, Any]: - """ - Parse a model completion text into a structured assistant message. - - This function takes the raw text output from the model (a single assistant turn) - and extracts: - - reasoning (thinking block) - - content (summary/response) - - tool_calls (if any) - - NOTE: This function is designed to parse only correctly formatted strings and - will raise ValueError for malformed output. - - Args: - text: The raw completion text (including EOS token). - thinking_mode: Either "chat" or "thinking". - - Returns: - Dict with keys: "role", "content", "reasoning", "tool_calls". - tool_calls are in OpenAI format. - """ - summary_content, reasoning = "", "" - tool_calls: List[Dict[str, str]] = [] - index, stop_token = 0, None - tool_calls_start_token = f"\n\n<{dsml_token}{tool_calls_block_name}" - - is_thinking = thinking_mode == "thinking" - is_tool_calling = False - - if is_thinking: - index, content_delta, stop_token = _read_until_stop(index, text, [thinking_end_token, tool_calls_start_token]) - reasoning = content_delta - if stop_token != thinking_end_token: - raise ValueError("Invalid thinking format: missing ") - - index, content_delta, stop_token = _read_until_stop(index, text, [eos_token, tool_calls_start_token]) - summary_content = content_delta - if stop_token == tool_calls_start_token: - is_tool_calling = True - else: - if stop_token != eos_token: - raise ValueError("Invalid format: missing EOS token") - - if is_tool_calling: - index, stop_token, tool_calls = parse_tool_calls(index, text) - - index, tool_ends_text, stop_token = _read_until_stop(index, text, [eos_token]) - if tool_ends_text: - raise ValueError("Unexpected content after tool calls") - - if len(text) != index or stop_token not in [eos_token, None]: - raise ValueError("Unexpected content at end") - - for sp_token in [bos_token, eos_token, thinking_start_token, thinking_end_token, dsml_token]: - if sp_token in summary_content or sp_token in reasoning: - raise ValueError(f"Unexpected special token '{sp_token}' in content") - - return { - "role": "assistant", - "content": summary_content, - "reasoning": reasoning, - "tool_calls": tool_calls_to_openai_format(tool_calls) - } - # fmt: on From 11d291511a35bfa2a9ecb8fd21ee1c48b3a78d6b Mon Sep 17 00:00:00 2001 From: mosya415 Date: Fri, 17 Jul 2026 19:45:41 +0300 Subject: [PATCH 02/51] [Bugfix][Tool Parser] Preserve whitespace in parameter values (MiniMax M2, Qwen3, MiniCPM5 XML) (#48846) Signed-off-by: mosya415 <263250241+mosya415@users.noreply.github.com> Signed-off-by: Ben Browning <56071+bbrowning@users.noreply.github.com> Co-authored-by: mosya415 <263250241+mosya415@users.noreply.github.com> Co-authored-by: Ben Browning <56071+bbrowning@users.noreply.github.com> --- tests/parser/engine/test_qwen3.py | 4 +- .../test_minicpm5xml_tool_parser.py | 52 ++++++++++++++++++ .../test_minimax_m2_tool_parser.py | 34 ++++++++++++ .../test_qwen3coder_tool_parser.py | 55 +++++++++++++++++++ vllm/parser/minimax_m2.py | 7 ++- vllm/parser/qwen3.py | 13 ++++- vllm/tool_parsers/minicpm5xml_tool_parser.py | 8 ++- 7 files changed, 164 insertions(+), 9 deletions(-) diff --git a/tests/parser/engine/test_qwen3.py b/tests/parser/engine/test_qwen3.py index 38450762421..07a025876ae 100644 --- a/tests/parser/engine/test_qwen3.py +++ b/tests/parser/engine/test_qwen3.py @@ -454,10 +454,10 @@ class TestStreaming: args_after_partial_tag = collect_tool_arguments(results[:4]) assert " None: "nums": [7, 8, 9], "exact": False, } + + +def make_tools_write() -> list[ChatCompletionToolsParam]: + return [ + _tool( + "write_file", + { + "type": "object", + "properties": {"content": {"type": "string"}}, + "required": ["content"], + }, + ) + ] + + +class TestParameterWhitespace: + """CDATA is verbatim, so its whitespace must survive.""" + + def test_cdata_whitespace_preserved(self, parser: ToolParser) -> None: + request = make_request(make_tools_write()) + text = ( + '' + '' + "\n" + ) + + out = parser.extract_tool_calls(text, request) + + assert len(out.tool_calls) == 1 + args = json.loads(out.tool_calls[0].function.arguments) + assert args["content"] == " def foo():\n pass\n" + + def test_cdata_whitespace_preserved_streaming(self, parser: ToolParser) -> None: + """Value split across chunks, i.e. the partial path.""" + request = make_request(make_tools_write()) + chunks = [ + '', + '\n", + ] + + reconstructor = run_tool_extraction_streaming( + parser, + chunks, + request, + assert_one_tool_per_delta=False, + ) + + assert len(reconstructor.tool_calls) == 1 + assert json.loads(reconstructor.tool_calls[0].function.arguments) == { + "content": " def foo():\n pass\n" + } diff --git a/tests/tool_parsers/test_minimax_m2_tool_parser.py b/tests/tool_parsers/test_minimax_m2_tool_parser.py index 029ee21ae1f..bb4ffdc4377 100644 --- a/tests/tool_parsers/test_minimax_m2_tool_parser.py +++ b/tests/tool_parsers/test_minimax_m2_tool_parser.py @@ -567,3 +567,37 @@ class TestNoneStringPreservation: assert len(tc) == 1 parsed = json.loads(tc[0]["arguments"]) assert parsed["value"] == "nil" + + +class TestParameterWhitespace: + """Parameter values must preserve surrounding whitespace.""" + + def test_whitespace_preserved(self, parser): + results = _feed( + parser, + [ + '' + ' hi ' + "", + ], + ) + tc = _collect_tool_calls(results) + assert len(tc) == 1 + assert json.loads(tc[0]["arguments"]) == {"msg": " hi "} + + def test_whitespace_preserved_across_chunks(self, parser): + """Value split before arrives, i.e. the partial path.""" + results = _feed( + parser, + [ + '' + ' def foo():\n', + " return 1\n", + "", + ], + ) + tc = _collect_tool_calls(results) + assert len(tc) == 1 + assert json.loads(tc[0]["arguments"]) == { + "content": " def foo():\n return 1\n" + } diff --git a/tests/tool_parsers/test_qwen3coder_tool_parser.py b/tests/tool_parsers/test_qwen3coder_tool_parser.py index 1f5e51412b9..7222625dbb0 100644 --- a/tests/tool_parsers/test_qwen3coder_tool_parser.py +++ b/tests/tool_parsers/test_qwen3coder_tool_parser.py @@ -1482,3 +1482,58 @@ def test_adjust_request_required_prefers_structural_tag( out = TestParser(MagicMock(), tools=sample_tools).adjust_request(req) assert out.structured_outputs is not None assert out.structured_outputs.structural_tag is not None + + +WRITE_FILE_TOOLS = [ + ChatCompletionToolsParam( + type="function", + function={ + "name": "write_file", + "parameters": { + "type": "object", + "properties": {"content": {"type": "string"}}, + }, + }, + ) +] + +WRITE_FILE_OUTPUT = ( + "\n\n" + "\n" + " def foo():\n return 1\n\n" + "\n\n" +) + +EXPECTED_CONTENT = " def foo():\n return 1\n" + + +class TestParameterWhitespace: + """Wrapping newlines are markup; the rest of the value is preserved.""" + + def test_whitespace_preserved(self, qwen3_tokenizer): + parser = Qwen3EngineToolParser(qwen3_tokenizer, tools=WRITE_FILE_TOOLS) + request = ChatCompletionRequest( + model=MODEL, messages=[], tools=WRITE_FILE_TOOLS + ) + + extracted = parser.extract_tool_calls(WRITE_FILE_OUTPUT, request=request) + + args = json.loads(extracted.tool_calls[0].function.arguments) + assert args["content"] == EXPECTED_CONTENT + + def test_whitespace_preserved_streaming(self, qwen3_tokenizer): + """The partial path must not strip the value either.""" + parser = Qwen3EngineToolParser(qwen3_tokenizer, tools=WRITE_FILE_TOOLS) + request = ChatCompletionRequest( + model=MODEL, messages=[], tools=WRITE_FILE_TOOLS + ) + + streamed = "" + for delta_message in stream_delta_message_generator( + parser, qwen3_tokenizer, WRITE_FILE_OUTPUT, request + ): + for tool_call in delta_message.tool_calls or []: + if tool_call.function and tool_call.function.arguments: + streamed += tool_call.function.arguments + + assert json.loads(streamed)["content"] == EXPECTED_CONTENT diff --git a/vllm/parser/minimax_m2.py b/vllm/parser/minimax_m2.py index 86fa1d1bad1..ab1b7eb1142 100644 --- a/vllm/parser/minimax_m2.py +++ b/vllm/parser/minimax_m2.py @@ -69,7 +69,9 @@ def _minimax_m2_arg_converter(raw_args: str, partial: bool) -> str: ).strip() if not name: continue - params[name] = match.group("value").strip() + # Keep the value verbatim (like the glm47 converter): it is inline + # between `>` and ``, so surrounding whitespace is data. + params[name] = match.group("value") if partial: remaining = _PARAM_RE.sub("", raw_args) @@ -82,7 +84,8 @@ def _minimax_m2_arg_converter(raw_args: str, partial: bool) -> str: or "" ).strip() if name: - params[name] = match.group("value").strip() + # Verbatim, same as the complete-match loop above. + params[name] = match.group("value") return json.dumps(params, ensure_ascii=False) diff --git a/vllm/parser/qwen3.py b/vllm/parser/qwen3.py index f80aa6ff7a2..0041a46eaae 100644 --- a/vllm/parser/qwen3.py +++ b/vllm/parser/qwen3.py @@ -56,13 +56,22 @@ _PARAM_RE = re.compile( _PARTIAL_PARAM_RE = re.compile(r"<\s*parameter\s*=\s*([^>]+)>(.*)$", re.DOTALL) +def _trim_wrapping_newlines(value: str) -> str: + """Strip one leading and one trailing newline (the Qwen3 template markup).""" + if value.startswith("\n"): + value = value[1:] + if value.endswith("\n"): + value = value[:-1] + return value + + def _qwen3_arg_converter(raw_args: str, partial: bool) -> str: params: dict[str, object] = {} for match in _PARAM_RE.finditer(raw_args): name = match.group(1) value = match.group(2) - params[name] = value.strip() + params[name] = _trim_wrapping_newlines(value) if partial: remaining = _PARAM_RE.sub("", raw_args) @@ -71,7 +80,7 @@ def _qwen3_arg_converter(raw_args: str, partial: bool) -> str: name = m.group(1) value = m.group(2) if name: - params[name] = value.strip() + params[name] = _trim_wrapping_newlines(value) return json.dumps(params, ensure_ascii=False) diff --git a/vllm/tool_parsers/minicpm5xml_tool_parser.py b/vllm/tool_parsers/minicpm5xml_tool_parser.py index a5b5252415c..c41eb7c2de8 100644 --- a/vllm/tool_parsers/minicpm5xml_tool_parser.py +++ b/vllm/tool_parsers/minicpm5xml_tool_parser.py @@ -360,7 +360,7 @@ def _parse_function_block( if not key: has_invalid_param = True break - val_text = (param.text or "").strip() + val_text = param.text or "" if not _add_argument( func_name or "", key, @@ -392,7 +392,8 @@ def _parse_function_block( val_text = pm.group(2) or "" if val_text.startswith(""): val_text = val_text[len("")] - val_text = val_text.strip() + else: + val_text = val_text.strip() if not _add_argument( func_name or "", key, @@ -445,7 +446,8 @@ def _parse_partial_params( val_text = pm.group(2) or "" if val_text.startswith(""): val_text = val_text[len("")] - val_text = val_text.strip() + else: + val_text = val_text.strip() _add_argument( func_name, key, From efed8a1e8345de2b4398a48243301d6a5b1a361e Mon Sep 17 00:00:00 2001 From: Fangzhou Ai <31551580+Fangzhou-Ai@users.noreply.github.com> Date: Fri, 17 Jul 2026 13:24:59 -0400 Subject: [PATCH 03/51] [ROCm][Perf][DSV4] Improve sparse decode reduction occupancy on gfx950 (#48788) Signed-off-by: fai --- vllm/v1/attention/ops/rocm_aiter_mla_sparse.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py index c3faee40246..96f7693430a 100644 --- a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py @@ -2205,7 +2205,7 @@ def _rocm_sparse_attn_decode_ragged_triton( num_warps=4, ) - _sparse_attn_decode_reduce_kernel[(num_queries, heads_blocks)]( + _sparse_attn_decode_reduce_kernel[(num_queries, num_heads)]( part_m, part_l, part_acc, @@ -2221,7 +2221,7 @@ def _rocm_sparse_attn_decode_ragged_triton( num_heads, HAS_ATTN_SINK=has_attn_sink, COMB_DIM=comb_dim, - BLOCK_H=block_h, + BLOCK_H=1, NUM_SPLITS=num_splits, SPLITS_PAD=triton.next_power_of_2(num_splits), num_warps=4, From bf578e1abdffc2d25232783ff59a3132279e6bdd Mon Sep 17 00:00:00 2001 From: labAxiaoming <34019940+labAxiaoming@users.noreply.github.com> Date: Sat, 18 Jul 2026 01:44:45 +0800 Subject: [PATCH 04/51] [Bugfix][GLM4V] Fix video dummy profiling and memory usage (#48729) Signed-off-by: xiaoming <1259730330@qq.com> --- .../multimodal/processing/test_glm4_1v.py | 52 ++++++++++++++ vllm/model_executor/models/glm4_1v.py | 69 ++++++------------- 2 files changed, 72 insertions(+), 49 deletions(-) diff --git a/tests/models/multimodal/processing/test_glm4_1v.py b/tests/models/multimodal/processing/test_glm4_1v.py index 5798c566347..0cafa261a34 100644 --- a/tests/models/multimodal/processing/test_glm4_1v.py +++ b/tests/models/multimodal/processing/test_glm4_1v.py @@ -1,9 +1,15 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from unittest.mock import Mock + import pytest from vllm.assets.video import VideoAsset +from vllm.model_executor.models.glm4_1v import ( + Glm4vForConditionalGeneration, + Glm4vProcessingInfo, +) from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import batched_tensors_equal from vllm.multimodal.video import DynamicVideoBackend, VideoBackend @@ -11,6 +17,52 @@ from vllm.multimodal.video import DynamicVideoBackend, VideoBackend from ...utils import build_model_context +@pytest.mark.parametrize( + ( + "max_video_pixels", + "max_tokens", + "expected_num_frames", + ), + [ + (47_040_000, 124_988, 11), + (47_040_000, 30_000, 24), + (100_352_000, 124_988, 21), + (100_352_000, 30_000, 7), + (100_352_000, 0, 1), + ], +) +def test_get_max_video_frames_matches_glm_resize( + max_video_pixels: int, + max_tokens: int, + expected_num_frames: int, +): + info = Mock(spec=Glm4vProcessingInfo) + info.get_image_size_with_most_features.return_value = (2184, 2184) + info._get_video_max_pixels.return_value = max_video_pixels + vision_config = info.get_hf_config.return_value.vision_config + vision_config.patch_size = 14 + vision_config.spatial_merge_size = 2 + vision_config.temporal_patch_size = 2 + info._get_vision_info.side_effect = lambda **kwargs: ( + Glm4vProcessingInfo._get_vision_info(info, **kwargs) + ) + + num_frames = Glm4vProcessingInfo._get_max_video_frames( + info, + max_tokens=max_tokens, + ) + + assert num_frames == expected_num_frames + assert info._get_video_max_pixels.call_count == 1 + assert info._get_vision_info.call_count == 600 + + +def test_encoder_cudagraph_uses_model_video_frame_limit(): + model = Mock() + + assert Glm4vForConditionalGeneration.get_max_frames_per_video(model) == 600 + + @pytest.mark.parametrize("model_id", ["zai-org/GLM-4.1V-9B-Thinking"]) @pytest.mark.parametrize("expected_toks_per_frame", [299]) @pytest.mark.parametrize( diff --git a/vllm/model_executor/models/glm4_1v.py b/vllm/model_executor/models/glm4_1v.py index d64a4ef5130..810d9de87b4 100644 --- a/vllm/model_executor/models/glm4_1v.py +++ b/vllm/model_executor/models/glm4_1v.py @@ -40,6 +40,7 @@ import transformers from einops import rearrange from packaging.version import Version from transformers import BatchFeature, Glm4vProcessor +from transformers.image_processing_base import ImageProcessingMixin from transformers.models.glm4v.configuration_glm4v import ( Glm4vTextConfig, Glm4vVisionConfig, @@ -49,6 +50,7 @@ from transformers.models.glm4v.image_processing_glm4v import ( smart_resize, ) from transformers.models.glm4v.video_processing_glm4v import Glm4vVideoProcessor +from transformers.video_processing_utils import BaseVideoProcessor from transformers.video_utils import VideoMetadata from vllm.config import VllmConfig @@ -94,6 +96,8 @@ from vllm.multimodal.processing import ( PromptUpdateDetails, ) from vllm.sequence import IntermediateTensors +from vllm.transformers_utils.processor import get_processor_cls_name_from_config +from vllm.transformers_utils.utils import convert_model_repo_to_path from vllm.utils.tensor_schema import TensorSchema, TensorShape from vllm.v1.attention.backends.registry import AttentionBackendEnum from vllm.v1.worker.encoder_cudagraph_defs import EncoderCudaGraphReplayBuffers @@ -982,11 +986,6 @@ class Glm4vProcessingInfo(BaseProcessingInfo): return self.get_hf_processor(**kwargs).video_processor def _get_processor_class_name(self) -> str | None: - from vllm.transformers_utils.processor import ( - get_processor_cls_name_from_config, - ) - from vllm.transformers_utils.utils import convert_model_repo_to_path - return get_processor_cls_name_from_config( convert_model_repo_to_path(self.ctx.model_config.model), revision=self.ctx.model_config.revision, @@ -1109,8 +1108,6 @@ class Glm4vProcessingInfo(BaseProcessingInfo): image_processor_config = self.ctx.get_hf_image_processor_config() if not image_processor_config.get("size"): - from transformers.image_processing_base import ImageProcessingMixin - image_processor_config, _ = ImageProcessingMixin.get_image_processor_dict( self.ctx.model_config.model, revision=self.ctx.model_config.revision, @@ -1122,8 +1119,6 @@ class Glm4vProcessingInfo(BaseProcessingInfo): return self._get_longest_edge(size, "GLM4V image processor size") def _get_video_max_pixels(self) -> int: - from transformers.video_processing_utils import BaseVideoProcessor - mm_kwargs = self.ctx.get_merged_mm_kwargs({}) if (override_max_pixels := mm_kwargs.get("max_pixels")) is not None: return int(override_max_pixels) @@ -1175,39 +1170,29 @@ class Glm4vProcessingInfo(BaseProcessingInfo): image_height=target_height, ) - def get_num_video_tokens( - self, - *, - image_width: int, - image_height: int, - num_frames: int, - ) -> int: - _, num_video_tokens = self._get_vision_info( - image_width=image_width, - image_height=image_height, - num_frames=num_frames, - max_image_pixels=28 * 28 * 2 * 30000, - ) - return num_video_tokens - def _get_max_video_frames(self, max_tokens: int) -> int: target_width, target_height = self.get_image_size_with_most_features() - num_frames = 0 + max_video_pixels = self._get_video_max_pixels() + num_frames_with_most_features = 1 + max_vision_tokens = 0 - while True: - next_num_frames = num_frames + 1 - next_max_tokens = self.get_num_video_tokens( + for num_frames in range(1, _MAX_FRAMES_PER_VIDEO + 1): + _, num_vision_tokens = self._get_vision_info( image_width=target_width, image_height=target_height, - num_frames=next_num_frames, + num_frames=num_frames, + max_image_pixels=max_video_pixels, ) - if next_max_tokens > max_tokens or next_max_tokens == 0: - break - num_frames = next_num_frames + if ( + 0 < num_vision_tokens <= max_tokens + and num_vision_tokens > max_vision_tokens + ): + num_frames_with_most_features = num_frames + max_vision_tokens = num_vision_tokens - return num_frames + return num_frames_with_most_features def get_num_frames_with_most_features( self, @@ -1215,15 +1200,9 @@ class Glm4vProcessingInfo(BaseProcessingInfo): mm_counts: Mapping[str, int], ) -> int: max_images = mm_counts.get("image", 0) - max_videos = mm_counts.get("video", 0) max_image_tokens = self.get_max_image_tokens() * max_images - max_total_frames = self._get_max_video_frames(seq_len - max_image_tokens) - max_frames_per_video = min( - max_total_frames // max(max_videos, 1), _MAX_FRAMES_PER_VIDEO - ) - - return max(max_frames_per_video, 1) + return self._get_max_video_frames(seq_len - max_image_tokens) def _get_video_second_idx_glm4v( self, metadata: dict[str, Any], total_frames: int @@ -2001,15 +1980,7 @@ class Glm4vForConditionalGeneration( raise AssertionError("This line should be unreachable.") def get_max_frames_per_video(self) -> int: - mm_registry = MULTIMODAL_REGISTRY - info = mm_registry.get_processing_info(self.model_config) - max_frames_per_video = info.get_num_frames_with_most_features( - seq_len=self.model_config.max_model_len, - mm_counts={"video": self.multimodal_config.get_limit_per_prompt("video")}, - ) - # Small 'max_frames_per_video' will cause 'tensor mismatch' in PR#43403 - # 16 is the default 'num_frames' of '_get_vision_info' - return max(max_frames_per_video, 16) + return _MAX_FRAMES_PER_VIDEO def get_encoder_cudagraph_budget_range( self, From 5784507da458656ef8248119590822445d4c67f2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicol=C3=B2=20Lucchesi?= Date: Fri, 17 Jul 2026 21:19:02 +0200 Subject: [PATCH 05/51] [Attention] Allow selecting a different attention backend per KV-cache group (#48012) Signed-off-by: NickLucche Co-authored-by: Claude Opus 4.8 --- tests/v1/attention/test_backend_per_kind.py | 67 +++++++++++++ .../test_attention_backend_per_kind.py | 94 +++++++++++++++++++ vllm/config/attention.py | 38 ++++++++ vllm/v1/attention/selector.py | 62 +++++++++++- 4 files changed, 258 insertions(+), 3 deletions(-) create mode 100644 tests/v1/attention/test_backend_per_kind.py create mode 100644 tests/v1/e2e/general/test_attention_backend_per_kind.py diff --git a/tests/v1/attention/test_backend_per_kind.py b/tests/v1/attention/test_backend_per_kind.py new file mode 100644 index 00000000000..5b197f7d159 --- /dev/null +++ b/tests/v1/attention/test_backend_per_kind.py @@ -0,0 +1,67 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Tests for per-KV-group attention backend selection (backend_per_kind).""" + +import pytest + +from vllm.config.attention import AttentionConfig +from vllm.v1.attention.backend import AttentionType +from vllm.v1.attention.backends.registry import AttentionBackendEnum +from vllm.v1.attention.selector import get_attn_spec_kind +from vllm.v1.kv_cache_interface import KVCacheSpecKind + + +@pytest.mark.parametrize( + "signals,expected", + [ + (dict(use_mla=False, has_sliding_window=False), "full"), + (dict(use_mla=True, has_sliding_window=False), "mla"), + (dict(use_mla=True, has_sliding_window=True), "sw_mla"), + (dict(use_mla=False, has_sliding_window=True), "sw"), + ], +) +def test_get_attn_spec_kind_decoder(signals, expected): + kind_by_name = { + "full": KVCacheSpecKind.FULL_ATTENTION, + "mla": KVCacheSpecKind.MLA_ATTENTION, + "sw_mla": KVCacheSpecKind.SLIDING_WINDOW_MLA, + "sw": KVCacheSpecKind.SLIDING_WINDOW, + } + kind = get_attn_spec_kind(attn_type=AttentionType.DECODER, **signals) + assert kind is kind_by_name[expected] + + +@pytest.mark.parametrize( + "attn_type,expected", + [ + (AttentionType.ENCODER_ONLY, KVCacheSpecKind.ENCODER_ONLY_ATTENTION), + (AttentionType.ENCODER_DECODER, KVCacheSpecKind.CROSS_ATTENTION), + ], +) +def test_get_attn_spec_kind_attn_type(attn_type, expected): + kind = get_attn_spec_kind( + use_mla=False, + has_sliding_window=False, + attn_type=attn_type, + ) + assert kind is expected + + +def test_backend_per_kind_parses_strings(): + cfg = AttentionConfig( + backend_per_kind={ + "mla_attention": "FLASHINFER_MLA", + "sliding_window_mla": "triton_mla", # case-insensitive + } + ) + assert cfg.backend_per_kind["mla_attention"] is AttentionBackendEnum.FLASHINFER_MLA + assert cfg.backend_per_kind["sliding_window_mla"] is AttentionBackendEnum.TRITON_MLA + + +def test_backend_per_kind_rejects_unknown_kind(): + with pytest.raises(ValueError, match="Unknown KV cache group kind"): + AttentionConfig(backend_per_kind={"not_a_kind": "TRITON_MLA"}) + + +def test_backend_per_kind_defaults_empty(): + assert AttentionConfig().backend_per_kind == {} diff --git a/tests/v1/e2e/general/test_attention_backend_per_kind.py b/tests/v1/e2e/general/test_attention_backend_per_kind.py new file mode 100644 index 00000000000..3281247b1c4 --- /dev/null +++ b/tests/v1/e2e/general/test_attention_backend_per_kind.py @@ -0,0 +1,94 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""End-to-end checks that ``backend_per_kind`` selects the requested attention +backend for each KV-cache group at runtime. + +Uses ``google/gemma-3-1b-it``, which interleaves full-attention and +sliding-window layers, so the model produces separate ``full_attention`` and +``sliding_window`` KV-cache groups. +""" + +import pytest + +from vllm import LLM +from vllm.config.attention import AttentionConfig +from vllm.platforms import current_platform + +MODEL = "google/gemma-3-1b-it" + + +def _collect_group_backends(worker) -> list[tuple[str, str]]: + """Runs on the worker: returns (spec_kind, backend_name) per attn group.""" + from vllm.v1.kv_cache_interface import get_kv_cache_spec_kind + + out: list[tuple[str, str]] = [] + for kv_group in worker.model_runner.attn_groups: + for attn_group in kv_group: + kind = get_kv_cache_spec_kind(attn_group.kv_cache_spec) + out.append((kind.value, attn_group.backend.get_name())) + return out + + +@pytest.mark.skipif( + not current_platform.is_cuda(), reason="backend names are CUDA-specific" +) +@pytest.mark.parametrize( + "backend_per_kind", + [ + {"full_attention": "FLASH_ATTN", "sliding_window": "TRITON_ATTN"}, + # Swapped, to prove the mapping is causal rather than the default. + {"full_attention": "TRITON_ATTN", "sliding_window": "FLASH_ATTN"}, + ], +) +def test_backend_per_kind_splits_groups(backend_per_kind, monkeypatch): + # collective_rpc ships the callable to the EngineCore subprocess; the + # secure msgpack encoder can't serialize functions, so opt into the + # pickle fallback (same pattern as test_pooling_chunked_prefill). + monkeypatch.setenv("VLLM_ALLOW_INSECURE_SERIALIZATION", "1") + llm = LLM( + model=MODEL, + attention_config=AttentionConfig(backend_per_kind=backend_per_kind), + enforce_eager=True, + max_model_len=2048, + gpu_memory_utilization=0.4, + ) + + group_backends = llm.llm_engine.collective_rpc(_collect_group_backends)[0] + kinds = {kind for kind, _ in group_backends} + + # gemma3 must actually split into both kinds for this test to be meaningful. + assert "full_attention" in kinds + assert "sliding_window" in kinds + + for kind, backend_name in group_backends: + if kind in backend_per_kind: + assert backend_name == backend_per_kind[kind], ( + f"{kind} group used {backend_name}, expected {backend_per_kind[kind]}" + ) + + +@pytest.mark.skipif( + not current_platform.is_cuda(), reason="backend names are CUDA-specific" +) +def test_backend_per_kind_overrides_global_backend(monkeypatch): + """A per-kind entry wins over the global ``backend`` for its kind; other + kinds fall back to the global backend.""" + monkeypatch.setenv("VLLM_ALLOW_INSECURE_SERIALIZATION", "1") + llm = LLM( + model=MODEL, + attention_config=AttentionConfig( + backend="FLASH_ATTN", + backend_per_kind={"sliding_window": "TRITON_ATTN"}, + ), + enforce_eager=True, + max_model_len=2048, + gpu_memory_utilization=0.4, + ) + + group_backends = llm.llm_engine.collective_rpc(_collect_group_backends)[0] + + for kind, backend_name in group_backends: + if kind == "sliding_window": + assert backend_name == "TRITON_ATTN" + elif kind == "full_attention": + assert backend_name == "FLASH_ATTN" diff --git a/vllm/config/attention.py b/vllm/config/attention.py index e2ad4729a1e..78bc1f9537d 100644 --- a/vllm/config/attention.py +++ b/vllm/config/attention.py @@ -1,6 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from dataclasses import field from typing import Any, Literal from pydantic import field_validator @@ -19,6 +20,17 @@ class AttentionConfig: backend: AttentionBackendEnum | None = None """Attention backend to use. Use "auto" or None for automatic selection.""" + backend_per_kind: dict[str, AttentionBackendEnum] = field(default_factory=dict) + """Per-KV-cache-group attention backend overrides, keyed by + `KVCacheSpecKind` (e.g. `{"mla_attention": "FLASHINFER_MLA", + "sliding_window_mla": "TRITON_MLA"}`). This lets a model that splits its + layers across multiple KV-cache groups (e.g. interleaved full and + sliding-window attention) use a different backend per group. + + An entry overrides `backend` for layers of the matching kind; kinds not + listed fall back to `backend` (or automatic selection). A selected backend + that is invalid for that kind raises at startup.""" + flash_attn_version: Literal[2, 3, 4] | None = None """Force vllm to use a specific flash-attention version (2, 3, or 4). Only valid when using the flash-attention backend.""" @@ -120,3 +132,29 @@ class AttentionConfig: if isinstance(value, str): return MLAPrefillBackendEnum[value.upper()] return value + + @field_validator("backend_per_kind", mode="before") + @classmethod + def validate_backend_per_kind_before(cls, value: Any) -> Any: + """Parse the `backend_per_kind` map from strings. + + Keys must be valid `KVCacheSpecKind` values; values are parsed like + `backend` (enum name, case-insensitive). + """ + from vllm.v1.kv_cache_interface import KVCacheSpecKind + + if not isinstance(value, dict): + return value + valid_kinds = {kind.value for kind in KVCacheSpecKind} + parsed: dict[str, AttentionBackendEnum] = {} + for kind, backend in value.items(): + if kind not in valid_kinds: + raise ValueError( + f"Unknown KV cache group kind '{kind}' in " + f"backend_per_kind. Valid kinds are: " + f"{', '.join(sorted(valid_kinds))}." + ) + if isinstance(backend, str): + backend = AttentionBackendEnum[backend.upper()] + parsed[kind] = backend + return parsed diff --git a/vllm/v1/attention/selector.py b/vllm/v1/attention/selector.py index 917f2f6f709..387eac34c9f 100644 --- a/vllm/v1/attention/selector.py +++ b/vllm/v1/attention/selector.py @@ -2,7 +2,7 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from functools import cache -from typing import NamedTuple, cast, get_args +from typing import TYPE_CHECKING, NamedTuple, cast, get_args import torch @@ -15,6 +15,9 @@ from vllm.v1.attention.backends.registry import ( MambaAttentionBackendEnum, ) +if TYPE_CHECKING: + from vllm.v1.kv_cache_interface import KVCacheSpecKind + logger = init_logger(__name__) @@ -53,6 +56,46 @@ class AttentionSelectorConfig(NamedTuple): ) +def get_attn_spec_kind( + use_mla: bool, + has_sliding_window: bool, + attn_type: str, +) -> "KVCacheSpecKind": + """Derive the KV-cache group kind a layer belongs to from its signals. + + Mirrors ``get_kv_cache_spec_kind`` (which derives the kind from the + produced ``KVCacheSpec``) so users can target groups by kind when + setting ``AttentionConfig.backend_per_kind``. + + ``SINK_FULL_ATTENTION`` is intentionally not derived here: it is produced + only by the ``StaticSinkAttention`` layer, whereas a plain ``Attention`` + layer with attention sinks (e.g. gpt-oss) still yields a + ``FullAttentionSpec``/``SlidingWindowSpec``. Sinks therefore do not change + the kind. + + Args: + use_mla: Whether the layer uses multi-head latent attention. + has_sliding_window: Whether the layer applies a sliding window. + attn_type: The layer's ``AttentionType``. + + Returns: + The ``KVCacheSpecKind`` the layer maps to. + """ + from vllm.v1.kv_cache_interface import KVCacheSpecKind + + if attn_type == AttentionType.ENCODER_ONLY: + return KVCacheSpecKind.ENCODER_ONLY_ATTENTION + if attn_type == AttentionType.ENCODER_DECODER: + return KVCacheSpecKind.CROSS_ATTENTION + if use_mla: + if has_sliding_window: + return KVCacheSpecKind.SLIDING_WINDOW_MLA + return KVCacheSpecKind.MLA_ATTENTION + if has_sliding_window: + return KVCacheSpecKind.SLIDING_WINDOW + return KVCacheSpecKind.FULL_ATTENTION + + def get_attn_backend( head_size: int, dtype: torch.dtype, @@ -91,6 +134,7 @@ def get_attn_backend( kv_transfer_config is not None and kv_transfer_config.is_kv_transfer_instance ) + attn_type = attn_type or AttentionType.DECODER attn_selector_config = AttentionSelectorConfig( head_size=head_size, dtype=dtype, @@ -101,15 +145,27 @@ def get_attn_backend( use_sparse=use_sparse, use_mm_prefix=use_mm_prefix, use_per_head_quant_scales=use_per_head_quant_scales, - attn_type=attn_type or AttentionType.DECODER, + attn_type=attn_type, has_sliding_window=has_sliding_window, use_non_causal=vllm_config.attention_config.use_non_causal, use_batch_invariant=envs.VLLM_BATCH_INVARIANT, use_kv_connector=use_kv_connector, ) + # A per-KV-group override (keyed by KVCacheSpecKind) takes precedence over + # the global backend; kinds not present in the map fall back to it. + attention_config = vllm_config.attention_config + backend = attention_config.backend + if attention_config.backend_per_kind: + kind = get_attn_spec_kind( + use_mla=use_mla, + has_sliding_window=has_sliding_window, + attn_type=attn_type, + ) + backend = attention_config.backend_per_kind.get(kind.value, backend) + return _cached_get_attn_backend( - backend=vllm_config.attention_config.backend, + backend=backend, attn_selector_config=attn_selector_config, num_heads=num_heads, ) From c4cd2bd5440bca29039e5b6f33b695a95721bd3e Mon Sep 17 00:00:00 2001 From: limeward <32970461+edwinlim0919@users.noreply.github.com> Date: Fri, 17 Jul 2026 12:35:04 -0700 Subject: [PATCH 06/51] [Bugfix] MoRIIO toy P/D proxy: fix DP-rank index aliasing + harden for high-concurrency bursts (#46115) Signed-off-by: Edwin Lim Co-authored-by: Claude Co-authored-by: QinPR <1905873179@qq.com> Co-authored-by: Peiran Qin <66068739+QinPR@users.noreply.github.com> --- .../moriio_toy_proxy_server.py | 76 +++++-- .../unit/test_moriio_proxy_routing.py | 205 ++++++++++++++++++ 2 files changed, 268 insertions(+), 13 deletions(-) create mode 100644 tests/v1/kv_connector/unit/test_moriio_proxy_routing.py diff --git a/examples/disaggregated/disaggregated_serving/moriio_toy_proxy_server.py b/examples/disaggregated/disaggregated_serving/moriio_toy_proxy_server.py index 6f7d210ff3c..d2d5be76077 100644 --- a/examples/disaggregated/disaggregated_serving/moriio_toy_proxy_server.py +++ b/examples/disaggregated/disaggregated_serving/moriio_toy_proxy_server.py @@ -210,8 +210,31 @@ async def stream_decode_response(session, response, request_id): await session.close() -def example_round_robin_dp_loader(request_number, dp_size): - return request_nums % dp_size +def flat_interleaved_dp_route(request_number, instances): + """Flat round-robin over the full (instance, dp_rank) slot space. + + ONE counter over (n_instances * dp_size) slots, so instance-selection and + DP-rank-selection are derived from the SAME index and can never alias. The + previous scheme computed instance = req % n and rank = req % dp from the + same counter with n | dp, which locked each instance to a stride-n subset + of its ranks (e.g. 2 prefill instances -> 4 of 8 ranks each -> half the + GPUs never receive a request, so the deployment falsely appears not to + scale). + + Interleaved order — inst0_r0, inst1_r0, inst0_r1, inst1_r1, ... — so + consecutive requests alternate instances AND every rank gets walked. + + Assumes homogeneous dp_size across a role's instances (true for the + DP<->DP and DP<->TP deployments this proxy targets). Returns + (instance_index, dp_rank); dp_rank is None when dp_size == 1 (e.g. a TP + decode), which avoids forwarding an out-of-range data-parallel rank. + """ + n = len(instances) + dp = instances[0]["dp_size"] + slot = (request_number - 1) % (n * dp) + inst_idx = slot % n + dp_rank = (slot // n) if dp > 1 else None + return inst_idx, dp_rank @app.route("/health", methods=["GET"]) @@ -252,18 +275,21 @@ async def handle_request(api: str, request: Request): 503, ) ) - pid = request_nums % len(prefill_instances) - did = request_nums % len(decode_instances) + # Flat interleaved round-robin (see flat_interleaved_dp_route): ONE + # counter over the full (instance, dp_rank) slot space per role, so + # instance-selection and DP-rank-selection derive from the same index + # and can never alias. The old scheme keyed both on request_nums with + # n_instances | dp_size, stranding half the ranks (e.g. in 2P_DP8EP). + pid, selected_prefill_dp_rank = flat_interleaved_dp_route( + request_nums, prefill_instances + ) + # Decode instance selection uses the same interleaved walk; in READ + # mode the decode reads KV from selected_prefill_dp_rank, so the + # decode's own dp_rank is not forwarded here. + did, _ = flat_interleaved_dp_route(request_nums, decode_instances) prefill_instance_endpoint = prefill_instances[pid] decode_instance_endpoint = decode_instances[did] - selected_prefill_dp_rank = None - if prefill_instance_endpoint["dp_size"] > 1: - selected_prefill_dp_rank = example_round_robin_dp_loader( - request_nums // len(prefill_instance_endpoint), - prefill_instance_endpoint["dp_size"], - ) - # Embed both zmq_addresses in the request_id so the connector can parse # the peer's host/ports from it, similar to P2P-NCCL uid = str(uuid.uuid4()).replace("-", "") @@ -427,9 +453,33 @@ if __name__ == "__main__": args = parser.parse_args() t = start_service_discovery("0.0.0.0", 36367) - app.debug = True + # High-concurrency hardening. Quart's app.run() uses a shallow listen + # backlog (100) and, with app.debug=True, adds per-request overhead that + # starves the single accept loop. Under a burst of ~512 simultaneous client + # connections the backlog overflows and the kernel RSTs the excess, so + # clients see "ClientOSError: [Errno 104] Connection reset by peer" before + # any response (~16% request loss at c=512). Serve via hypercorn with debug + # OFF and a deep backlog so the burst QUEUES (higher TTFT) instead of being + # reset -> 100% request success. + app.debug = False app.config["BODY_TIMEOUT"] = 360000 app.config["RESPONSE_TIMEOUT"] = 360000 - app.run(host="0.0.0.0", port=args.port) + import asyncio + import os + + from hypercorn.asyncio import serve as _hypercorn_serve + from hypercorn.config import Config as _HypercornConfig + + _hcfg = _HypercornConfig() + _hcfg.bind = [f"0.0.0.0:{args.port}"] + # Deep listen backlog so a wide connection burst queues, not RSTs. NOTE: + # effective backlog is capped by the host's net.core.somaxconn (proxy runs + # --network host); kernel 6.x defaults to 4096. Override via + # PROXY_LISTEN_BACKLOG. + _hcfg.backlog = int(os.environ.get("PROXY_LISTEN_BACKLOG", "4096")) + # Long-lived SSE streams (8k1k decode ~5 min): never reap on keepalive. + _hcfg.keep_alive_timeout = 360000.0 + + asyncio.run(_hypercorn_serve(app, _hcfg)) t.join() diff --git a/tests/v1/kv_connector/unit/test_moriio_proxy_routing.py b/tests/v1/kv_connector/unit/test_moriio_proxy_routing.py new file mode 100644 index 00000000000..71ee0ef5d79 --- /dev/null +++ b/tests/v1/kv_connector/unit/test_moriio_proxy_routing.py @@ -0,0 +1,205 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Hardware-fair request routing in the MoRIIO toy P/D proxy. + +Exercises the REAL ``flat_interleaved_dp_route`` from the toy proxy (loaded from +its ``examples/`` path with heavy deps stubbed) — no routing logic is copied +here. This is the request-distribution half of end-to-end fairness: the proxy +must hand every prefill/decode (instance, dp_rank) slot an equal share so no GPU +is starved. The connector-side read routing is covered in +``test_moriio_routing_fairness.py``. + +Regression target: the previous scheme derived ``instance = req % n_instances`` +and ``dp_rank = req % dp_size`` from the same counter, so when +``n_instances | dp_size`` each instance was locked to a stride-``n`` subset of +its ranks — e.g. 2 prefill instances x DP8 stranded 4 of every node's 8 GPUs. +``flat_interleaved_dp_route`` walks ONE counter over the full +``(instance, dp_rank)`` slot space, so the two selections can never alias. + +Role shapes below are exactly those in the RFC (#46107) deployments: + (1, 1) 1P/1D TP8 (2, 1) 2P/2D TP8 (4, 1) 4D TP8 + (1, 8) 1D DP8EP (2, 8) 2P DP8EP (3, 8) 3D DP8EP +""" + +import contextlib +import importlib.util +import sys +import types +from collections import Counter +from pathlib import Path +from typing import cast + +import pytest + +_MISSING = object() + +PROXY_REL = "examples/disaggregated/disaggregated_serving/moriio_toy_proxy_server.py" + + +def _module(name, **attrs): + module = types.ModuleType(name) + for attr, value in attrs.items(): + setattr(module, attr, value) + return module + + +def _package(name): + module = _module(name) + module.__path__ = [] + return module + + +class _QuartStub: + def __init__(self, *a, **k): + pass + + def route(self, *a, **k): + return lambda fn: fn + + def post(self, *a, **k): + return self.route() + + +async def _make_response_stub(value): + return value + + +@contextlib.contextmanager +def _proxy_import_stubs(): + """Stub the proxy's external deps so the module imports for a pure unit test. + + Only third-party/vllm imports are stubbed; the routing function under test + is executed as-is from the real module. + """ + common = "vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_common" + + class _MoRIIOConstants: + TRANSFER_PREFIX = "moriio-transfer" + + stubs = { + "aiohttp": _module("aiohttp"), + "msgpack": _module("msgpack"), + "zmq": _module("zmq"), + "quart": _module( + "quart", + Quart=_QuartStub, + Request=object, + make_response=_make_response_stub, + request=object(), + ), + "vllm": _package("vllm"), + "vllm.distributed": _package("vllm.distributed"), + "vllm.distributed.kv_transfer": _package("vllm.distributed.kv_transfer"), + "vllm.distributed.kv_transfer.kv_connector": _package( + "vllm.distributed.kv_transfer.kv_connector" + ), + "vllm.distributed.kv_transfer.kv_connector.v1": _package( + "vllm.distributed.kv_transfer.kv_connector.v1" + ), + "vllm.distributed.kv_transfer.kv_connector.v1.moriio": _package( + "vllm.distributed.kv_transfer.kv_connector.v1.moriio" + ), + common: _module(common, MoRIIOConstants=_MoRIIOConstants), + } + saved = {} + for name, module in stubs.items(): + if name not in sys.modules: + saved[name] = _MISSING + sys.modules[name] = module + try: + yield + finally: + for name, previous in saved.items(): + if previous is _MISSING: + sys.modules.pop(name, None) + else: + sys.modules[name] = cast(types.ModuleType, previous) + + +def _load_proxy_module(): + path = Path(__file__).parents[4] / PROXY_REL + spec = importlib.util.spec_from_file_location("moriio_proxy_under_test", path) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + with _proxy_import_stubs(): + spec.loader.exec_module(module) + return module + + +@pytest.fixture(scope="module") +def route(): + return _load_proxy_module().flat_interleaved_dp_route + + +def _instances(n: int, dp_size: int): + # Only dp_size is read by the router; tp_size carried for realism. + return [{"dp_size": dp_size, "tp_size": 8 // dp_size} for _ in range(n)] + + +# (n_instances, dp_size) for every distinct P/D role shape in the RFC configs. +ROLE_SHAPES = [(1, 1), (2, 1), (4, 1), (1, 8), (2, 8), (3, 8)] +SHAPE_IDS = [f"n{n}_dp{dp}" for n, dp in ROLE_SHAPES] + + +def _route_n(route, instances, count): + # Proxy uses 1-indexed request numbers (slot = (request_number - 1) % ...). + return [route(rn, instances) for rn in range(1, count + 1)] + + +@pytest.mark.parametrize(("n", "dp"), ROLE_SHAPES, ids=SHAPE_IDS) +def test_full_slot_space_is_covered_uniformly(route, n, dp): + instances = _instances(n, dp) + period = n * dp + # Three full cycles -> every (instance, dp_rank) slot must be hit the same + # number of times (exactly uniform, no starved slot, no aliasing). + hits = Counter(_route_n(route, instances, period * 3)) + + expected: set[tuple[int, int | None]] + if dp == 1: + expected = {(inst, None) for inst in range(n)} + else: + expected = {(inst, r) for inst in range(n) for r in range(dp)} + assert set(hits) == expected, f"missing slots: {expected - set(hits)}" + assert max(hits.values()) == min(hits.values()) + + +@pytest.mark.parametrize(("n", "dp"), ROLE_SHAPES, ids=SHAPE_IDS) +def test_instance_and_dp_rank_marginals_are_balanced(route, n, dp): + instances = _instances(n, dp) + routed = _route_n(route, instances, n * dp * 5) + + inst_counts = Counter(inst for inst, _ in routed) + assert set(inst_counts) == set(range(n)) + assert max(inst_counts.values()) == min(inst_counts.values()) + + dp_counts = Counter(r for _, r in routed) + if dp == 1: + assert set(dp_counts) == {None} + else: + assert set(dp_counts) == set(range(dp)) + assert max(dp_counts.values()) == min(dp_counts.values()) + + +def test_tp_instance_forwards_no_dp_rank(route): + # dp_size == 1 (a TP instance) must yield dp_rank None so the proxy never + # forwards an out-of-range data-parallel rank. + assert all(r is None for _, r in _route_n(route, _instances(2, 1), 8)) + + +def test_two_instance_dp8_gives_every_node_all_ranks(route): + # Direct regression for the stranded-GPU bug: with 2 prefill instances x DP8 + # each instance must receive ALL 8 dp ranks (16 distinct slots), not a + # stride-2 subset of 4. + routed = _route_n(route, _instances(2, 8), 2 * 8) + per_instance: dict[int, set] = {0: set(), 1: set()} + for inst, dp_rank in routed: + per_instance[inst].add(dp_rank) + assert per_instance[0] == set(range(8)) + assert per_instance[1] == set(range(8)) + + +def test_consecutive_requests_alternate_instances(route): + # Interleaved order spreads consecutive requests across instances rather + # than filling one instance's ranks before moving on. + insts = [inst for inst, _ in _route_n(route, _instances(3, 8), 3)] + assert insts == [0, 1, 2] From cc25f028b76844024fddd0b4ca18d4ca169b32be Mon Sep 17 00:00:00 2001 From: Michael Goin Date: Fri, 17 Jul 2026 16:30:02 -0400 Subject: [PATCH 07/51] [Loader] Improve InstantTensor loading (#46868) Signed-off-by: mgoin Co-authored-by: OpenAI Codex --- requirements/test/cpu.txt | 2 +- requirements/test/cuda.in | 2 +- requirements/test/cuda.txt | 2 +- requirements/test/nightly-torch.txt | 2 +- requirements/test/rocm.in | 2 +- requirements/test/rocm.txt | 2 +- setup.py | 2 +- .../instanttensor_loader/test_weight_utils.py | 2 -- .../model_loader/weight_utils.py | 25 +++++++++++++++---- 9 files changed, 27 insertions(+), 14 deletions(-) diff --git a/requirements/test/cpu.txt b/requirements/test/cpu.txt index 143155ad7a6..d7cec776697 100644 --- a/requirements/test/cpu.txt +++ b/requirements/test/cpu.txt @@ -379,7 +379,7 @@ inflect==5.6.2 # via datamodel-code-generator iniconfig==2.0.0 # via pytest -instanttensor==0.1.5 +instanttensor==0.1.9 # via -r requirements/test/cuda.in interegular==0.3.3 # via lm-format-enforcer diff --git a/requirements/test/cuda.in b/requirements/test/cuda.in index bd6f179c105..d91cba0601a 100644 --- a/requirements/test/cuda.in +++ b/requirements/test/cuda.in @@ -58,7 +58,7 @@ arctic-inference == 0.1.1; platform_machine == "x86_64" # Required for suffix de numba == 0.65.0 # Required for N-gram speculative decoding runai-model-streamer[s3,gcs,azure]==0.15.7 fastsafetensors>=0.3.2 -instanttensor>=0.1.5; platform_machine == "x86_64" +instanttensor>=0.1.9; platform_machine == "x86_64" decord==0.6.0; platform_machine == "x86_64" # terratorch is temporarily disabled while PyPI has the `lightning` package # in `quarantined` status (every published terratorch version transitively diff --git a/requirements/test/cuda.txt b/requirements/test/cuda.txt index e4da9ac06e7..c4a9ea1ae2d 100644 --- a/requirements/test/cuda.txt +++ b/requirements/test/cuda.txt @@ -398,7 +398,7 @@ inflect==5.6.2 # via datamodel-code-generator iniconfig==2.0.0 # via pytest -instanttensor==0.1.5 +instanttensor==0.1.9 # via -r requirements/test/cuda.in interegular==0.3.3 # via lm-format-enforcer diff --git a/requirements/test/nightly-torch.txt b/requirements/test/nightly-torch.txt index dfefa8239c5..b308a440139 100644 --- a/requirements/test/nightly-torch.txt +++ b/requirements/test/nightly-torch.txt @@ -44,5 +44,5 @@ numba == 0.65.0 # Required for N-gram speculative decoding numpy runai-model-streamer[s3,gcs,azure]==0.15.7 fastsafetensors>=0.3.2 -instanttensor>=0.1.5 +instanttensor>=0.1.9 pydantic>=2.12 # 2.11 leads to error on python 3.13 diff --git a/requirements/test/rocm.in b/requirements/test/rocm.in index c5f8f85f2f3..7cb3b91f7bd 100644 --- a/requirements/test/rocm.in +++ b/requirements/test/rocm.in @@ -54,7 +54,7 @@ arctic-inference==0.1.1 # Required for suffix decoding test numba==0.65.0 # Required for N-gram speculative decoding runai-model-streamer[s3,gcs,azure]==0.15.7 fastsafetensors>=0.3.2 -instanttensor>=0.1.5 +instanttensor>=0.1.9 decord==0.6.0 # Prithvi tests diff --git a/requirements/test/rocm.txt b/requirements/test/rocm.txt index 6cd27abe17b..f274a93362c 100644 --- a/requirements/test/rocm.txt +++ b/requirements/test/rocm.txt @@ -391,7 +391,7 @@ inflect==7.5.0 # via datamodel-code-generator iniconfig==2.3.0 # via pytest -instanttensor==0.1.6 +instanttensor==0.1.9 # via -r requirements/test/rocm.in interegular==0.3.3 # via lm-format-enforcer diff --git a/setup.py b/setup.py index a8685157933..40d4ca103be 100644 --- a/setup.py +++ b/setup.py @@ -1268,7 +1268,7 @@ setup( "bench": ["pandas", "matplotlib", "seaborn", "datasets", "scipy", "plotly"], "tensorizer": ["tensorizer==2.10.1"], "fastsafetensors": ["fastsafetensors >= 0.3.2"], - "instanttensor": ["instanttensor >= 0.1.5"], + "instanttensor": ["instanttensor >= 0.1.9"], "runai": ["runai-model-streamer[s3,gcs,azure] >= 0.15.7"], "audio": [ "av", diff --git a/tests/model_executor/model_loader/instanttensor_loader/test_weight_utils.py b/tests/model_executor/model_loader/instanttensor_loader/test_weight_utils.py index 992a83e0eea..09e0a19d11b 100644 --- a/tests/model_executor/model_loader/instanttensor_loader/test_weight_utils.py +++ b/tests/model_executor/model_loader/instanttensor_loader/test_weight_utils.py @@ -33,8 +33,6 @@ def test_instanttensor_model_loader(): hf_safetensors_tensors = {} for name, tensor in instanttensor_weights_iterator(safetensors, True): - # Copy the tensor immediately as it is a reference to the internal - # buffer of instanttensor. instanttensor_tensors[name] = tensor.to("cpu") for name, tensor in safetensors_weights_iterator(safetensors, True): diff --git a/vllm/model_executor/model_loader/weight_utils.py b/vllm/model_executor/model_loader/weight_utils.py index 8e673ae7fcb..db161a58988 100644 --- a/vllm/model_executor/model_loader/weight_utils.py +++ b/vllm/model_executor/model_loader/weight_utils.py @@ -1120,7 +1120,7 @@ def instanttensor_weights_iterator( import instanttensor except ImportError as e: raise ImportError( - "Please install instanttensor via `pip install instanttensor`" + "Please install instanttensor via `pip install vllm[instanttensor]`" ) from e if not current_platform.is_cuda(): @@ -1136,18 +1136,33 @@ def instanttensor_weights_iterator( device = current_platform.current_device() + # copy=True yields tensors that own their memory, staying valid after the + # context exits or InstantTensor reuses its buffer. with instanttensor.safe_open( - hf_weights_files, framework="pt", device=device, process_group=process_group + hf_weights_files, + framework="pt", + device=device, + process_group=process_group, + copy=True, ) as f: - yield from tqdm( - f.tensors(), + # Track bytes so the bar reports load throughput (GB/s). + pbar = tqdm( + total=f.total_tensor_size, desc="Loading safetensors using InstantTensor loader", disable=not enable_tqdm(use_tqdm_on_load), bar_format=_BAR_FORMAT, position=tqdm._get_free_pos(), - total=len(f.keys()), + unit="B", + unit_scale=True, + unit_divisor=1024, mininterval=1.0, ) + try: + for name, tensor in f.tensors(): + pbar.update(tensor.numel() * tensor.element_size()) + yield name, tensor + finally: + pbar.close() def pt_weights_iterator( From b5433b6f5079feb32f9f278cf4ae23bd87375148 Mon Sep 17 00:00:00 2001 From: Wentao Ye <44945378+yewentao256@users.noreply.github.com> Date: Fri, 17 Jul 2026 16:35:06 -0400 Subject: [PATCH 08/51] [Perf] Optimize dsv4 routing using specialized kernel, 2.94% E2E TPOT improvement (#48660) Signed-off-by: yewentao256 --- .../moe/topk_softplus_sqrt_kernels.cu | 78 +++++++++++ tests/kernels/moe/test_topk_softplus_sqrt.py | 45 +++++++ .../layers/fused_moe/router/dsv4_topk.py | 121 ++++++++++++++++++ .../router/fused_topk_bias_router.py | 20 +++ 4 files changed, 264 insertions(+) create mode 100644 vllm/model_executor/layers/fused_moe/router/dsv4_topk.py diff --git a/csrc/libtorch_stable/moe/topk_softplus_sqrt_kernels.cu b/csrc/libtorch_stable/moe/topk_softplus_sqrt_kernels.cu index 785bbf2f6e0..b6878eb2d2f 100644 --- a/csrc/libtorch_stable/moe/topk_softplus_sqrt_kernels.cu +++ b/csrc/libtorch_stable/moe/topk_softplus_sqrt_kernels.cu @@ -71,6 +71,73 @@ __device__ __forceinline__ float toFloat(T value) { } } +#ifndef USE_ROCM +// Adapted from: +// https://github.com/sgl-project/sglang/blob/main/python/sglang/jit_kernel/csrc/deepseek_v4/hash_topk.cuh +template +__launch_bounds__(128) __global__ + void dsv4HashTopkSoftplusSqrt(const float* input, float* output, + OutIndType* indices, int num_rows, + int num_experts, float routed_scaling_factor, + const HashIndType* input_ids, + const HashIndType* tid2eid) { + const int warp = (blockIdx.x * blockDim.x + threadIdx.x) / 32; + const int lane = threadIdx.x % 32; + if (warp >= num_rows) return; + const int64_t token_id = load_index_as_int64(input_ids, warp); + + #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900) + cudaGridDependencySynchronize(); + #endif + int expert = 0; + float weight = 0.f; + if (lane < 6) { + // only load and calculate for 6 experts + expert = static_cast(tid2eid[token_id * 6 + lane]); + const float x = input[warp * num_experts + expert]; + weight = sqrtf(fmaxf(x, 0.f) + __logf(1.f + __expf(-fabsf(x)))); + } + float weight_sum = weight; + #pragma unroll + for (int mask = 16; mask > 0; mask >>= 1) { + // sum in warp + weight_sum += VLLM_SHFL_XOR_SYNC(weight_sum, mask); + } + + #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900) + cudaTriggerProgrammaticLaunchCompletion(); + #endif + if (lane < 6) { + const int offset = warp * 6 + lane; + output[offset] = + weight * routed_scaling_factor / (weight_sum > 0.f ? weight_sum : 1.f); + indices[offset] = static_cast(expert); + } +} + +template +void launchDsv4HashTopk(const float* input, float* output, OutIndType* indices, + int num_rows, int num_experts, + double routed_scaling_factor, + const HashIndType* input_ids, + const HashIndType* tid2eid, cudaStream_t stream) { + if (num_rows == 0) return; + auto* kernel = &dsv4HashTopkSoftplusSqrt; + cudaLaunchConfig_t config = {}; + config.gridDim = (num_rows + 3) / 4; + config.blockDim = 128; + config.stream = stream; + cudaLaunchAttribute attr; + attr.id = cudaLaunchAttributeProgrammaticStreamSerialization; + attr.val.programmaticStreamSerializationAllowed = 1; + config.attrs = &attr; + config.numAttrs = 1; + const float scale = static_cast(routed_scaling_factor); + cudaLaunchKernelEx(&config, kernel, input, output, indices, num_rows, + num_experts, scale, input_ids, tid2eid); +} +#endif + // ====================== TopK softplus_sqrt things // =============================== @@ -556,6 +623,17 @@ void topkGatingSoftplusSqrtKernelLauncher( const float* correction_bias, const bool use_hash, const HashIndType* input_ids, const HashIndType* tid2eid, cudaStream_t stream) { +#ifndef USE_ROCM + if constexpr (std::is_same_v) { + if (use_hash && topk == 6 && renormalize && + (num_experts == 256 || num_experts == 384)) { + launchDsv4HashTopk( + gating_output, topk_weights, topk_indices, num_tokens, num_experts, + routed_scaling_factor, input_ids, tid2eid, stream); + return; + } + } +#endif static constexpr int WARPS_PER_TB = 4; static constexpr int BYTES_PER_LDG_POWER_OF_2 = 16; // for bfloat16 dtype, we need 4 bytes loading to make sure num_experts diff --git a/tests/kernels/moe/test_topk_softplus_sqrt.py b/tests/kernels/moe/test_topk_softplus_sqrt.py index 46ca934c146..5c6691bbc32 100644 --- a/tests/kernels/moe/test_topk_softplus_sqrt.py +++ b/tests/kernels/moe/test_topk_softplus_sqrt.py @@ -10,6 +10,7 @@ from vllm.model_executor.layers.fused_moe.config import ( RoutingMethodType, get_routing_method_type, ) +from vllm.model_executor.layers.fused_moe.router.dsv4_topk import dsv4_topk from vllm.model_executor.layers.fused_moe.router.fused_topk_bias_router import ( fused_topk_bias, ) @@ -186,3 +187,47 @@ def test_fused_topk_softplus_sqrt_hash( sorted_w_ref = topk_weights_ref.gather(1, idx_ref) sorted_w = topk_weights.gather(1, idx_ops) torch.testing.assert_close(sorted_w_ref, sorted_w, atol=2e-2, rtol=1e-2) + + +@pytest.mark.skipif( + not current_platform.is_cuda(), + reason="The DeepSeek V4 fast path is CUDA-only.", +) +@pytest.mark.parametrize( + ("num_tokens", "num_experts", "indices_type"), + [ + (0, 256, torch.uint32), + (17, 256, torch.uint32), + (17, 384, torch.int64), + ], +) +def test_dsv4_fast_topk( + num_tokens: int, + num_experts: int, + indices_type: torch.dtype, +): + torch.manual_seed(0) + gating_output = torch.randn( + (num_tokens, num_experts), dtype=torch.float32, device="cuda" + ) + correction_bias = torch.randn(num_experts, dtype=torch.float32, device="cuda") + + topk_weights_ref, topk_ids_ref = _torch_topk_softplus_sqrt( + gating_output=gating_output, + topk=6, + renormalize=True, + routed_scaling_factor=1.5, + e_score_correction_bias=correction_bias, + ) + topk_weights, topk_ids = dsv4_topk( + gating_output, correction_bias, indices_type, 1.5 + ) + + assert topk_ids.dtype == indices_type + torch.testing.assert_close(topk_ids_ref.to(indices_type), topk_ids, atol=0, rtol=0) + torch.testing.assert_close( + topk_weights_ref, + topk_weights, + atol=2e-5, + rtol=2e-5, + ) diff --git a/vllm/model_executor/layers/fused_moe/router/dsv4_topk.py b/vllm/model_executor/layers/fused_moe/router/dsv4_topk.py new file mode 100644 index 00000000000..dd861d4ebcf --- /dev/null +++ b/vllm/model_executor/layers/fused_moe/router/dsv4_topk.py @@ -0,0 +1,121 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import torch + +from vllm.platforms import current_platform +from vllm.triton_utils import tl, triton + +_TOPK = 6 + +# Adapted from: +# https://github.com/sgl-project/sglang/blob/main/python/sglang/jit_kernel/moe_fused_gate.py + + +def can_use_dsv4_topk( + gating_output: torch.Tensor, + correction_bias: torch.Tensor | None, + topk: int, + renormalize: bool, + indices_dtype: torch.dtype, +) -> bool: + return ( + current_platform.is_cuda() + and gating_output.dtype == torch.float32 + and gating_output.ndim == 2 + and gating_output.shape[1] in (256, 384) + and gating_output.is_contiguous() + and correction_bias is not None + and correction_bias.dtype == torch.float32 + and correction_bias.shape == (gating_output.shape[1],) + and correction_bias.is_contiguous() + and topk == _TOPK + and renormalize + and indices_dtype in (torch.int32, torch.uint32, torch.int64) + ) + + +if current_platform.is_cuda(): + + @triton.jit + def _dsv4_topk_kernel( + gating_output_ptr, + correction_bias_ptr, + topk_weights_ptr, + topk_ids_ptr, + routed_scaling_factor, + NUM_EXPERTS: tl.constexpr, + BLOCK_N: tl.constexpr, + launch_pdl: tl.constexpr, + ): + row = tl.program_id(0) + expert_offsets = tl.arange(0, BLOCK_N) + expert_mask = expert_offsets < NUM_EXPERTS + bias = tl.load( + correction_bias_ptr + expert_offsets, mask=expert_mask, other=0.0 + ).to(tl.float32) + + if launch_pdl: + tl.extra.cuda.gdc_wait() + + logits = tl.load( + gating_output_ptr + row * NUM_EXPERTS + expert_offsets, + mask=expert_mask, + other=0.0, + ).to(tl.float32) + weights = tl.sqrt(tl.where(logits > 20.0, logits, tl.log(1.0 + tl.exp(logits)))) + current = tl.where(expert_mask, weights + bias, -float("inf")) + current = tl.where(current == current, current, -1e30) + + topk_offsets = tl.arange(0, 8) + selected_weights = tl.zeros([8], dtype=tl.float32) + selected_ids = tl.zeros([8], dtype=tl.int32) + for slot in tl.static_range(6): + max_value = tl.max(current, axis=0) + candidate = tl.where(current == max_value, expert_offsets, NUM_EXPERTS) + expert_id = tl.min(candidate, axis=0).to(tl.int32) + selected_weight = tl.sum( + tl.where(expert_offsets == expert_id, weights, 0.0), axis=0 + ) + is_slot = topk_offsets == slot + selected_weights = tl.where(is_slot, selected_weight, selected_weights) + selected_ids = tl.where(is_slot, expert_id, selected_ids) + current = tl.where(expert_offsets == expert_id, -float("inf"), current) + + weight_sum = tl.sum(selected_weights, axis=0) + selected_weights *= routed_scaling_factor / tl.where( + weight_sum > 0.0, weight_sum, 1.0 + ) + output_mask = topk_offsets < 6 + output_offsets = row * 6 + topk_offsets + + if launch_pdl: + tl.extra.cuda.gdc_launch_dependents() + + tl.store(topk_weights_ptr + output_offsets, selected_weights, mask=output_mask) + tl.store(topk_ids_ptr + output_offsets, selected_ids, mask=output_mask) + + +def dsv4_topk( + gating_output: torch.Tensor, + correction_bias: torch.Tensor, + indices_dtype: torch.dtype, + routed_scaling_factor: float, +) -> tuple[torch.Tensor, torch.Tensor]: + num_tokens, num_experts = gating_output.shape + shape = (num_tokens, _TOPK) + topk_weights = gating_output.new_empty(shape, dtype=torch.float32) + topk_ids = gating_output.new_empty(shape, dtype=indices_dtype) + if num_tokens > 0: + _dsv4_topk_kernel[(num_tokens,)]( + gating_output, + correction_bias, + topk_weights, + topk_ids, + routed_scaling_factor, + NUM_EXPERTS=num_experts, + BLOCK_N=triton.next_power_of_2(num_experts), + num_warps=1, + launch_pdl=current_platform.is_arch_support_pdl(), + ) + return topk_weights, topk_ids diff --git a/vllm/model_executor/layers/fused_moe/router/fused_topk_bias_router.py b/vllm/model_executor/layers/fused_moe/router/fused_topk_bias_router.py index 1ddcaa50e83..f89eb4910b3 100644 --- a/vllm/model_executor/layers/fused_moe/router/fused_topk_bias_router.py +++ b/vllm/model_executor/layers/fused_moe/router/fused_topk_bias_router.py @@ -14,6 +14,10 @@ from vllm.model_executor.layers.fused_moe.config import ( get_routing_method_type, ) from vllm.model_executor.layers.fused_moe.router.base_router import BaseRouter +from vllm.model_executor.layers.fused_moe.router.dsv4_topk import ( + can_use_dsv4_topk, + dsv4_topk, +) def vllm_topk_softmax( @@ -182,6 +186,22 @@ def fused_topk_bias( "Number of tokens mismatch" ) + output_indices_dtype = torch.int32 if indices_type is None else indices_type + if scoring_func == "sqrtsoftplus" and can_use_dsv4_topk( + gating_output, + e_score_correction_bias, + topk, + renormalize, + output_indices_dtype, + ): + assert e_score_correction_bias is not None + return dsv4_topk( + gating_output, + e_score_correction_bias, + output_indices_dtype, + routed_scaling_factor, + ) + M, _ = hidden_states.size() topk_weights = torch.empty( From fcd2255d16bd3c62493bae5dee769ce998098f21 Mon Sep 17 00:00:00 2001 From: devalshahamd Date: Fri, 17 Jul 2026 13:38:59 -0700 Subject: [PATCH 09/51] [Hardware][GPU] Profiler config additional to increase it scope and annotation details (#37524) Signed-off-by: devalshahamd Signed-off-by: Deval Shah Signed-off-by: Deval Shah Co-authored-by: Deval Shah --- tests/v1/worker/test_gpu_profiler.py | 64 ++++++++++++++- vllm/config/profiler.py | 17 ++++ vllm/v1/worker/gpu_model_runner.py | 74 +++++++++++++++--- vllm/v1/worker/gpu_worker.py | 111 +++++++++++++++++++++++---- 4 files changed, 240 insertions(+), 26 deletions(-) diff --git a/tests/v1/worker/test_gpu_profiler.py b/tests/v1/worker/test_gpu_profiler.py index ca22f3c9da6..8ff89354ea9 100644 --- a/tests/v1/worker/test_gpu_profiler.py +++ b/tests/v1/worker/test_gpu_profiler.py @@ -1,10 +1,15 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from unittest.mock import MagicMock + import pytest -from vllm.config import ProfilerConfig +from vllm.config import CUDAGraphMode, ProfilerConfig from vllm.config.profiler import _is_uri_path from vllm.profiler.wrapper import WorkerProfiler +from vllm.v1.core.sched.output import CachedRequestData +from vllm.v1.worker.gpu_model_runner import GPUModelRunner +from vllm.v1.worker.gpu_worker import Worker class ConcreteWorkerProfiler(WorkerProfiler): @@ -236,3 +241,60 @@ class TestIsUriPath: def test_is_uri_path(self, path, expected): """Test that _is_uri_path correctly identifies URI vs local paths.""" assert _is_uri_path(path) == expected + + +class TestAnnotateProfile: + """Tests for Worker.annotate_profile() annotation string formatting.""" + + def _annotate(self, detailed: bool) -> str: + worker = MagicMock() + worker.vllm_config.profiler_config.detailed_trace_annotation = detailed + worker.profiler = MagicMock() + + ctx_req = MagicMock(req_id="ctx1", num_computed_tokens=0) + cached = CachedRequestData( + req_ids=["gen1"], + resumed_req_ids=set(), + new_token_ids=[], + all_token_ids={}, + new_block_ids=[], + num_computed_tokens=[10], + num_output_tokens=[1], + ) + sched = MagicMock( + scheduled_new_reqs=[ctx_req], + scheduled_cached_reqs=cached, + num_scheduled_tokens={"ctx1": 4, "gen1": 1}, + ) + + Worker.annotate_profile(worker, sched) + return worker.profiler.annotate_context_manager.call_args[0][0] + + def test_simple_format_mixed(self): + assert self._annotate(detailed=False) == ( + "execute_context_1(4)_generation_1(1)" + ) + + def test_detailed_format_mixed(self): + # ctx1: sq=4, sk=4, sqsq=16, sqsk=16 | gen1: sq=1, sk=11, sqsq=1, sqsk=11 | bs=5 + assert self._annotate(detailed=True) == ( + "execute_5_context_1(sq4sk4sqsq16sqsk16)_generation_1(sq1sk11sqsq1sqsk11)" + ) + + +def test_profiler_entered_during_capture(): + """Profiler is used as a context manager in _warmup_and_capture, + confirming it is active during the actual graph capture run.""" + runner = MagicMock() + runner.compilation_config.cudagraph_num_of_warmups = 0 + mock_profiler = MagicMock() + + GPUModelRunner._warmup_and_capture( + runner, + desc=MagicMock(num_tokens=4, uniform=True), + cudagraph_runtime_mode=CUDAGraphMode.FULL, + profiler=mock_profiler, + ) + + mock_profiler.__enter__.assert_called_once() + mock_profiler.__exit__.assert_called_once() diff --git a/vllm/config/profiler.py b/vllm/config/profiler.py index 68fa78854b4..f0e29d08f38 100644 --- a/vllm/config/profiler.py +++ b/vllm/config/profiler.py @@ -66,6 +66,17 @@ class ProfilerConfig: """If `True`, enables memory profiling in the torch profiler. Disabled by default.""" + capture_torch_profiler: bool = False + """If `True`, enables a torch profiler during CUDA graph capture on rank 0. + Traces are saved to a `capture_traces` subdirectory under `torch_profiler_dir`. + Requires `profiler` to be set to 'torch'.""" + + detailed_trace_annotation: bool = False + """If `True`, uses detailed annotations with roofline metrics (sk, sqsq, + sqsk) in profiler trace events. If `False`, uses simple annotations with + only context/generation request counts and token counts. + Disabled by default.""" + ignore_frontend: bool = False """If `True`, disables the front-end profiling of AsyncLLM when using the 'torch' profiler. This is needed to reduce overhead when using delay/limit options, @@ -144,4 +155,10 @@ class ProfilerConfig: if profiler_dir and not _is_uri_path(profiler_dir): self.torch_profiler_dir = os.path.abspath(os.path.expanduser(profiler_dir)) + if self.capture_torch_profiler and self.profiler != "torch": + raise ValueError( + "capture_torch_profiler is only applicable when profiler is " + "set to 'torch'" + ) + return self diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index 79a905adada..2c7adaaf2ae 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -8,7 +8,7 @@ import threading import time from collections import defaultdict from collections.abc import Callable, Iterable, Iterator, Sequence -from contextlib import contextmanager +from contextlib import AbstractContextManager, contextmanager, nullcontext from copy import copy, deepcopy from dataclasses import dataclass, replace from functools import reduce @@ -6756,6 +6756,44 @@ class GPUModelRunner( # Capture the large shapes first so that the smaller shapes # can reuse the memory pool allocated for the large shapes. set_cudagraph_capturing_enabled(True) + + # Setup torch profiler for graph capture traces (conditional) + from vllm.distributed.parallel_state import get_world_group + + local_rank = get_world_group().local_rank + enable_profiler = ( + local_rank == 0 + ) and self.vllm_config.profiler_config.capture_torch_profiler + if enable_profiler: + trace_dir = ( + self.vllm_config.profiler_config.torch_profiler_dir + "/capture_traces" + ) + profiler = torch.profiler.profile( + activities=[ + torch.profiler.ProfilerActivity.CPU, + torch.profiler.ProfilerActivity.CUDA, + ], + record_shapes=True, + profile_memory=True, + with_stack=True, + on_trace_ready=torch.profiler.tensorboard_trace_handler( + trace_dir, + worker_name=f"graph_capture_rank_{local_rank}", + use_gzip=True, + ), + ) + logger.info_once( + "Rank %d: Torch profiler enabled for CUDA graph capture, " + "traces will be saved to: %s", + local_rank, + trace_dir, + ) + else: + profiler = nullcontext() + logger.info_once( + "Rank %d: Torch profiler disabled for CUDA graph capture", local_rank + ) + with self._freeze_gc(), graph_capture(device=self.device): torch.accelerator.synchronize() torch.accelerator.empty_cache() @@ -6768,6 +6806,7 @@ class GPUModelRunner( self._capture_cudagraphs( batch_descriptors=batch_descs, cudagraph_runtime_mode=runtime_mode, + profiler=profiler, ) torch.accelerator.synchronize() @@ -6811,7 +6850,10 @@ class GPUModelRunner( profile_seq_lens: int | None = None, allow_microbatching: bool = False, num_warmups: int | None = None, + profiler: AbstractContextManager[Any] | None = None, ): + if profiler is None: + profiler = nullcontext() if num_warmups is None: num_warmups = self.compilation_config.cudagraph_num_of_warmups force_attention = cudagraph_runtime_mode == CUDAGraphMode.FULL @@ -6827,22 +6869,29 @@ class GPUModelRunner( num_active_loras=desc.num_active_loras, profile_seq_lens=profile_seq_lens, ) - self._dummy_run( - desc.num_tokens, - cudagraph_runtime_mode=cudagraph_runtime_mode, - uniform_decode=desc.uniform, - allow_microbatching=allow_microbatching, - skip_eplb=True, - remove_lora=False, - num_active_loras=desc.num_active_loras, - is_graph_capturing=True, - profile_seq_lens=profile_seq_lens, - ) + with ( + profiler, + torch.profiler.record_function( + f"capture_{desc.num_tokens}_{cudagraph_runtime_mode.name}" + ), + ): + self._dummy_run( + desc.num_tokens, + cudagraph_runtime_mode=cudagraph_runtime_mode, + uniform_decode=desc.uniform, + allow_microbatching=allow_microbatching, + skip_eplb=True, + remove_lora=False, + num_active_loras=desc.num_active_loras, + is_graph_capturing=True, + profile_seq_lens=profile_seq_lens, + ) def _capture_cudagraphs( self, batch_descriptors: list[BatchDescriptor], cudagraph_runtime_mode: CUDAGraphMode, + profiler: AbstractContextManager[Any] | None = None, ): assert ( cudagraph_runtime_mode != CUDAGraphMode.NONE @@ -6885,6 +6934,7 @@ class GPUModelRunner( batch_desc, cudagraph_runtime_mode=cudagraph_runtime_mode, allow_microbatching=allow_microbatching, + profiler=profiler, ) torch.accelerator.synchronize() self.maybe_remove_all_loras(self.lora_config) diff --git a/vllm/v1/worker/gpu_worker.py b/vllm/v1/worker/gpu_worker.py index 5fb0c387737..9c20df0d18c 100644 --- a/vllm/v1/worker/gpu_worker.py +++ b/vllm/v1/worker/gpu_worker.py @@ -975,19 +975,104 @@ class Worker(WorkerBase): iteration_details = compute_iteration_details(scheduler_output) - annotation = "".join( - [ - "execute_context_", - str(iteration_details.num_ctx_requests), - "(", - str(iteration_details.num_ctx_tokens), - ")_generation_", - str(iteration_details.num_generation_requests), - "(", - str(iteration_details.num_generation_tokens), - ")", - ] - ) + if self.vllm_config.profiler_config.detailed_trace_annotation: + # Compute roofline-model metrics per request, split by phase + # (context vs generation). These help estimate compute and + # memory intensity from the trace. + # + # Per-request quantities: + # query_len = number of scheduled (new) tokens for this request + # seq_len = total sequence length (computed + scheduled tokens) + # + # Aggregated across requests in each phase + # (ctx_=context, gen_=generation): + # seq_len_sum = sum of seq_len (total KV length) + # qq_compute = sum of query_len*query_len + # (proxy for QK^T compute cost) + # qk_compute = sum of query_len*seq_len + # (proxy for QK^T compute cost for decode and + # chunked prefill) + # total_scheduled_tokens = scheduled tokens across all requests + ctx_seq_len_sum = 0 + ctx_qq_compute = 0 + ctx_qk_compute = 0 + gen_seq_len_sum = 0 + gen_qq_compute = 0 + gen_qk_compute = 0 + total_scheduled_tokens = 0 + + # Build a map of req_id -> num_computed_tokens for all requests + new_req_ids = { + new_req.req_id for new_req in scheduler_output.scheduled_new_reqs + } + num_computed_tokens_ids = { + new_req.req_id: new_req.num_computed_tokens + for new_req in scheduler_output.scheduled_new_reqs + } + for req_id, num_computed_tokens in zip( + scheduler_output.scheduled_cached_reqs.req_ids, + scheduler_output.scheduled_cached_reqs.num_computed_tokens, + ): + num_computed_tokens_ids[req_id] = num_computed_tokens + + # Accumulate per-phase metrics + for req_id, num_tokens in scheduler_output.num_scheduled_tokens.items(): + query_len = num_tokens + total_scheduled_tokens += query_len + seq_len = num_computed_tokens_ids.get(req_id, 0) + query_len + if ( + scheduler_output.scheduled_cached_reqs.is_context_phase(req_id) + or req_id in new_req_ids + ): + ctx_seq_len_sum += seq_len + ctx_qq_compute += query_len * query_len + ctx_qk_compute += query_len * seq_len + else: + gen_seq_len_sum += seq_len + gen_qq_compute += query_len * query_len + gen_qk_compute += query_len * seq_len + annotation = "".join( + [ + "execute_", + str(total_scheduled_tokens), + "_context_", + str(iteration_details.num_ctx_requests), + "(sq", + str(iteration_details.num_ctx_tokens), + "sk", + str(ctx_seq_len_sum), + "sqsq", + str(ctx_qq_compute), + "sqsk", + str(ctx_qk_compute), + ")_generation_", + str(iteration_details.num_generation_requests), + "(sq", + str(iteration_details.num_generation_tokens), + "sk", + str(gen_seq_len_sum), + "sqsq", + str(gen_qq_compute), + "sqsk", + str(gen_qk_compute), + ")", + ] + ) + else: + annotation = "".join( + [ + "execute_context_", + str(iteration_details.num_ctx_requests), + "(", + str(iteration_details.num_ctx_tokens), + ")", + "_generation_", + str(iteration_details.num_generation_requests), + "(", + str(iteration_details.num_generation_tokens), + ")", + ] + ) return self.profiler.annotate_context_manager(annotation) @torch.inference_mode() From 088c0be2684fda35b18645bb36f0944a5c5fc941 Mon Sep 17 00:00:00 2001 From: "Kevin H. Luu" Date: Fri, 17 Jul 2026 13:44:48 -0700 Subject: [PATCH 10/51] [CI] Fix macOS wheel release annotation context (#48771) Signed-off-by: khluu Co-authored-by: Codex --- .buildkite/release-pipeline.yaml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.buildkite/release-pipeline.yaml b/.buildkite/release-pipeline.yaml index 136c96d218a..1be6f60fe25 100644 --- a/.buildkite/release-pipeline.yaml +++ b/.buildkite/release-pipeline.yaml @@ -137,7 +137,7 @@ steps: - 'mv artifacts/reassembled/wheel "artifacts/dist/$$wheel_name"' - "aws sts get-caller-identity" - "VLLM_WHEEL_PLATFORM=macos bash .buildkite/scripts/upload-nightly-wheels.sh" - - 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "s3://vllm-wheels/$$BUILDKITE_COMMIT/$(cd artifacts/dist && echo *.whl)"' + - 'bash .buildkite/scripts/annotate-build-artifact.sh "$$BUILDKITE_LABEL" "s3://vllm-wheels/$$BUILDKITE_COMMIT/$(cd artifacts/dist && echo *.whl)" release-wheels' plugins: - aws-assume-role-with-web-identity#v1.6.0: role-arn: arn:aws:iam::936637512419:role/vllm-release-macos-wheel-uploader From 41ea2dd44a3a20c46ebeb985de0022c7673fb953 Mon Sep 17 00:00:00 2001 From: aoshen02 Date: Sat, 18 Jul 2026 04:58:59 +0800 Subject: [PATCH 11/51] [Bugfix][V1/V2] Fix prompt_logprobs to respect logprobs_mode (#47680) Signed-off-by: Wojciech Wais Signed-off-by: Federico Kamelhar <209537060+fede-kamel@users.noreply.github.com> Signed-off-by: Allen Shen Co-authored-by: Wojciech Wais Co-authored-by: Federico Kamelhar <209537060+fede-kamel@users.noreply.github.com> --- tests/v1/sample/test_logprobs.py | 38 +++++++++++++++++++ vllm/config/model.py | 2 + vllm/config/vllm.py | 7 ---- vllm/model_executor/models/diffusion_gemma.py | 6 ++- vllm/v1/sample/rejection_sampler.py | 10 ++++- vllm/v1/worker/gpu/model_runner.py | 5 ++- vllm/v1/worker/gpu/sample/logprob.py | 17 ++++++--- vllm/v1/worker/gpu/sample/prompt_logprob.py | 28 ++++++++------ vllm/v1/worker/gpu/sample/sampler.py | 14 ++++--- .../gpu/spec_decode/rejection_sampler.py | 8 ++-- vllm/v1/worker/gpu_model_runner.py | 11 ++++-- 11 files changed, 106 insertions(+), 40 deletions(-) diff --git a/tests/v1/sample/test_logprobs.py b/tests/v1/sample/test_logprobs.py index fba240fea6a..aa17d2a1004 100644 --- a/tests/v1/sample/test_logprobs.py +++ b/tests/v1/sample/test_logprobs.py @@ -564,6 +564,44 @@ def test_logprobs_mode(logprobs_mode: LogprobsMode): cleanup_dist_env_and_memory() +def test_prompt_logprobs_mode(): + """prompt_logprobs must respect logprobs_mode: *_logits and *_logprobs + must return different values. Prompt tokens skip sampling processors, + so processed_* == raw_* on the prompt side.""" + from vllm import LLM + + values: dict[str, float] = {} + for mode in get_args(LogprobsMode): + llm = LLM( + "facebook/opt-125m", + enable_prefix_caching=False, + gpu_memory_utilization=0.05, + max_model_len=16, + logprobs_mode=mode, + ) + try: + results = llm.generate( + ["Hello world"], + sampling_params=SamplingParams( + max_tokens=1, prompt_logprobs=0, temperature=0 + ), + ) + assert results[0].prompt_logprobs is not None + assert results[0].prompt_logprobs[1] is not None + tok_id = results[0].prompt_token_ids[1] + values[mode] = results[0].prompt_logprobs[1][tok_id].logprob + finally: + del llm + torch.accelerator.empty_cache() + cleanup_dist_env_and_memory() + + assert values["raw_logprobs"] <= 0 + assert values["processed_logprobs"] <= 0 + assert values["raw_logits"] != values["raw_logprobs"] + assert values["processed_logits"] == values["raw_logits"] + assert values["processed_logprobs"] == values["raw_logprobs"] + + class TestCorrectDecodedToken: """Unit tests for _correct_decoded_token method in LogprobsProcessor. diff --git a/vllm/config/model.py b/vllm/config/model.py index e36b672cd82..6b032ae7621 100644 --- a/vllm/config/model.py +++ b/vllm/config/model.py @@ -233,6 +233,8 @@ class ModelConfig: Raw means the values before applying any logit processors, like bad words. Processed means the values after applying all processors, including temperature and top_k/top_p. + Note: for prompt_logprobs, processed_* and raw_* yield identical results + because prompt tokens do not go through sampling processors. """ use_fp64_gumbel: bool = False """Whether to use FP64 (instead of FP32) random noise for Gumbel-max and diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index 2e57cb1ec8e..cca9adbcf58 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -2202,13 +2202,6 @@ class VllmConfig: if model_config is not None and model_config.enable_prompt_embeds: unsupported.append("prompt embeds") - if ( - model_config is not None - and model_config.runner_type == "generate" - and model_config.logprobs_mode in ("raw_logits", "processed_logits") - ): - unsupported.append(f"logprobs mode '{model_config.logprobs_mode}'") - if self.cache_config.kv_sharing_fast_prefill: # Will be added by https://github.com/vllm-project/vllm/pull/35045 unsupported.append("KV sharing fast prefill") diff --git a/vllm/model_executor/models/diffusion_gemma.py b/vllm/model_executor/models/diffusion_gemma.py index 11a10131df1..70566871e09 100644 --- a/vllm/model_executor/models/diffusion_gemma.py +++ b/vllm/model_executor/models/diffusion_gemma.py @@ -53,7 +53,7 @@ from vllm.v1.worker.gpu.attn_utils import build_attn_metadata from vllm.v1.worker.gpu.buffer_utils import UvaBackedTensor, async_copy_to_gpu from vllm.v1.worker.gpu.input_batch import InputBatch from vllm.v1.worker.gpu.model_states.interface import ModelState -from vllm.v1.worker.gpu.sample.logprob import compute_topk_logprobs +from vllm.v1.worker.gpu.sample.logprob import compute_topk_scores from vllm.v1.worker.gpu.sample.output import SamplerOutput from vllm.v1.worker.gpu.sample.penalties import use_penalty from vllm.v1.worker.gpu.states import RequestState @@ -1064,6 +1064,7 @@ class DiffusionSampler: ): self.sampling_states = sampler.sampling_states self.req_states = sampler.req_states + self.logits_mode = sampler.logprobs_mode in ("raw_logits", "processed_logits") # Self-conditioning soft embed = probs @ embed_weight * normalizer, # computed in the sampler (see _compiled_sample_step). ``embed_weight`` # is the vocab-parallel shard; [sc_vocab_start, sc_vocab_end) is this @@ -1359,10 +1360,11 @@ class DiffusionSampler: # positions are never emitted. k_i = int(valid_canvas_len_np[start_req + li]) pos = li * CL - self._pending_logprobs[slot.item()] = compute_topk_logprobs( + self._pending_logprobs[slot.item()] = compute_topk_scores( flat_logits[pos : pos + k_i], max_num_logprobs, argmax_tokens[local_idx][:k_i], + logits_mode=self.logits_mode, ) # Commit steps: is_committing was True at entry. Reassemble previously diff --git a/vllm/v1/sample/rejection_sampler.py b/vllm/v1/sample/rejection_sampler.py index 69309213836..0f56b8a4a2c 100644 --- a/vllm/v1/sample/rejection_sampler.py +++ b/vllm/v1/sample/rejection_sampler.py @@ -67,8 +67,14 @@ class RejectionSampler(nn.Module): self.sampler = sampler self.use_fp64_gumbel = getattr(sampler, "use_fp64_gumbel", False) logprobs_mode = self.sampler.logprobs_mode - self.is_processed_logprobs_mode = logprobs_mode.startswith("processed") - self.is_logits_logprobs_mode = logprobs_mode.endswith("logits") + self.is_processed_logprobs_mode = logprobs_mode in ( + "processed_logprobs", + "processed_logits", + ) + self.is_logits_logprobs_mode = logprobs_mode in ( + "raw_logits", + "processed_logits", + ) self.synthetic_conditional_rates: torch.Tensor | None = None if ( diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 2bb52e2fd89..518d12a6b28 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -341,7 +341,10 @@ class GPUModelRunner(LoRAModelRunnerMixin): self.speculative_config, self.device, ) - self.prompt_logprobs_worker = PromptLogprobsWorker(self.max_num_reqs) + self.prompt_logprobs_worker = PromptLogprobsWorker( + self.max_num_reqs, + logprobs_mode=self.model_config.logprobs_mode, + ) self.structured_outputs_worker = StructuredOutputsWorker( max_num_logits=self.max_num_reqs * self.decode_query_len, vocab_size=self.vocab_size, diff --git a/vllm/v1/worker/gpu/sample/logprob.py b/vllm/v1/worker/gpu/sample/logprob.py index cb2cf1a590e..c0f31cdf9f0 100644 --- a/vllm/v1/worker/gpu/sample/logprob.py +++ b/vllm/v1/worker/gpu/sample/logprob.py @@ -106,7 +106,7 @@ def compute_token_logprobs( return logprobs -def compute_topk_logprobs( +def compute_topk_scores( logits: torch.Tensor, num_logprobs: int, sampled_token_ids: torch.Tensor, @@ -114,6 +114,7 @@ def compute_topk_logprobs( logprob_token_ids_state: "LogprobTokenIdsState | None" = None, expanded_idx_mapping: torch.Tensor | None = None, max_per_req_token_ids: int = 0, + logits_mode: bool = False, ) -> LogprobsTensors: assert num_logprobs >= 0 batch_size, vocab_size = logits.shape @@ -124,7 +125,10 @@ def compute_topk_logprobs( if num_logprobs > 0: topk_indices = torch.topk(logits, num_logprobs, dim=-1).indices logprob_token_ids = torch.cat((logprob_token_ids, topk_indices), dim=1) - logprobs = compute_token_logprobs(logits, logprob_token_ids) + if logits_mode: + scores = logits.gather(-1, logprob_token_ids).to(torch.float32) + else: + scores = compute_token_logprobs(logits, logprob_token_ids) else: # Some requests specified logprob_token_ids. Build the [batch_size, # 1 + max_cols] token_ids matrix and validity mask on the GPU via a @@ -158,8 +162,11 @@ def compute_topk_logprobs( NUM_TOPK=num_logprobs, PADDED_COLS=triton.next_power_of_2(num_cols), ) - logprobs = compute_token_logprobs(logits, logprob_token_ids) - logprobs = logprobs.masked_fill(~valid_mask, float("-inf")) + if logits_mode: + scores = logits.gather(-1, logprob_token_ids).to(torch.float32) + else: + scores = compute_token_logprobs(logits, logprob_token_ids) + scores = scores.masked_fill(~valid_mask, float("-inf")) token_ranks = torch.empty(batch_size, dtype=torch.int64, device=logits.device) _ranks_kernel[(batch_size,)]( @@ -172,7 +179,7 @@ def compute_topk_logprobs( ) return LogprobsTensors( logprob_token_ids=logprob_token_ids, - logprobs=logprobs, + logprobs=scores, selected_token_ranks=token_ranks, cu_num_generated_tokens=cu_num_logits, ) diff --git a/vllm/v1/worker/gpu/sample/prompt_logprob.py b/vllm/v1/worker/gpu/sample/prompt_logprob.py index b89ebac35d9..4d4cc244825 100644 --- a/vllm/v1/worker/gpu/sample/prompt_logprob.py +++ b/vllm/v1/worker/gpu/sample/prompt_logprob.py @@ -5,16 +5,18 @@ from collections.abc import Callable import numpy as np import torch +from vllm.config.model import LogprobsMode from vllm.sampling_params import SamplingParams from vllm.triton_utils import tl, triton from vllm.v1.outputs import LogprobsTensors from vllm.v1.worker.gpu.input_batch import InputBatch -from vllm.v1.worker.gpu.sample.logprob import compute_topk_logprobs +from vllm.v1.worker.gpu.sample.logprob import compute_topk_scores class PromptLogprobsWorker: - def __init__(self, max_num_reqs: int): + def __init__(self, max_num_reqs: int, logprobs_mode: LogprobsMode = "raw_logprobs"): self.max_num_reqs = max_num_reqs + self.logprobs_mode = logprobs_mode self.uses_prompt_logprobs = np.zeros(self.max_num_reqs, dtype=bool) self.num_prompt_logprobs = np.zeros(self.max_num_reqs, dtype=np.int32) @@ -82,6 +84,7 @@ class PromptLogprobsWorker: hidden_states[: input_batch.num_tokens], logits_fn, max_num_prompt_logprobs, + self.logprobs_mode, ) ) @@ -206,33 +209,36 @@ def compute_prompt_logprobs_with_chunking( prompt_hidden_states: torch.Tensor, logits_fn: Callable[[torch.Tensor], torch.Tensor], num_prompt_logprobs: int, + logprobs_mode: LogprobsMode = "raw_logprobs", ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: # Since materializing the full prompt logits can take too much memory, # we compute it in chunks. CHUNK_SIZE = 1024 token_ids = [] - logprobs = [] + scores = [] ranks = [] + logits_mode = logprobs_mode in ("raw_logits", "processed_logits") prompt_token_ids = prompt_token_ids.to(torch.int64) for start_idx in range(0, prompt_token_ids.shape[0], CHUNK_SIZE): end_idx = start_idx + CHUNK_SIZE # NOTE(woosuk): logits_fn can be slow because it involves all-gather. prompt_logits = logits_fn(prompt_hidden_states[start_idx:end_idx]) - requested_num_prompt_logprobs = ( + requested_num = ( prompt_logits.shape[-1] if num_prompt_logprobs == -1 else num_prompt_logprobs ) - prompt_logprobs = compute_topk_logprobs( + result = compute_topk_scores( prompt_logits, - requested_num_prompt_logprobs, + requested_num, prompt_token_ids[start_idx:end_idx], + logits_mode=logits_mode, ) - token_ids.append(prompt_logprobs.logprob_token_ids) - logprobs.append(prompt_logprobs.logprobs) - ranks.append(prompt_logprobs.selected_token_ranks) + token_ids.append(result.logprob_token_ids) + scores.append(result.logprobs) + ranks.append(result.selected_token_ranks) token_ids = torch.cat(token_ids, dim=0) if len(token_ids) > 1 else token_ids[0] - logprobs = torch.cat(logprobs, dim=0) if len(logprobs) > 1 else logprobs[0] + scores = torch.cat(scores, dim=0) if len(scores) > 1 else scores[0] ranks = torch.cat(ranks, dim=0) if len(ranks) > 1 else ranks[0] - return token_ids, logprobs, ranks + return token_ids, scores, ranks diff --git a/vllm/v1/worker/gpu/sample/sampler.py b/vllm/v1/worker/gpu/sample/sampler.py index b269de9eaed..f0a83c92efb 100644 --- a/vllm/v1/worker/gpu/sample/sampler.py +++ b/vllm/v1/worker/gpu/sample/sampler.py @@ -19,7 +19,7 @@ from vllm.v1.worker.gpu.sample.gumbel import gumbel_sample from vllm.v1.worker.gpu.sample.logit_bias import LogitBiasState from vllm.v1.worker.gpu.sample.logprob import ( LogprobTokenIdsState, - compute_topk_logprobs, + compute_topk_scores, ) from vllm.v1.worker.gpu.sample.output import SamplerOutput from vllm.v1.worker.gpu.sample.penalties import PenaltiesState @@ -38,8 +38,6 @@ class Sampler: num_speculative_tokens: int = 1, use_fp64_gumbel: bool = False, ): - if logprobs_mode not in ("processed_logprobs", "raw_logprobs"): - raise NotImplementedError(f"Unsupported logprobs_mode: {logprobs_mode}") self.logprobs_mode = logprobs_mode self.compute_nans = envs.VLLM_COMPUTE_NANS_IN_LOGITS # False by default. self.use_fp64_gumbel = use_fp64_gumbel @@ -102,12 +100,12 @@ class Sampler: ) if return_logprobs: - if self.logprobs_mode == "processed_logprobs": + if self.logprobs_mode in ("processed_logprobs", "processed_logits"): logits = processed_logits expanded_logits = logits.shape[0] != idx_mapping_np.shape[0] cu_num_logits = cu_num_logits_np.tolist() if expanded_logits else None num_logprobs = max_num_logprobs if max_num_logprobs != NO_LOGPROBS else 0 - logprobs_tensors = compute_topk_logprobs( + logprobs_tensors = compute_topk_scores( logits, num_logprobs, sampled, @@ -115,6 +113,7 @@ class Sampler: logprob_token_ids_state=self.logprob_token_ids_state, expanded_idx_mapping=input_batch.expanded_idx_mapping, max_per_req_token_ids=max_per_req_token_ids, + logits_mode=self.logprobs_mode in ("raw_logits", "processed_logits"), ) else: logprobs_tensors = None @@ -222,7 +221,10 @@ class Sampler: # any greedy requests or per-request seeds, or if post-processed # logprobs need to be returned for any requests. (top_k is None and top_p is None) - or (return_logprobs and self.logprobs_mode == "processed_logprobs") + or ( + return_logprobs + and self.logprobs_mode in ("processed_logprobs", "processed_logits") + ) or self.sampling_states.any_greedy(idx_mapping_np) or self.sampling_states.any_explicit_seed(idx_mapping_np) ) diff --git a/vllm/v1/worker/gpu/spec_decode/rejection_sampler.py b/vllm/v1/worker/gpu/spec_decode/rejection_sampler.py index c56252d55d7..4753d281746 100644 --- a/vllm/v1/worker/gpu/spec_decode/rejection_sampler.py +++ b/vllm/v1/worker/gpu/spec_decode/rejection_sampler.py @@ -11,7 +11,7 @@ from vllm.v1.worker.gpu.input_batch import ( get_num_sampled_and_rejected, ) from vllm.v1.worker.gpu.metrics.logits import get_num_nans -from vllm.v1.worker.gpu.sample.logprob import compute_topk_logprobs +from vllm.v1.worker.gpu.sample.logprob import compute_topk_scores from vllm.v1.worker.gpu.sample.output import SamplerOutput from vllm.v1.worker.gpu.sample.sampler import Sampler from vllm.v1.worker.gpu.sample.states import NO_LOGPROBS @@ -91,11 +91,13 @@ class RejectionSampler: num_warps=1, ) expanded_logits = num_logits != input_batch.idx_mapping.shape[0] - return compute_topk_logprobs( + return compute_topk_scores( logits, max_num_logprobs, flat_sampled, input_batch.cu_num_logits_np.tolist() if expanded_logits else None, + logits_mode=self.sampler.logprobs_mode + in ("raw_logits", "processed_logits"), ) def __call__( @@ -139,7 +141,7 @@ class RejectionSampler: sampled, num_sampled, processed_logits - if self.sampler.logprobs_mode == "processed_logprobs" + if self.sampler.logprobs_mode in ("processed_logprobs", "processed_logits") else logits, ) diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index 2c7adaaf2ae..f3f0afb1301 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -5635,10 +5635,15 @@ class GPUModelRunner( # to gather the logprob for. tgt_token_ids = prompt_token_ids[start_tok : start_tok + num_logits] - # Compute prompt logprobs. - logprobs = self.sampler.compute_logprobs(logits) + # Compute prompt scores respecting logprobs_mode. + # NOTE: prompt tokens skip sampling processors, so + # processed_* and raw_* yield the same scores here. + if self.model_config.logprobs_mode in ("raw_logits", "processed_logits"): + scores = logits.to(torch.float32) + else: + scores = self.sampler.compute_logprobs(logits) token_ids, logprobs, ranks, _ = self.sampler.gather_logprobs( - logprobs, num_prompt_logprobs, tgt_token_ids + scores, num_prompt_logprobs, tgt_token_ids ) # Transfer GPU->CPU async. From c9be3a8aa1f01f7efdc99dd9439451617ef54a00 Mon Sep 17 00:00:00 2001 From: Shangdi Yu Date: Fri, 17 Jul 2026 14:39:09 -0700 Subject: [PATCH 12/51] [Kernel][Helion] Disable warp specialization in rms_norm_per_block_quant B200 configs (#48797) Signed-off-by: Shangdi Yu Co-authored-by: Claude Opus 4.8 (1M context) --- .../rms_norm_per_block_quant/nvidia_b200.json | 94 +++++++++---------- 1 file changed, 47 insertions(+), 47 deletions(-) diff --git a/vllm/kernels/helion/configs/rms_norm_per_block_quant/nvidia_b200.json b/vllm/kernels/helion/configs/rms_norm_per_block_quant/nvidia_b200.json index 6acf8f29dab..afb30aab42d 100644 --- a/vllm/kernels/helion/configs/rms_norm_per_block_quant/nvidia_b200.json +++ b/vllm/kernels/helion/configs/rms_norm_per_block_quant/nvidia_b200.json @@ -655,7 +655,7 @@ "range_warp_specializes": [ null, null, - true, + false, null ], "range_num_stages": [], @@ -1705,7 +1705,7 @@ "range_warp_specializes": [ null, false, - true, + false, null ], "range_num_stages": [], @@ -1838,14 +1838,14 @@ ], "range_unroll_factors": [ 0, - 4, + 2, 0, 1 ], "range_warp_specializes": [ null, null, - true, + false, null ], "range_num_stages": [], @@ -1872,7 +1872,7 @@ "last", "" ], - "num_warps": 2, + "num_warps": 4, "num_stages": 6, "indexing": [ "tensor_descriptor", @@ -1881,7 +1881,7 @@ "tensor_descriptor", "tensor_descriptor", "tensor_descriptor", - "pointer", + "tensor_descriptor", "tensor_descriptor", "pointer" ], @@ -2610,49 +2610,49 @@ ], "range_unroll_factors": [ 0, - 3, + 0, 0, 0 ], "range_warp_specializes": [ null, + false, null, - true, null ], "range_num_stages": [], "range_multi_buffers": [ null, - true, + null, null, null ], "range_flattens": [ null, + false, null, - true, null ], "static_ranges": [ true ], "load_eviction_policies": [ - "last", - "first", - "first", "first", "", + "", + "first", + "first", "last" ], - "num_warps": 1, - "num_stages": 3, + "num_warps": 4, + "num_stages": 7, "indexing": [ - "tensor_descriptor", - "tensor_descriptor", - "tensor_descriptor", - "tensor_descriptor", "pointer", "tensor_descriptor", + "pointer", + "pointer", + "tensor_descriptor", + "pointer", "tensor_descriptor", "tensor_descriptor", "tensor_descriptor" @@ -2680,14 +2680,14 @@ ], "range_unroll_factors": [ 0, - 4, + 0, 0, 0 ], "range_warp_specializes": [ - null, null, true, + null, null ], "range_num_stages": [], @@ -2707,14 +2707,14 @@ true ], "load_eviction_policies": [ - "", + "last", "", "first", "last", "last", "" ], - "num_warps": 1, + "num_warps": 2, "num_stages": 3, "indexing": [ "pointer", @@ -2723,7 +2723,7 @@ "tensor_descriptor", "tensor_descriptor", "pointer", - "tensor_descriptor", + "pointer", "pointer", "pointer" ], @@ -2809,7 +2809,7 @@ }, "config": { "block_sizes": [ - 8192, + 2048, 8 ], "loop_orders": [ @@ -2821,13 +2821,13 @@ "range_unroll_factors": [ 0, 3, - 0, - 2 + 1, + 0 ], "range_warp_specializes": [ null, null, - true, + false, null ], "range_num_stages": [], @@ -2841,10 +2841,10 @@ null, null, false, - false + null ], "static_ranges": [ - false + true ], "load_eviction_policies": [ "", @@ -2854,7 +2854,7 @@ "", "last" ], - "num_warps": 1, + "num_warps": 2, "num_stages": 3, "indexing": [ "pointer", @@ -2884,20 +2884,20 @@ ], "loop_orders": [ [ - 0, - 1 + 1, + 0 ] ], "range_unroll_factors": [ 0, - 3, + 1, 0, - 2 + 0 ], "range_warp_specializes": [ null, + false, null, - true, null ], "range_num_stages": [], @@ -2908,37 +2908,37 @@ null ], "range_flattens": [ - null, null, false, - false + false, + null ], "static_ranges": [ - false + true ], "load_eviction_policies": [ + "last", "", "", "first", - "last", - "", + "first", "last" ], - "num_warps": 1, - "num_stages": 3, + "num_warps": 2, + "num_stages": 7, "indexing": [ "pointer", "pointer", - "tensor_descriptor", + "pointer", + "pointer", "tensor_descriptor", "pointer", "tensor_descriptor", - "tensor_descriptor", "pointer", - "pointer" + "tensor_descriptor" ], "atomic_indexing": [], "pid_type": "flat" } } -] \ No newline at end of file +] From fae543015cd49dfb4db3afb3a2e96d2eccb3e1db Mon Sep 17 00:00:00 2001 From: Wang Xingda Date: Sat, 18 Jul 2026 06:09:14 +0800 Subject: [PATCH 13/51] [Frontend]Flatten beam-search beams with itertools.chain instead of sum (#48829) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Wang Xingda Co-authored-by: 王兴达 Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> --- vllm/entrypoints/generate/beam_search/offline.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/vllm/entrypoints/generate/beam_search/offline.py b/vllm/entrypoints/generate/beam_search/offline.py index b38830d6e41..e1bceb5d23d 100644 --- a/vllm/entrypoints/generate/beam_search/offline.py +++ b/vllm/entrypoints/generate/beam_search/offline.py @@ -207,7 +207,9 @@ class BeamSearchOfflineMixin(OfflineInferenceMixin): Returns True if all beams are exhausted and search should stop. """ all_beams: list[BeamSearchSequence] = list( - sum((instance.beams for instance in instances_batch), []) + itertools.chain.from_iterable( + instance.beams for instance in instances_batch + ) ) pos = [0] + list( itertools.accumulate(len(instance.beams) for instance in instances_batch) From 02c01f442b8bfd0affa1f2ceedda01f2e1384303 Mon Sep 17 00:00:00 2001 From: Michael Goin Date: Fri, 17 Jul 2026 21:13:14 -0400 Subject: [PATCH 14/51] [Model] Use standard ModelOpt config for Inkling NVFP4 (#48990) Signed-off-by: mgoin --- tests/config/test_model_arch_config.py | 17 +++++ .../models/inkling/test_moe_weight_layout.py | 25 ++++++++ vllm/models/inkling/nvfp4.py | 64 ------------------- vllm/models/inkling/nvidia/model.py | 17 +---- vllm/models/inkling/nvidia/moe.py | 33 ++-------- vllm/models/inkling/nvidia/mtp.py | 1 - .../model_arch_config_convertor.py | 9 ++- 7 files changed, 56 insertions(+), 110 deletions(-) delete mode 100644 vllm/models/inkling/nvfp4.py diff --git a/tests/config/test_model_arch_config.py b/tests/config/test_model_arch_config.py index 46790be6e4e..212f3a9a254 100644 --- a/tests/config/test_model_arch_config.py +++ b/tests/config/test_model_arch_config.py @@ -131,6 +131,23 @@ def test_head_size_falls_back_when_head_dim_is_zero(): assert convertor.get_head_size() == 128 +def test_legacy_modelopt_config_without_producer_is_normalized(): + quantization_config = { + "quantization": { + "quant_algo": "NVFP4", + "group_size": 16, + "kv_cache_quant_algo": None, + "exclude_modules": [], + "modelopt_quant_config": {"quant_cfg": {}}, + } + } + hf_config = PretrainedConfig(quantization_config=quantization_config) + + convertor = ModelArchConfigConvertorBase(hf_config, hf_config) + + assert convertor.get_quantization_config()["quant_method"] == "modelopt_fp4" + + @pytest.mark.parametrize("model", BASE_MODELS_TO_TEST) def test_base_model_arch_config(model: str): """Test model architecture config for base models.""" diff --git a/tests/models/inkling/test_moe_weight_layout.py b/tests/models/inkling/test_moe_weight_layout.py index e6c35c57b4d..53a37e3e60b 100644 --- a/tests/models/inkling/test_moe_weight_layout.py +++ b/tests/models/inkling/test_moe_weight_layout.py @@ -7,6 +7,7 @@ import pytest import torch from vllm.lora.utils import get_supported_lora_modules +from vllm.model_executor.layers.quantization.modelopt import ModelOptNvFp4Config from vllm.models.inkling.nvidia import moe from vllm.models.inkling.nvidia.model import _TmlForCausalLMBase from vllm.platforms import current_platform @@ -87,6 +88,30 @@ def test_custom_embedding_is_not_a_lora_target() -> None: assert "lm_head" in supported +def test_inkling_mapper_maps_modelopt_exclusions() -> None: + quant_config = ModelOptNvFp4Config.from_config( + { + "quantization": { + "quant_algo": "NVFP4", + "group_size": 16, + "kv_cache_quant_algo": None, + "exclude_modules": [ + "model.llm.layers.2.mlp.experts", + "model.llm.layers.2.mlp.shared_experts", + ], + } + } + ) + + quant_config.apply_vllm_mapper( + _TmlForCausalLMBase.hf_to_vllm_mapper.get_unstacked_mapper() + ) + + assert quant_config.is_layer_excluded("model.layers.2.mlp.experts") + assert quant_config.is_layer_excluded("model.layers.2.mlp.shared_experts") + assert not quant_config.is_layer_excluded("model.layers.3.mlp.experts") + + @pytest.mark.parametrize(("projection", "amax"), [("w13", 4.375), ("w2", 2960.0)]) def test_moe_loads_calibrated_input_scale(projection: str, amax: float) -> None: experts = SimpleNamespace( diff --git a/vllm/models/inkling/nvfp4.py b/vllm/models/inkling/nvfp4.py deleted file mode 100644 index 129ab82aa14..00000000000 --- a/vllm/models/inkling/nvfp4.py +++ /dev/null @@ -1,64 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""NVFP4 (ModelOpt) support for the Inkling mixture-of-experts. - -Only the routed MoE experts are quantized in the Inkling checkpoint; -attention, the dense MLP, and the shared "sink" experts stay bf16 (they are -in the checkpoint ``exclude_modules``). The routed experts are served by -vLLM's standard ModelOpt NVFP4 fused-MoE stack (see ``moe.py``); this module -keeps the checkpoint detection. -""" - -from __future__ import annotations - -FLOAT8_E4M3_MAX = 448.0 -FLOAT4_E2M1_MAX = 6.0 - - -class InklingNvfp4Config: - """Lightweight NVFP4 descriptor parsed from the checkpoint quant config. - - Holds the (mapped) ``exclude_modules`` so the model can decide, per MoE - layer and per expert group, whether the weights are NVFP4 or plain bf16. - """ - - def __init__(self, group_size: int, exclude_modules: list[str]) -> None: - self.group_size = group_size - self.exclude_modules = set(exclude_modules) - - @staticmethod - def _is_nvfp4(quant_cfg: dict) -> bool: - wq = quant_cfg["modelopt_quant_config"]["quant_cfg"]["*weight_quantizer"] - return tuple(wq["num_bits"]) == (2, 1) and tuple( - wq["block_sizes"].get("scale_bits", []) - ) == (4, 3) - - @classmethod - def from_hf_config(cls, hf_config) -> InklingNvfp4Config | None: - quant_cfg = getattr(hf_config, "quantization_config", None) - text_config = getattr(hf_config, "text_config", None) - if quant_cfg is None and text_config is not None: - quant_cfg = getattr(text_config, "quantization_config", None) - if quant_cfg is None: - return None - # ModelOpt <=0.29 nests everything under "quantization". - if "quantization" in quant_cfg: - quant_cfg = quant_cfg["quantization"] - if not cls._is_nvfp4(quant_cfg): - return None - group_size = quant_cfg.get("group_size", 16) - if group_size != 16: - raise ValueError("Inkling NVFP4 only supports group size 16") - exclude = list(quant_cfg.get("exclude_modules", []) or []) - return cls(group_size=group_size, exclude_modules=exclude) - - def experts_quantized(self, layer_id: int) -> bool: - """Whether the routed experts of ``layer_id`` are NVFP4 (vs excluded).""" - return f"model.llm.layers.{layer_id}.mlp.experts" not in self.exclude_modules - - def shared_experts_quantized(self, layer_id: int) -> bool: - """Whether the shared sink experts of ``layer_id`` are NVFP4.""" - return ( - f"model.llm.layers.{layer_id}.mlp.shared_experts" - not in self.exclude_modules - ) diff --git a/vllm/models/inkling/nvidia/model.py b/vllm/models/inkling/nvidia/model.py index bf7f700dee6..a8506d92327 100644 --- a/vllm/models/inkling/nvidia/model.py +++ b/vllm/models/inkling/nvidia/model.py @@ -47,7 +47,6 @@ from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.sequence import IntermediateTensors from ..configs import InklingMMConfig, InklingModelConfig -from ..nvfp4 import InklingNvfp4Config from .attention import InklingAttention, compute_log_scaling_tau from .layernorm import InklingRMSNorm from .logits_processor import InklingLogitsProcessor @@ -123,7 +122,6 @@ class InklingDecoderLayer(nn.Module): is_local: bool, quant_config: QuantizationConfig | None, prefix: str, - nvfp4_config: InklingNvfp4Config | None = None, force_dense_mlp: bool = False, ) -> None: super().__init__() @@ -172,14 +170,10 @@ class InklingDecoderLayer(nn.Module): prefix=f"{prefix}.mlp", ) else: - # InklingMoE decides per layer (from the checkpoint exclude list) - # whether the routed experts are NVFP4 or bf16; the shared sink - # experts are always bf16. self.mlp = InklingMoE( config, - layer_id, prefix=f"{prefix}.mlp", - nvfp4_config=nvfp4_config, + quant_config=quant_config, ) # Short convolution on the attention-output and MLP-output residual @@ -261,7 +255,6 @@ class InklingModel(nn.Module): config: InklingModelConfig, quant_config: QuantizationConfig | None, prefix: str, - nvfp4_config: InklingNvfp4Config | None = None, ) -> None: super().__init__() self.config = config @@ -278,7 +271,7 @@ class InklingModel(nn.Module): def get_layer(prefix: str) -> InklingDecoderLayer: idx = _layer_id(prefix + ".") or int(prefix.split(".")[-1]) return InklingDecoderLayer( - config, idx, idx in local_ids, quant_config, prefix, nvfp4_config + config, idx, idx in local_ids, quant_config, prefix ) self.start_layer, self.end_layer, self.layers = make_layers( @@ -408,11 +401,6 @@ class _TmlForCausalLMBase(nn.Module, SupportsPP, SupportsLoRA): ) -> None: quant_config = vllm_config.quant_config self.config = text_config - # NVFP4 experts are detected directly from the checkpoint quant config; - # only the MoE experts are quantized (attention/dense MLP stay bf16). - self.nvfp4_config = InklingNvfp4Config.from_hf_config( - vllm_config.model_config.hf_config - ) # Read by the MRV2 runner to publish per-request short-conv metadata. # Short convolution is intrinsic to Inkling, so this is always set. self.uses_sconv = True @@ -420,7 +408,6 @@ class _TmlForCausalLMBase(nn.Module, SupportsPP, SupportsLoRA): config=text_config, quant_config=quant_config, prefix=maybe_prefix(prefix, "model"), - nvfp4_config=self.nvfp4_config, ) initialize_lamport_rs_conv( text_config.hidden_size, diff --git a/vllm/models/inkling/nvidia/moe.py b/vllm/models/inkling/nvidia/moe.py index 255c9328c8a..6d0a550b03d 100644 --- a/vllm/models/inkling/nvidia/moe.py +++ b/vllm/models/inkling/nvidia/moe.py @@ -43,20 +43,19 @@ from vllm.utils.multi_stream_utils import maybe_execute_in_parallel from vllm.utils.torch_utils import aux_stream from ..configs import InklingModelConfig -from ..nvfp4 import FLOAT4_E2M1_MAX, FLOAT8_E4M3_MAX if TYPE_CHECKING: from vllm.model_executor.layers.fused_moe.routed_experts import ( RoutedExperts, ) - - from ..nvfp4 import InklingNvfp4Config + from vllm.model_executor.layers.quantization import QuantizationConfig # --------------------------------------------------------------------------- # Gate / expert selection # --------------------------------------------------------------------------- _INKLING_LL_BF16_MAX_TOKENS = 64 +_NVFP4_INPUT_SCALE_DENOMINATOR = torch.finfo(torch.float8_e4m3fn).max * 6.0 def _linear_with_fp32_out(x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: @@ -412,10 +411,9 @@ class InklingMoE(nn.Module): def __init__( self, config: InklingModelConfig, - layer_id: int, *, prefix: str = "", - nvfp4_config: InklingNvfp4Config | None = None, + quant_config: QuantizationConfig | None = None, ) -> None: super().__init__() # Overfit to the served checkpoint: sigmoid gate renormalized after @@ -436,22 +434,6 @@ class InklingMoE(nn.Module): use_gate_bias=config.use_gate_bias, ) - moe_quant_config = None - if nvfp4_config is not None and nvfp4_config.experts_quantized(layer_id): - from vllm.model_executor.layers.quantization.modelopt import ( - ModelOptNvFp4Config, - ) - - # The Inkling checkpoint is ModelOpt NVFP4; exclusion is decided per - # layer right here, so no exclude list is needed. - moe_quant_config = ModelOptNvFp4Config( - quant_method="NVFP4", - is_checkpoint_nvfp4_serialized=True, - kv_cache_quant_algo=None, - exclude_modules=[], - group_size=nvfp4_config.group_size, - ) - # TRTLLM MoE kernels assume equal, contiguous per-rank expert slabs # (local_expert_offset = ep_rank * local_num_experts), so pad the # expert count to a multiple of the EP size. A no-op for the usual @@ -464,7 +446,7 @@ class InklingMoE(nn.Module): hidden_size=config.hidden_size, intermediate_size=config.intermediate_size, renormalize=False, - quant_config=moe_quant_config, + quant_config=quant_config, prefix=f"{prefix}.experts", custom_routing_function=self._select_routed, router_logits_dtype=torch.float32, @@ -476,11 +458,6 @@ class InklingMoE(nn.Module): self.experts.moe_config.skip_final_all_reduce = True self._routed_sel: tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None = None - # The sinks are always bf16; fail loudly on a checkpoint that - # quantizes them instead of silently misloading. - assert nvfp4_config is None or not nvfp4_config.shared_experts_quantized( - layer_id - ), f"layer {layer_id}: NVFP4 shared experts are not supported" sink_experts_cls = ( InklingSinkExpertsLinear @@ -589,7 +566,7 @@ class InklingMoE(nn.Module): f"bad {projection} input_amax: {amax}" ) input_scale = getattr(experts, f"{projection}_input_scale") - input_scale.data.fill_(amax / (FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX)) + input_scale.data.fill_(amax / _NVFP4_INPUT_SCALE_DENOMINATOR) return [f"experts.routed_experts.{projection}_input_scale"] param = getattr(experts, key) diff --git a/vllm/models/inkling/nvidia/mtp.py b/vllm/models/inkling/nvidia/mtp.py index 868ede3bed6..a2559d4cf05 100644 --- a/vllm/models/inkling/nvidia/mtp.py +++ b/vllm/models/inkling/nvidia/mtp.py @@ -73,7 +73,6 @@ class InklingMTPDepthLayer(nn.Module): is_local=is_local, quant_config=None, prefix=f"{prefix}.transformer_block", - nvfp4_config=None, force_dense_mlp=True, ) diff --git a/vllm/transformers_utils/model_arch_config_convertor.py b/vllm/transformers_utils/model_arch_config_convertor.py index b6f73f39f3d..bd146dff7dc 100644 --- a/vllm/transformers_utils/model_arch_config_convertor.py +++ b/vllm/transformers_utils/model_arch_config_convertor.py @@ -215,8 +215,13 @@ class ModelArchConfigConvertorBase: else: # Set quant_method for ModelOpt models. producer_name = quant_cfg.get("producer", {}).get("name") - if producer_name == "modelopt": - quant_algo = quant_cfg.get("quantization", {}).get("quant_algo") + modelopt_quant_cfg = quant_cfg.get("quantization", {}) + is_legacy_modelopt = ( + isinstance(modelopt_quant_cfg, dict) + and "modelopt_quant_config" in modelopt_quant_cfg + ) + if producer_name == "modelopt" or is_legacy_modelopt: + quant_algo = modelopt_quant_cfg.get("quant_algo") if quant_algo is not None: quant_algo_upper = str(quant_algo).upper() if quant_algo_upper in { From 425c4eafb064f3804538a93584ec07c196d14cf9 Mon Sep 17 00:00:00 2001 From: Michael Goin Date: Fri, 17 Jul 2026 21:15:20 -0400 Subject: [PATCH 15/51] [Sampler] Stop upcasting logits to fp32 in apply_sampling_params (#48641) Signed-off-by: mgoin Signed-off-by: Nick Hill Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> Co-authored-by: Nick Hill --- tests/v1/sample/test_topk_topp_sampler.py | 38 +++++++++++++++++++++++ vllm/v1/sample/ops/topk_topp_sampler.py | 10 +++--- vllm/v1/sample/ops/topk_topp_triton.py | 32 +++++++++++-------- vllm/v1/worker/gpu/sample/logit_bias.py | 4 ++- vllm/v1/worker/gpu/sample/sampler.py | 26 +++++++++++----- 5 files changed, 85 insertions(+), 25 deletions(-) diff --git a/tests/v1/sample/test_topk_topp_sampler.py b/tests/v1/sample/test_topk_topp_sampler.py index 8a3d313f1d5..439ce0ea340 100644 --- a/tests/v1/sample/test_topk_topp_sampler.py +++ b/tests/v1/sample/test_topk_topp_sampler.py @@ -5,11 +5,13 @@ import torch from torch import Generator from tests.utils import large_gpu_mark +from vllm.model_executor.layers.vocab_parallel_embedding import pad_vocab_size from vllm.platforms import current_platform from vllm.triton_utils import HAS_TRITON from vllm.utils.torch_utils import set_random_seed from vllm.v1.sample.ops.topk_topp_sampler import ( apply_top_k_top_p_pytorch, + flashinfer_sample, random_sample, ) from vllm.v1.sample.sampler import Sampler @@ -1043,3 +1045,39 @@ class TestFlashInferDistributionMatch: f"{label}: distribution differs from theoretical: " f"chi2={chi2:.2f} p_value={p_value:.2e} alpha={self.ALPHA}" ) + + +@pytest.mark.skipif( + not FLASHINFER_TOPK_TOPP_SUPPORTED, + reason="FlashInfer top-k/top-p sampler is not available on this platform.", +) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("k, p", [(20, 0.95), (20, None), (None, 0.95)]) +def test_flashinfer_sample_padded_vocab( + dtype: torch.dtype, k: int | None, p: float | None +): + """flashinfer_sample must accept the logits the sampler actually hands it. + + compute_logits slices the padding off the vocab, so for a vocab that isn't a + multiple of 64 (e.g. opt's 50272) the logits are a strided view in the model + dtype, while FlashInfer requires contiguous fp32. + """ + torch.set_default_device(DEVICE_TYPE) + batch_size = 8 + org_vocab_size = 50272 + padded_vocab_size = pad_vocab_size(org_vocab_size) + assert padded_vocab_size != org_vocab_size + + logits = torch.randn(batch_size, padded_vocab_size, dtype=dtype)[ + ..., :org_vocab_size + ] + # A single row stays contiguous despite the padded stride, hence batch_size > 1. + assert not logits.is_contiguous() + + token_ids = flashinfer_sample( + logits, + torch.full((batch_size,), k, dtype=torch.int32) if k is not None else None, + torch.full((batch_size,), p, dtype=torch.float32) if p is not None else None, + ) + assert token_ids.shape == (batch_size,) + assert torch.all((token_ids >= 0) & (token_ids < org_vocab_size)) diff --git a/vllm/v1/sample/ops/topk_topp_sampler.py b/vllm/v1/sample/ops/topk_topp_sampler.py index 69b35830add..0f3dc79b98a 100644 --- a/vllm/v1/sample/ops/topk_topp_sampler.py +++ b/vllm/v1/sample/ops/topk_topp_sampler.py @@ -392,8 +392,9 @@ def apply_top_k_top_p_pytorch( logits_sort.masked_fill_(top_k_mask, -float("inf")) if p is not None: - # Apply top-p. - probs_sort = logits_sort.softmax(dim=-1) + # Apply top-p. The cumsum below runs over the whole vocab, so accumulating + # in a low-precision dtype makes the nucleus undershoot p. + probs_sort = logits_sort.softmax(dim=-1, dtype=torch.float32) probs_sum = torch.cumsum(probs_sort, dim=-1, out=probs_sort) top_p_mask = probs_sum <= 1 - p.unsqueeze(dim=1) # at least one @@ -500,9 +501,10 @@ def flashinfer_sample( probs, k, deterministic=True ) else: - # Both top-k and top-p. + # Both top-k and top-p. FlashInfer requires contiguous fp32 logits; the + # branches above get that from softmax(). next_token_ids = flashinfer.sampling.top_k_top_p_sampling_from_logits( - logits, k, p, deterministic=True + logits.float().contiguous(), k, p, deterministic=True ) return next_token_ids.view(-1) diff --git a/vllm/v1/sample/ops/topk_topp_triton.py b/vllm/v1/sample/ops/topk_topp_triton.py index c284ff61876..1b9058ba160 100755 --- a/vllm/v1/sample/ops/topk_topp_triton.py +++ b/vllm/v1/sample/ops/topk_topp_triton.py @@ -131,7 +131,7 @@ def _topk_topp_kernel( mask_n = offs < VOCAB_SIZE logits_blk0 = tl.load( LOGITS_ROW + offs, mask=mask_n, other=-float("inf") - ) + ).to(tl.float32) # Exclude -inf values (e.g. from grammar bitmasks) from # statistics to avoid NaN in pivot computation. finite_mask = (logits_blk0 > -float("inf")) & mask_n @@ -164,7 +164,7 @@ def _topk_topp_kernel( mask_n = offs_n < VOCAB_SIZE logits_blk = tl.load( LOGITS_ROW + offs_n, mask=mask_n, other=-float("inf") - ) + ).to(tl.float32) max_logit = tl.maximum(max_logit, tl.max(logits_blk)) # Exclude -inf from min to keep binary search bounds @@ -305,7 +305,7 @@ def _topk_topp_kernel( mask_n = offs_n < VOCAB_SIZE logits_blk2 = tl.load( LOGITS_ROW + offs_n, mask=mask_n, other=-float("inf") - ) + ).to(tl.float32) above_0 = logits_blk2 > k_pivot_0 above_1 = logits_blk2 > k_pivot_1 @@ -457,7 +457,7 @@ def _topk_topp_kernel( LOGITS_ROW + offs_n, mask=mask_n, other=-float("inf"), - ) + ).to(tl.float32) outlier_mask = (probs_blk > min_logit) & mask_n @@ -600,7 +600,7 @@ def _topk_topp_kernel( mask_n = offs < VOCAB_SIZE logits_blk0 = tl.load( LOGITS_ROW + offs, mask=mask_n, other=-float("inf") - ) + ).to(tl.float32) # Exclude -inf values (e.g. from grammar bitmasks) from # statistics to avoid NaN in pivot computation. finite_mask = (logits_blk0 > -float("inf")) & mask_n @@ -626,7 +626,7 @@ def _topk_topp_kernel( mask_n = offs_n < VOCAB_SIZE logits_blk = tl.load( LOGITS_ROW + offs_n, mask=mask_n, other=-float("inf") - ) + ).to(tl.float32) max_logit = tl.maximum(max_logit, tl.max(logits_blk)) # Exclude -inf from min to keep binary search bounds # finite (avoids NaN pivots). @@ -660,7 +660,7 @@ def _topk_topp_kernel( probs_blk = tl.load( LOGITS_ROW + offs_n, mask=mask_n, other=-float("inf") - ) + ).to(tl.float32) probs_blk = tl.exp(probs_blk - max_sample) probs_blk = probs_blk / sum_exp_logits @@ -754,7 +754,7 @@ def _topk_topp_kernel( probs_blk = tl.load( LOGITS_ROW + offs_n, mask=mask_n, other=-float("inf") - ) + ).to(tl.float32) probs_blk = tl.exp(probs_blk - max_sample) probs_blk = probs_blk / sum_exp_logits tl.store(BUFFER_ROW + offs_n, probs_blk, mask=mask_n) @@ -835,7 +835,7 @@ def _topk_topp_kernel( mask_n = offs_n < VOCAB_SIZE logits_blk = tl.load( LOGITS_ROW + offs_n, mask=mask_n, other=-float("inf") - ) + ).to(tl.float32) keep_mask = (logits_blk > final_pivot) & mask_n # Duplicate logit handling @@ -878,7 +878,7 @@ def apply_top_k_top_p_triton( The masked logits tensor. It may or may not be modified in-place. """ assert logits.ndim == 2 - assert logits.dtype == torch.float32 + assert logits.dtype in (torch.float32, torch.bfloat16, torch.float16) batch_size, vocab_size = logits.shape topk_enabled = k is not None topp_enabled = p is not None @@ -911,7 +911,9 @@ def apply_top_k_top_p_triton( buffer = _TRITON_BUFFER_CACHE.get(buf_key) if buffer is None or buffer.shape[0] < NUM_PROGRAMS: size = min(next_power_of_2(NUM_PROGRAMS), num_sm) - buffer = logits.new_empty((size, vocab_size)) + buffer = torch.empty( + (size, vocab_size), dtype=torch.float32, device=logits.device + ) _TRITON_BUFFER_CACHE[buf_key] = buffer if buffer.shape[0] > NUM_PROGRAMS: buffer = buffer[:NUM_PROGRAMS] @@ -919,8 +921,12 @@ def apply_top_k_top_p_triton( # Cache lookup table entries on each device. tables = _TRITON_TABLE_CACHE.get(logits.device) if tables is None: - normal_cdf_to_sigma_table = logits.new_tensor(_NORMAL_CDF_TO_SIGMA_TABLE) - percentile_to_std_table = logits.new_tensor(_PERCENTILE_TO_STD_TABLE) + normal_cdf_to_sigma_table = torch.tensor( + _NORMAL_CDF_TO_SIGMA_TABLE, dtype=torch.float32, device=logits.device + ) + percentile_to_std_table = torch.tensor( + _PERCENTILE_TO_STD_TABLE, dtype=torch.float32, device=logits.device + ) _TRITON_TABLE_CACHE[logits.device] = ( normal_cdf_to_sigma_table, percentile_to_std_table, diff --git a/vllm/v1/worker/gpu/sample/logit_bias.py b/vllm/v1/worker/gpu/sample/logit_bias.py index 6c95ed7aacb..1812251d47f 100644 --- a/vllm/v1/worker/gpu/sample/logit_bias.py +++ b/vllm/v1/worker/gpu/sample/logit_bias.py @@ -218,7 +218,9 @@ def _bias_kernel( mask=mask, ) bias = tl.load(bias_ptr + req_state_idx * bias_stride + block, mask=mask) - logits = tl.load(logits_ptr + token_idx * logits_stride + token_ids, mask=mask) + logits = tl.load( + logits_ptr + token_idx * logits_stride + token_ids, mask=mask + ).to(tl.float32) logits += bias tl.store(logits_ptr + token_idx * logits_stride + token_ids, logits, mask=mask) diff --git a/vllm/v1/worker/gpu/sample/sampler.py b/vllm/v1/worker/gpu/sample/sampler.py index f0a83c92efb..d9d67541107 100644 --- a/vllm/v1/worker/gpu/sample/sampler.py +++ b/vllm/v1/worker/gpu/sample/sampler.py @@ -83,11 +83,7 @@ class Sampler: # that num_nans is computed before applying penalties and temperature. num_nans = get_num_nans(logits) if self.compute_nans else None - max_num_logprobs = self.sampling_states.max_num_logprobs(idx_mapping_np) - max_per_req_token_ids = self.logprob_token_ids_state.max_num_token_ids( - idx_mapping_np - ) - return_logprobs = max_num_logprobs != NO_LOGPROBS or max_per_req_token_ids > 0 + return_logprobs = self.returns_logprobs(idx_mapping_np) sampled, processed_logits = self.sample( logits, @@ -102,6 +98,10 @@ class Sampler: if return_logprobs: if self.logprobs_mode in ("processed_logprobs", "processed_logits"): logits = processed_logits + max_num_logprobs = self.sampling_states.max_num_logprobs(idx_mapping_np) + max_per_req_token_ids = self.logprob_token_ids_state.max_num_token_ids( + idx_mapping_np + ) expanded_logits = logits.shape[0] != idx_mapping_np.shape[0] cu_num_logits = cu_num_logits_np.tolist() if expanded_logits else None num_logprobs = max_num_logprobs if max_num_logprobs != NO_LOGPROBS else 0 @@ -142,6 +142,13 @@ class Sampler: ) return sampler_output + def returns_logprobs(self, idx_mapping_np: np.ndarray) -> bool: + """Whether any request in the batch produces logprobs this step.""" + return ( + self.sampling_states.max_num_logprobs(idx_mapping_np) != NO_LOGPROBS + or self.logprob_token_ids_state.max_num_token_ids(idx_mapping_np) > 0 + ) + def apply_sampling_params( self, logits: torch.Tensor, @@ -152,8 +159,13 @@ class Sampler: expanded_local_pos: torch.Tensor, skip_top_k_top_p: bool = False, ) -> torch.Tensor: - # Copy logits to a new FP32 tensor. - logits = torch.empty_like(logits, dtype=torch.float32).copy_(logits) + # The ops below upcast to fp32 internally, so the input dtype is kept and + # mutated in place. Only raw_logprobs reads the unmodified logits + # afterward, so copy just for that case. + if self.logprobs_mode.startswith("raw_") and self.returns_logprobs( + idx_mapping_np + ): + logits = logits.clone() # Apply logit bias (e.g., allowed_token_ids, min_tokens) in place. self.logit_bias_state.apply_logit_bias( From da64db78b989d820effd090d8ac4b65ade2cbe4f Mon Sep 17 00:00:00 2001 From: Jee Jee Li Date: Sat, 18 Jul 2026 10:26:14 +0800 Subject: [PATCH 16/51] [LoRA] Optimize TrtLlmLoRAExperts (#48759) Signed-off-by: Jee Jee Li --- .../fused_moe/experts/lora_experts_mixin.py | 12 +- .../fused_moe/experts/trtllm_lora_moe.py | 250 +++++++++++++++--- 2 files changed, 221 insertions(+), 41 deletions(-) diff --git a/vllm/model_executor/layers/fused_moe/experts/lora_experts_mixin.py b/vllm/model_executor/layers/fused_moe/experts/lora_experts_mixin.py index b8dc80ed181..0aa6f6c30b6 100644 --- a/vllm/model_executor/layers/fused_moe/experts/lora_experts_mixin.py +++ b/vllm/model_executor/layers/fused_moe/experts/lora_experts_mixin.py @@ -46,6 +46,7 @@ class LoRAExpertsMixin: num_tokens: int, top_k_num: int, add_inputs: bool = True, + swap_w13_slices: bool = False, ) -> tuple[ torch.Tensor | None, torch.Tensor | None, @@ -53,6 +54,7 @@ class LoRAExpertsMixin: torch.Tensor | None, ]: w13_lora_a_stacked = lora_context.w13_lora_a_stacked + w13_lora_b_stacked = lora_context.w13_lora_b_stacked if lora_context.enable_moe_shared_loras: # w13 lora_A is shared across experts (collapsed expert-dim 1); # broadcast to local_num_experts via a stride-0 view. The kernel @@ -61,11 +63,19 @@ class LoRAExpertsMixin: a.expand(-1, lora_context.local_num_experts, -1, -1) for a in w13_lora_a_stacked ) + if swap_w13_slices: + # The expand kernel writes slice j into the j-th half of y's last + # dim. Reversing the (gate, up) slice tuples makes it emit + # [up, gate] order directly -- used by the FlashInfer trtllm path, + # whose SwiGLU expects the up half first, to avoid an out-of-place + # concat swap afterwards. + w13_lora_a_stacked = w13_lora_a_stacked[::-1] + w13_lora_b_stacked = w13_lora_b_stacked[::-1] return lora_context.punica_wrapper.add_lora_w13( y, x, w13_lora_a_stacked, - lora_context.w13_lora_b_stacked, + w13_lora_b_stacked, topk_ids, topk_weights, expert_map, diff --git a/vllm/model_executor/layers/fused_moe/experts/trtllm_lora_moe.py b/vllm/model_executor/layers/fused_moe/experts/trtllm_lora_moe.py index 5630f6ab9b5..1a947756ce1 100644 --- a/vllm/model_executor/layers/fused_moe/experts/trtllm_lora_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/trtllm_lora_moe.py @@ -31,6 +31,7 @@ from vllm.model_executor.layers.fused_moe.config import ( FusedMoEQuantConfig, RoutingMethodType, ) +from vllm.model_executor.layers.fused_moe.experts.lora_context import MoELoRAContext from vllm.model_executor.layers.fused_moe.experts.lora_experts_mixin import ( LoRAExpertsMixin, ) @@ -41,9 +42,75 @@ from vllm.model_executor.layers.fused_moe.utils import ( trtllm_moe_pack_topk_ids_weights, ) from vllm.platforms import current_platform +from vllm.triton_utils import tl, triton from vllm.utils.flashinfer import has_flashinfer_trtllm_fused_moe +@triton.jit +def _unpermute_activation_kernel( + act_ptr, # act_permuted: (num_permuted, num_cols) + idx_ptr, # idx_map: (num_rows,), values in [0, num_permuted) or -1 + out_ptr, # out: (num_rows, num_cols) + num_cols, + stride_ar, + stride_or, + BLOCK_I: tl.constexpr, +): + row = tl.program_id(0) + col_offs = tl.program_id(1) * BLOCK_I + tl.arange(0, BLOCK_I) + col_mask = col_offs < num_cols + + idx = tl.load(idx_ptr + row) + out_ptrs = out_ptr + row * stride_or + col_offs + if idx >= 0: + vals = tl.load(act_ptr + idx * stride_ar + col_offs, mask=col_mask, other=0.0) + tl.store(out_ptrs, vals, mask=col_mask) + else: + zeros = tl.zeros((BLOCK_I,), dtype=out_ptr.dtype.element_ty) + tl.store(out_ptrs, zeros, mask=col_mask) + + +@triton.jit +def _finalize_lora_kernel( + gemm2_ptr, # (num_permuted, K) base FC2 output, permuted, unweighted + weight_ptr, # (num_tokens * top_k,) routing weights (expanded order) + idx_ptr, # (num_tokens * top_k,) expanded_idx -> permuted_idx or -1 + delta_ptr, # (num_tokens, top_k, K) W2 LoRA delta, already routing-weighted + out_ptr, # (num_tokens, K) + K, + stride_g0, + stride_d0, + stride_d1, + stride_o0, + scale, + TOP_K: tl.constexpr, + BLOCK_K: tl.constexpr, +): + token = tl.program_id(0) + col = tl.program_id(1) * BLOCK_K + tl.arange(0, BLOCK_K) + mask = col < K + + acc_base = tl.zeros((BLOCK_K,), dtype=tl.float32) + acc_delta = tl.zeros((BLOCK_K,), dtype=tl.float32) + for k in tl.static_range(TOP_K): + eid = token * TOP_K + k + pidx = tl.load(idx_ptr + eid) + if pidx >= 0: + w = tl.load(weight_ptr + eid).to(tl.float32) + base = tl.load(gemm2_ptr + pidx * stride_g0 + col, mask=mask, other=0.0).to( + tl.float32 + ) + acc_base += w * base + acc_delta += tl.load( + delta_ptr + token * stride_d0 + k * stride_d1 + col, mask=mask, other=0.0 + ).to(tl.float32) + + out = acc_base * scale + acc_delta + tl.store( + out_ptr + token * stride_o0 + col, out.to(out_ptr.dtype.element_ty), mask=mask + ) + + class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular): """LoRA-aware trtllm MoE experts""" @@ -107,7 +174,8 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular): return E, M, N, K, topk def finalize_weight_and_reduce_impl(self) -> mk.TopKWeightAndReduce: - # do_finalize=True: the kernel already does moe_sum, so this is a No-Op + # apply() writes the fully finalized result into `output` (fused base + # finalize + W2 LoRA reduction), so this is a No-Op. return TopKWeightAndReduceNoOP() @staticmethod @@ -146,9 +214,13 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular): ) -> list[torch.Tensor]: """Call the dtype-specific trtllm_*_routed_moe and return list[Tensor]. - Return contract (do_finalize=True): - gemm1_lora_delta is None -> [output] - otherwise -> [output, expanded_idx_to_permuted_idx, + The LoRA path always sets gemm1_lora_delta and runs with + do_finalize=False so the base finalize can be fused with the W2 LoRA + reduction (see _finalize_with_w2_lora). Return contract: + gemm1_lora_delta is None -> [output] (do_finalize=True) + otherwise -> [gemm2_output(permuted, unweighted), + expert_weights, + expanded_idx_to_permuted_idx, gemm1_activation_output(permuted)] """ raise NotImplementedError @@ -178,6 +250,27 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular): intermediate_size = self.intermediate_size_per_partition K = output.size(1) + # Routing is computed outside the MoE; pack it into the + # (eid<<16)|w.bf16 format the routed API expects. + packed_topk_ids = trtllm_moe_pack_topk_ids_weights(topk_ids, topk_weights) + + # ---- Base-model fast path ---- + # When no token in the batch selects a LoRA adapter, skip the LoRA machinery + # and run the plain base MoE with do_finalize=True, which writes the finalized + # result straight into `output`. + if self._batch_has_no_lora(lora_context): + self.invoke_routed_moe( + hidden_states=hidden_states, + w1=w1, + w2=w2, + packed_topk_ids=packed_topk_ids, + gemm1_lora_delta=None, # without LoRA, no delta + global_num_experts=global_num_experts, + a1q_scale=a1q_scale, + output=output, + ) + return + # The LoRA tile-config heuristic (try_get_optimal_moe_config) unpacks # w1/w2 as standard 3D MoE weights, but flashinfer stores shuffled # 4D BlockMajorK weights. add_lora_w13/add_lora_w2 only read .shape @@ -195,14 +288,13 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular): dtype=torch.bfloat16, ) - # Routing is computed outside the MoE; pack it into the - # (eid<<16)|w.bf16 format the routed API expects. - packed_topk_ids = trtllm_moe_pack_topk_ids_weights(topk_ids, topk_weights) - # ---- 1) W13 LoRA delta -> gemm1_lora_delta (bf16, [T, top_k, 2I]) ---- gemm1_lora_delta = None w13_meta = (None, None, None, None) + # zeros (not empty): under EP the punica expand kernel only writes + # slots whose expert is local to this rank; non-local (token, top_k) + # slots must stay 0 so they contribute no bias when fed to flashinfer. gemm1_lora_delta = torch.zeros( num_tokens, top_k, @@ -221,6 +313,10 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular): # add_inputs=False: write the pure delta only (the base is fused in # by the kernel) and do NOT multiply by the routing weight (it is a # pre-SwiGLU bias). + # swap_w13_slices=True: apply_w13_lora writes in vLLM's w13 order + # (gate=w1 first, up=w3 second), but FlashInfer's gemm1_lora_delta + # expects [up, gate]; reversing the slices emits that order directly, + # avoiding an out-of-place concat swap. w13_meta = self.apply_w13_lora( lora_context, y=gemm1_lora_delta, @@ -233,18 +329,7 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular): num_tokens=num_tokens, top_k_num=top_k, add_inputs=False, - ) - - # apply_w13_lora writes the delta in vLLM's w13 order (gate=w1 first, - # up=w3 second), but FlashInfer's gemm1_lora_delta expects the halves - # in [up, gate] order. Swap them so the delta lands on the matching - # SwiGLU branch. - gemm1_lora_delta = torch.cat( - [ - gemm1_lora_delta[..., intermediate_size:], - gemm1_lora_delta[..., :intermediate_size], - ], - dim=-1, + swap_w13_slices=True, ) # ---- 2) Call the routed flashinfer kernel ---- @@ -259,8 +344,14 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular): output=output, ) # ---- 3) W2 LoRA (computed out of kernel) ---- - expanded_idx_to_permuted_idx = ret[1] - gemm1_act_permuted = ret[2] # [max_padded, I], post-act + # do_finalize=False: flashinfer returns the *unfinalized* base output. + # ret = [gemm2_output(permuted, unweighted), + # expert_weights, expanded_idx_to_permuted_idx, + # gemm1_activation_output(permuted)] + gemm2_permuted = ret[0] + expert_weights = ret[1] + expanded_idx_to_permuted_idx = ret[2] + gemm1_act_permuted = ret[3] # [max_padded, I], post-act act = self._unpermute_activation( gemm1_act_permuted, expanded_idx_to_permuted_idx, @@ -298,11 +389,37 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular): top_k_num=top_k, add_inputs=False, ) - # The base output is already finalized (routing-weighted + summed over - # top_k); the W2 delta is likewise already routing-weighted, so sum it - # over top_k and add. - # TODO(verify): if routed_scaling_factor is not None, scale to match. - output.add_(w2_delta.sum(dim=1)) + # ---- 4) Fused finalize: reduce the base path over top_k (with routing + # weights) and add the already-weighted W2 delta, in a single kernel. + # This replaces flashinfer's internal finalize launch plus a separate + # w2_delta.sum(dim=1) + add_. + self._finalize_with_w2_lora( + output, + gemm2_permuted, + expert_weights, + expanded_idx_to_permuted_idx, + w2_delta, + num_tokens, + top_k, + scale=1.0, + ) + + @staticmethod + def _batch_has_no_lora(lora_context: MoELoRAContext) -> bool: + """True when no token in the batch selects a LoRA adapter. + + Mirrors the no-lora fast path in + ``PunicaWrapperGPU.add_lora_fused_moe``: the punica kernel metadata + carries a CPU ``no_lora_flag`` computed once per forward from the + token->LoRA mapping. Reading it is a host-only check (no device sync), + and under CUDA graphs the branch is frozen at capture time against the + graph's ``has_lora`` dispatch key, so it stays correct on replay. + """ + meta = getattr(lora_context.punica_wrapper, "token_mapping_meta", None) + if meta is None: + return False + flag = meta.no_lora_flag_cpu + return bool(flag.numel() == 1 and flag.item()) @staticmethod def _unpermute_activation( @@ -315,13 +432,63 @@ class _TrtLlmLoRAExpertsBase(LoRAExpertsMixin, mk.FusedMoEExpertsModular): """Permuted FC1 activation -> (num_tokens*top_k, I). expanded_idx = token*top_k + k; idx_map[expanded_idx] = permuted_idx or -1. - TODO optimize these operations + Fused gather + drop-masking: each output row copies the matching + permuted row, or is zeroed when idx_map < 0. """ + num_rows = num_tokens * top_k + out = torch.empty( + (num_rows, intermediate_size), + dtype=act_permuted.dtype, + device=act_permuted.device, + ) + BLOCK_I = 1024 + grid = (num_rows, triton.cdiv(intermediate_size, BLOCK_I)) + _unpermute_activation_kernel[grid]( + act_permuted, + idx_map, + out, + intermediate_size, + act_permuted.stride(0), + out.stride(0), + BLOCK_I=BLOCK_I, + ) + return out - valid = idx_map >= 0 - safe_idx = idx_map.clamp_min(0).long() - gathered = act_permuted[safe_idx] - return gathered * valid.unsqueeze(1).to(act_permuted.dtype) + @staticmethod + def _finalize_with_w2_lora( + output: torch.Tensor, + gemm2_permuted: torch.Tensor, + expert_weights: torch.Tensor, + idx_map: torch.Tensor, + w2_delta: torch.Tensor, + num_tokens: int, + top_k: int, + scale: float = 1.0, + ) -> None: + """Fused base finalize + W2 LoRA reduction, written into ``output``. + + For each token: sum the routing-weighted permuted base rows over top_k + (``expert_weights`` in expanded order, ``idx_map < 0`` dropped), scale by + ``scale``, and add the already-weighted ``w2_delta`` reduced over top_k. + """ + K = gemm2_permuted.size(1) + BLOCK_K = 512 + grid = (num_tokens, triton.cdiv(K, BLOCK_K)) + _finalize_lora_kernel[grid]( + gemm2_permuted, + expert_weights.reshape(-1), + idx_map, + w2_delta, + output, + K, + gemm2_permuted.stride(0), + w2_delta.stride(0), + w2_delta.stride(1), + output.stride(0), + scale, + TOP_K=top_k, + BLOCK_K=BLOCK_K, + ) # BF16 unquantized trtllm MoE + LoRA @@ -359,10 +526,12 @@ class TrtLlmBf16LoRAExperts(_TrtLlmLoRAExpertsBase): ) -> list[torch.Tensor]: import flashinfer - # Unlike the fp8/mxint4 routed APIs, trtllm_bf16_routed_moe has no - # `output=` kwarg: it returns the finalized tensor (or a list whose - # [0] is it when gemm1_lora_delta is set). Copy it into the caller's - # buffer so the modular-kernel output plumbing sees the result. + # With gemm1_lora_delta set (the LoRA path) run do_finalize=False and + # return the unfinalized permuted base output so apply() can fuse the + # finalize with the W2 LoRA reduction (see _finalize_with_w2_lora). + # Without a delta (base path), run do_finalize=True and hand flashinfer + # the caller's buffer via output= so it finalizes in place -- no copy. + do_finalize = gemm1_lora_delta is None ret = flashinfer.fused_moe.trtllm_bf16_routed_moe( topk_ids=packed_topk_ids, hidden_states=hidden_states, @@ -378,10 +547,11 @@ class TrtLlmBf16LoRAExperts(_TrtLlmLoRAExpertsBase): local_num_experts=self.local_num_experts, routed_scaling_factor=None, routing_method_type=self.routing_method_type, - do_finalize=True, + do_finalize=do_finalize, + output=output if do_finalize else None, ) - if isinstance(ret, (list, tuple)): - output.copy_(ret[0]) + if not do_finalize: + # [gemm2_output, expert_weights, expanded_idx, gemm1_activation] return list(ret) - output.copy_(ret) + # do_finalize=True finalized directly into `output`. return [output] From f12b80c6efa116c7c4020ee611d622ab7030ca40 Mon Sep 17 00:00:00 2001 From: xuebwang-amd Date: Sat, 18 Jul 2026 11:49:41 +0800 Subject: [PATCH 17/51] [ROCm][Bugfix] Fix GPT-OSS Quark MXFP4 MoE loading - emulation buffer not block-aligned (#43979) Signed-off-by: xuebwang-amd --- tests/kernels/moe/test_ocp_mx_moe.py | 46 +++++++++++++++++++ .../layers/fused_moe/oracle/mxfp4.py | 11 ++++- .../layers/quantization/quark/quark_moe.py | 9 ++-- 3 files changed, 59 insertions(+), 7 deletions(-) diff --git a/tests/kernels/moe/test_ocp_mx_moe.py b/tests/kernels/moe/test_ocp_mx_moe.py index 49682dbfcd6..2f819c09aaa 100644 --- a/tests/kernels/moe/test_ocp_mx_moe.py +++ b/tests/kernels/moe/test_ocp_mx_moe.py @@ -1512,3 +1512,49 @@ def test_rocm_mxfp4_moe_oracle( # Check accuracy using per-backend thresholds check_accuracy(ref, out, atol=0.1, rtol=config["rtol"], percent=config["percent"]) + + +# ----------------------------------------------------------------------------- +# MXFP4 emulation size-rounding tests +# ----------------------------------------------------------------------------- +# Emulation needs each per-partition dim rounded up to OCP_MX_BLOCK_SIZE (32); +# a non-block-aligned shard (e.g. GPT-OSS 2880 // 4 = 720) otherwise truncates +# the scale buffer and fails weight loading. +# NOTE: gated to ROCm since it is the emulation backend's current target; +# remove this skip if the backend is enabled on non-ROCm platforms. +@pytest.mark.skipif(not ROCM_AVAILABLE, reason="emulation backend targets ROCm") +@pytest.mark.parametrize( + "hidden_size,intermediate_size,expected_hidden,expected_intermediate", + [ + (2880, 720, 2880, 736), # GPT-OSS TP=4 shard: 720 -> round_up(720, 32) + (2880, 360, 2880, 384), # GPT-OSS TP=8 shard: 360 -> 384 + (2880, 2880, 2880, 2880), # already block-aligned: unchanged + (90, 90, 96, 96), # both dims unaligned + ], +) +def test_mxfp4_emulation_rounds_up_to_block_size( + hidden_size: int, + intermediate_size: int, + expected_hidden: int, + expected_intermediate: int, +): + """Emulation must block-align per-partition dims to OCP_MX_BLOCK_SIZE.""" + from vllm.model_executor.layers.fused_moe.oracle.mxfp4 import ( + Mxfp4MoeBackend, + mxfp4_round_up_hidden_size_and_intermediate_size, + ) + from vllm.model_executor.layers.quantization.utils.ocp_mx_utils import ( + OCP_MX_BLOCK_SIZE, + ) + + rounded_hidden, rounded_intermediate = ( + mxfp4_round_up_hidden_size_and_intermediate_size( + Mxfp4MoeBackend.EMULATION, hidden_size, intermediate_size + ) + ) + + assert rounded_hidden == expected_hidden + assert rounded_intermediate == expected_intermediate + # The block-scale buffer (dim // OCP_MX_BLOCK_SIZE) must not floor-truncate. + assert rounded_hidden % OCP_MX_BLOCK_SIZE == 0 + assert rounded_intermediate % OCP_MX_BLOCK_SIZE == 0 diff --git a/vllm/model_executor/layers/fused_moe/oracle/mxfp4.py b/vllm/model_executor/layers/fused_moe/oracle/mxfp4.py index 55a767f060c..921e8f114d5 100644 --- a/vllm/model_executor/layers/fused_moe/oracle/mxfp4.py +++ b/vllm/model_executor/layers/fused_moe/oracle/mxfp4.py @@ -27,6 +27,9 @@ from vllm.model_executor.layers.fused_moe.config import ( ocp_mx_moe_quant_config, ) from vllm.model_executor.layers.quantization.utils.mxfp4_utils import _swizzle_mxfp4 +from vllm.model_executor.layers.quantization.utils.ocp_mx_utils import ( + OCP_MX_BLOCK_SIZE, +) from vllm.model_executor.layers.quantization.utils.quant_utils import ( QuantKey, kFp8Dynamic128Sym, @@ -633,7 +636,13 @@ def mxfp4_round_up_hidden_size_and_intermediate_size( backend: Mxfp4MoeBackend, hidden_size: int, intermediate_size: int ) -> tuple[int, int]: """Round up hidden_size and intermediate_size based on backend requirements.""" - if backend == Mxfp4MoeBackend.DEEPGEMM_MXFP4: + if backend == Mxfp4MoeBackend.EMULATION: + # Emulation has no kernel tile; it only needs OCP MX block alignment so the + # per-block scale buffers (`dim // OCP_MX_BLOCK_SIZE`) aren't floor-truncated + # by a non-block-aligned TP/DP shard (e.g. 2880 // 4 = 720). + intermediate_size = round_up(intermediate_size, OCP_MX_BLOCK_SIZE) + hidden_size = round_up(hidden_size, OCP_MX_BLOCK_SIZE) + elif backend == Mxfp4MoeBackend.DEEPGEMM_MXFP4: # DeepGEMM requires M/N/K alignment intermediate_size = round_up(intermediate_size, 128) hidden_size = round_up(hidden_size, 128) diff --git a/vllm/model_executor/layers/quantization/quark/quark_moe.py b/vllm/model_executor/layers/quantization/quark/quark_moe.py index 7bdf963b512..15023d7ca39 100644 --- a/vllm/model_executor/layers/quantization/quark/quark_moe.py +++ b/vllm/model_executor/layers/quantization/quark/quark_moe.py @@ -1063,12 +1063,9 @@ class QuarkOCP_MX_MoEMethod(QuarkMoEMethod): act_dtype=act_dtype, moe_parallel_config=moe_parallel_config, ) - # In case quantization emulation backend is used, there is no need to apply - # MXFP4-specific padding logic as the compute happens in higher precision. - if ( - self.mxfp4_backend is not None - and self.mxfp4_backend != Mxfp4MoeBackend.EMULATION - ): + # Round per-partition sizes up to each backend's requirement. Emulation is + # handled inside the helper too (OCP MX block alignment), so no special-case. + if self.mxfp4_backend is not None: hidden_size, intermediate_size_per_partition = ( mxfp4_round_up_hidden_size_and_intermediate_size( self.mxfp4_backend, hidden_size, intermediate_size_per_partition From c71a583aa9f81400528e67e3d818f66b804e8340 Mon Sep 17 00:00:00 2001 From: Francesco Fusco Date: Sat, 18 Jul 2026 06:43:09 +0200 Subject: [PATCH 18/51] [Perf][Hybrid] Vectorize _copy_mamba_state_block to uint64 for temporal (#48110) --- vllm/v1/worker/mamba_utils.py | 70 +++++++++++++++++++++++++++-------- 1 file changed, 54 insertions(+), 16 deletions(-) diff --git a/vllm/v1/worker/mamba_utils.py b/vllm/v1/worker/mamba_utils.py index 8d8d3e62a9d..7f611a2eeff 100644 --- a/vllm/v1/worker/mamba_utils.py +++ b/vllm/v1/worker/mamba_utils.py @@ -109,24 +109,45 @@ def _copy_mamba_state_block( src_addr = state_base_addr + src_block_id * state_block_stride + src_offset num_elems_to_copy = (conv_width - token_bias).to(tl.int64) * state_inner_size copy_size = num_elems_to_copy * state_elem_size - else: - # Temporal state: copy state[bt[src_col + token_bias]] -> state[bt[dst_col]] - actual_src_block_id = tl.load(block_table_base + src_col + token_bias).to( - tl.int64 - ) - src_addr = state_base_addr + actual_src_block_id * state_block_stride - # Use natural block data size (inner_size * elem_size), NOT - # state_block_stride which is the page stride and can exceed the - # actual data when the state tensor uses as_strided page padding. - copy_size = state_inner_size * state_elem_size + offsets = tl.arange(0, COPY_BLOCK_SIZE) + for i in range(0, copy_size, COPY_BLOCK_SIZE): + mask = (i + offsets) < copy_size + curr_src = (src_addr + i + offsets).to(tl.pointer_type(tl.uint8)) + curr_dst = (dst_addr + i + offsets).to(tl.pointer_type(tl.uint8)) + data = tl.load(curr_src, mask=mask) + tl.store(curr_dst, data, mask=mask) + return + # Temporal state: copy state[bt[src_col + token_bias]] -> state[bt[dst_col]] + actual_src_block_id = tl.load(block_table_base + src_col + token_bias).to(tl.int64) + src_addr = state_base_addr + actual_src_block_id * state_block_stride + # Use natural block data size (inner_size * elem_size), NOT + # state_block_stride which is the page stride and can exceed the + # actual data when the state tensor uses as_strided page padding. + copy_size = state_inner_size * state_elem_size + + # Vectorize via uint64 (8B per thread → LDG.64/STG.64): both temporal + # and SD conv produce src/dst addresses aligned to a full token slice + # (inner_size * elem_size) and a copy_size that's a multiple of it, + # which is 8B-aligned for all state dtypes in use. A masked byte tail + # covers any remaining 0-7 bytes (only reachable for sub-8B slices). + copy_size_u64 = copy_size // 8 + src_u64 = src_addr.to(tl.pointer_type(tl.uint64)) + dst_u64 = dst_addr.to(tl.pointer_type(tl.uint64)) offsets = tl.arange(0, COPY_BLOCK_SIZE) - for i in range(0, copy_size, COPY_BLOCK_SIZE): - mask = (i + offsets) < copy_size - curr_src = (src_addr + i + offsets).to(tl.pointer_type(tl.uint8)) - curr_dst = (dst_addr + i + offsets).to(tl.pointer_type(tl.uint8)) - data = tl.load(curr_src, mask=mask) - tl.store(curr_dst, data, mask=mask) + for i in range(0, copy_size_u64, COPY_BLOCK_SIZE): + mask = (i + offsets) < copy_size_u64 + data = tl.load(src_u64 + i + offsets, mask=mask) + tl.store(dst_u64 + i + offsets, data, mask=mask) + + tail_start = copy_size_u64 * 8 + tail_bytes = copy_size - tail_start + tail_off = tl.arange(0, 8) + tail_src = (src_addr + tail_start).to(tl.pointer_type(tl.uint8)) + tail_dst = (dst_addr + tail_start).to(tl.pointer_type(tl.uint8)) + tail_mask = tail_off < tail_bytes + tail_data = tl.load(tail_src + tail_off, mask=tail_mask) + tl.store(tail_dst + tail_off, tail_data, mask=tail_mask) @triton.jit @@ -674,6 +695,23 @@ class MambaSpecDecodeGPUContext: self.state_inner_sizes[idx] = ( state[0].numel() if state.dim() > 1 else 1 ) + # Temporal copies are vectorized with uint64 + # loads/stores; base pointer and block stride must + # be 8B-aligned (tail loop handles copy_size % 8). + base_addr = state.data_ptr() + block_stride_bytes = block_stride_elems * state.element_size() + assert base_addr % 8 == 0, ( + f"layer {layer_name}: state.data_ptr() = " + f"{base_addr:#x} is not 8B-aligned; " + f"_copy_mamba_state_block uint64 " + f"vectorization requires it" + ) + assert block_stride_bytes % 8 == 0, ( + f"layer {layer_name}: block stride = " + f"{block_stride_bytes}B is not 8B-aligned; " + f"_copy_mamba_state_block uint64 " + f"vectorization requires it" + ) self.state_group_indices[idx] = group_local_idx idx += 1 From d96aee09518a57fdfe1e83369d2f4ab515f78c30 Mon Sep 17 00:00:00 2001 From: alexxu-roblox Date: Sat, 18 Jul 2026 01:40:06 -0700 Subject: [PATCH 19/51] [Bugfix] Re-sync parameter tp_rank after process_weights_after_loading (fix replicated / disable_tp weight reload) (#48025) Signed-off-by: Alex Xu Co-authored-by: YQ-Wang Co-authored-by: alexhxu Co-authored-by: Cursor --- vllm/model_executor/layers/linear.py | 18 +++++++++++------- .../model_loader/reload/layerwise.py | 5 +++++ vllm/model_executor/model_loader/utils.py | 7 +++++++ 3 files changed, 23 insertions(+), 7 deletions(-) diff --git a/vllm/model_executor/layers/linear.py b/vllm/model_executor/layers/linear.py index d92b9fc7d00..c0dbc776acb 100644 --- a/vllm/model_executor/layers/linear.py +++ b/vllm/model_executor/layers/linear.py @@ -281,6 +281,15 @@ class LinearBase(PluggableLayer): self.tp_size = get_tensor_model_parallel_world_size() if not disable_tp else 1 def update_param_tp_status(self): + # Single source of truth for a parameter's TP state. BasevLLMParameter + # stamps self.tp_rank with the *global* rank in __init__; this reconciles + # every child parameter to the *layer's* tp_rank/tp_size (which correctly + # accounts for disable_tp -> replicated weights with tp_rank == 0). + # + # Must be re-run whenever parameters are (re-)created after construction, + # e.g. after quant_method.process_weights_after_loading() swaps in fresh + # Parameters. Otherwise a later load_weights()/weight-refit would narrow a + # replicated weight at global_rank * shard_size and overflow. for param in self.parameters(): if isinstance(param, BasevLLMParameter): param.tp_rank = self.tp_rank @@ -910,7 +919,6 @@ class MergedColumnParallelLinear(ColumnParallelLinear): shard_id=loaded_shard_id, shard_offset=shard_offset, shard_size=shard_size, - tp_rank=self.tp_rank, ) def load_weights( @@ -1110,12 +1118,10 @@ class QKVParallelLinear(ColumnParallelLinear): # to ensure that any subsequent reduction (like .max()) # works correctly while preserving the parameter shape. for idx in range(param.data.shape[0]): - param.load_qkv_weight( - loaded_weight=loaded_weight, shard_id=idx, tp_rank=self.tp_rank - ) + param.load_qkv_weight(loaded_weight=loaded_weight, shard_id=idx) return elif type(param) in (RowvLLMParameter, BasevLLMParameter): - param.load_qkv_weight(loaded_weight=loaded_weight, tp_rank=self.tp_rank) + param.load_qkv_weight(loaded_weight=loaded_weight) return # TODO: @dsikka - move to parameter.py self._load_fused_module_from_checkpoint(param, loaded_weight) @@ -1139,7 +1145,6 @@ class QKVParallelLinear(ColumnParallelLinear): shard_id=loaded_shard_id, shard_offset=shard_offset, shard_size=shard_size, - tp_rank=self.tp_rank, ) def weight_loader( @@ -1493,7 +1498,6 @@ class MinimaxM3QKVParallelLinearWithIndexer(QKVParallelLinear): shard_id=loaded_shard_id, shard_offset=shard_offset, shard_size=shard_size, - tp_rank=self.tp_rank, ) def weight_loader( diff --git a/vllm/model_executor/model_loader/reload/layerwise.py b/vllm/model_executor/model_loader/reload/layerwise.py index eb609d63c50..92f454f9a5f 100644 --- a/vllm/model_executor/model_loader/reload/layerwise.py +++ b/vllm/model_executor/model_loader/reload/layerwise.py @@ -364,6 +364,11 @@ def _layerwise_process(layer: torch.nn.Module, info: LayerReloadingInfo): quant_method = getattr(layer, "quant_method", None) if isinstance(quant_method, QuantizeMethodBase): quant_method.process_weights_after_loading(layer) + # Re-reconcile parameter TP state: process_weights_after_loading may + # have re-created Parameters (stamped with the global rank), which would + # otherwise break replicated (disable_tp) weights on a subsequent reload. + if hasattr(layer, "update_param_tp_status"): + layer.update_param_tp_status() # Copy processed values into original tensor storage (preserves cudagraph refs) # this code is a no-op if not reloading (because kernel tensors is empty) diff --git a/vllm/model_executor/model_loader/utils.py b/vllm/model_executor/model_loader/utils.py index 6be057bff08..3367f4833e6 100644 --- a/vllm/model_executor/model_loader/utils.py +++ b/vllm/model_executor/model_loader/utils.py @@ -111,6 +111,13 @@ def process_weights_after_loading( # parameters onto device for processing and back off after. with device_loading_context(module, target_device): quant_method.process_weights_after_loading(module) + # process_weights_after_loading may swap in freshly-created + # Parameters (e.g. FP8 requantization), which are stamped with the + # global rank in BasevLLMParameter.__init__. Re-reconcile their TP + # state to the layer so a later weight reload / RL weight-refit + # narrows replicated (disable_tp) weights at the correct offset. + if hasattr(module, "update_param_tp_status"): + module.update_param_tp_status() # Repacking transients above can leave large amounts of memory in # the caching allocator, which starves the OS on UMA devices. release_device_memory_under_pressure(target_device) From c233d90aa826df072872df47b201450059be8e71 Mon Sep 17 00:00:00 2001 From: Harry Mellor <19981378+hmellor@users.noreply.github.com> Date: Sat, 18 Jul 2026 09:40:27 +0100 Subject: [PATCH 20/51] Remove even more unnecessary `load_weights` methods (#48496) Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com> --- vllm/model_executor/models/AXK1.py | 14 +- vllm/model_executor/models/aria.py | 150 ++--------- vllm/model_executor/models/bailing_moe.py | 116 ++------- vllm/model_executor/models/bloom.py | 43 ++-- vllm/model_executor/models/cohere2_moe.py | 4 +- vllm/model_executor/models/deepencoder.py | 13 +- vllm/model_executor/models/deepseek_mtp.py | 7 +- vllm/model_executor/models/deepseek_v2.py | 19 +- vllm/model_executor/models/ernie45_moe.py | 167 +++---------- vllm/model_executor/models/funaudiochat.py | 49 ++-- vllm/model_executor/models/gemma3n.py | 79 ++---- vllm/model_executor/models/glm4.py | 13 - vllm/model_executor/models/glm4_moe.py | 232 ++---------------- vllm/model_executor/models/glm4_moe_lite.py | 14 +- .../models/glm4_moe_lite_mtp.py | 3 +- vllm/model_executor/models/glm4_moe_mtp.py | 3 +- vllm/model_executor/models/glm_ocr_mtp.py | 3 +- vllm/model_executor/models/glmasr.py | 54 ++-- vllm/model_executor/models/gpt2.py | 43 ++-- vllm/model_executor/models/gpt_neox.py | 54 ++-- vllm/model_executor/models/hy_v3.py | 16 +- vllm/model_executor/models/hy_v3_mtp.py | 8 +- vllm/model_executor/models/intern_vit.py | 12 +- vllm/model_executor/models/interns1_vit.py | 13 +- vllm/model_executor/models/jamba.py | 99 +------- vllm/model_executor/models/kimi_linear.py | 14 +- vllm/model_executor/models/laguna.py | 207 +++------------- vllm/model_executor/models/lfm2_moe.py | 127 ++-------- vllm/model_executor/models/molmo.py | 17 +- vllm/model_executor/models/molmo2.py | 17 +- vllm/model_executor/models/nemotron_h.py | 112 +-------- vllm/model_executor/models/olmo_hybrid.py | 97 ++------ vllm/model_executor/models/openpangu.py | 191 +++----------- vllm/model_executor/models/param2moe.py | 217 +++------------- vllm/model_executor/models/qwen3_5.py | 19 +- vllm/model_executor/models/qwen3_5_mtp.py | 21 +- vllm/model_executor/models/qwen3_dflash.py | 69 ++---- vllm/model_executor/models/qwen3_next.py | 18 +- vllm/model_executor/models/qwen3_next_mtp.py | 19 +- vllm/model_executor/models/sarvam.py | 113 ++------- vllm/model_executor/models/step3p5.py | 16 +- vllm/model_executor/models/step3p5_mtp.py | 4 +- vllm/model_executor/models/telechat2.py | 92 +++---- vllm/model_executor/models/utils.py | 111 +++++++++ 44 files changed, 636 insertions(+), 2073 deletions(-) diff --git a/vllm/model_executor/models/AXK1.py b/vllm/model_executor/models/AXK1.py index a465c6b5632..7818200dff5 100644 --- a/vllm/model_executor/models/AXK1.py +++ b/vllm/model_executor/models/AXK1.py @@ -79,6 +79,7 @@ from .interfaces import MixtureOfExperts, SupportsEagle, SupportsLoRA, SupportsP from .utils import ( AutoWeightsLoader, PPMissingLayer, + get_spec_layer_idx_from_weight_name, is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, @@ -1143,16 +1144,3 @@ class AXK1ForCausalLM( def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: loader = AutoWeightsLoader(self) return loader.load_weights(weights) - - -def get_spec_layer_idx_from_weight_name( - config: AXK1Config, weight_name: str -) -> int | None: - if config.num_nextn_predict_layers and config.num_nextn_predict_layers > 0: - layer_idx = config.num_hidden_layers - for i in range(config.num_nextn_predict_layers): - if weight_name.startswith( - f"model.layers.{layer_idx + i}." - ) or weight_name.startswith(f"layers.{layer_idx + i}."): - return layer_idx + i - return None diff --git a/vllm/model_executor/models/aria.py b/vllm/model_executor/models/aria.py index 6b723883423..5117f541cfd 100644 --- a/vllm/model_executor/models/aria.py +++ b/vllm/model_executor/models/aria.py @@ -11,21 +11,13 @@ from transformers.models.aria.processing_aria import AriaProcessor from vllm.config import VllmConfig from vllm.config.multimodal import BaseDummyOptions -from vllm.distributed import get_tensor_model_parallel_rank from vllm.inputs import MultiModalDataDict from vllm.model_executor.layers.activation import get_act_fn -from vllm.model_executor.layers.fused_moe import ( - FusedMoE, - RoutedExperts, -) +from vllm.model_executor.layers.fused_moe import FusedMoE from vllm.model_executor.layers.linear import ColumnParallelLinear, RowParallelLinear from vllm.model_executor.layers.logits_processor import LogitsProcessor from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.layers.vocab_parallel_embedding import ParallelLMHead -from vllm.model_executor.model_loader.weight_utils import ( - default_weight_loader, - maybe_remap_kv_scale_name, -) from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import ( MultiModalFieldConfig, @@ -51,7 +43,6 @@ from .llama import LlamaDecoderLayer, LlamaMLP, LlamaModel from .utils import ( AutoWeightsLoader, WeightsMapper, - is_pp_missing_parameter, maybe_prefix, ) @@ -94,34 +85,18 @@ class AriaVisionTransformer(Idefics3VisionTransformer, SupportsQuant): # Identity layer self.post_layernorm = nn.Identity() - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ] - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - for name, loaded_weight in weights: - # NOTE: post_layernorm is not used in Aria - if "post_layernorm" in name: - continue + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_stacked={ + ".q_proj": (".qkv_proj", "q"), + ".k_proj": (".qkv_proj", "k"), + ".v_proj": (".qkv_proj", "v"), + } + ) - for param_name, weight_name, shard_id in stacked_params_mapping: - if weight_name not in name: - continue - name = name.replace(weight_name, param_name) - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - break - else: - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + # NOTE: post_layernorm is not used in Aria. + loader = AutoWeightsLoader(self, skip_substrs=["post_layernorm"]) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) class AriaProjectorMLP(nn.Module): @@ -217,39 +192,6 @@ class AriaProjector(nn.Module): return out -class AriaRoutedExperts(RoutedExperts): - def weight_loader( - self, param: nn.Parameter, loaded_weight: torch.Tensor, shard_id: str - ) -> None: - # Override the weight_loader to handle the expert weights in the Aria - # model, which are already packed with experts, and merge the gate and - # up weights for each expert. - # Note: Loading expert weights with quantization is not supported - tp_rank = get_tensor_model_parallel_rank() - tp_size = self.moe_config.tp_size - if shard_id == "w13": - # the shape of loaded_weight is - # (num_experts, hidden_size, 2 * moe_intermediate_size) - if tp_size > 1: - up, gate = loaded_weight.chunk(2, dim=-1) - up_current_rank = up.chunk(tp_size, dim=-1)[tp_rank] - gate_current_rank = gate.chunk(tp_size, dim=-1)[tp_rank] - up_and_gate = torch.cat( - [up_current_rank, gate_current_rank], dim=-1 - ).transpose(1, 2) - param.data.copy_(up_and_gate) - else: - param.data.copy_(loaded_weight.transpose(1, 2)) - elif shard_id == "w2": - # the shape of loaded_weight is - # (num_experts, moe_intermediate_size, hidden_size) - if tp_size > 1: - down_current_rank = loaded_weight.chunk(tp_size, dim=1)[tp_rank] - param.data.copy_(down_current_rank.transpose(1, 2)) - else: - param.data.copy_(loaded_weight.transpose(1, 2)) - - class AriaTextMoELayer(nn.Module): """ Mixture of Experts (MoE) Layer for the AriaMoE model. @@ -288,7 +230,6 @@ class AriaTextMoELayer(nn.Module): intermediate_size=config.intermediate_size, quant_config=quant_config, prefix=f"{prefix}.experts", - routed_experts_cls=AriaRoutedExperts, ) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: @@ -335,70 +276,23 @@ class AriaTextModel(LlamaModel, SupportsQuant): packed_modules_mapping = { "qkv_proj": ["q_proj", "k_proj", "v_proj"], "gate_up_proj": ["gate_proj", "up_proj"], - "experts.routed_experts.w13_weight": ["experts.fc1.weight"], - "experts.routed_experts.w2_weight": ["experts.fc2.weight"], } + # Aria packs all experts into single (transposed) fc1/fc2 tensors, which is + # exactly the pre-fused checkpoint layout FusedMoE self-loads once fc1/fc2 + # are renamed to the fused gate_up_proj/down_proj names. + hf_to_vllm_mapper = LlamaModel.hf_to_vllm_mapper | WeightsMapper( + orig_to_new_substr={ + "experts.fc1": "experts.gate_up_proj", + "experts.fc2": "experts.down_proj", + } + ) + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): super().__init__( vllm_config=vllm_config, prefix=prefix, layer_type=AriaTextDecoderLayer ) - # Adapted from LlamaModel.load_weights with the modification of adding - # the expert weights mapping to `stacked_params_mapping` - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - (".qkv_proj", ".q_proj", "q"), - (".qkv_proj", ".k_proj", "k"), - (".qkv_proj", ".v_proj", "v"), - (".gate_up_proj", ".gate_proj", 0), - (".gate_up_proj", ".up_proj", 1), - ("experts.routed_experts.w13_weight", "experts.fc1.weight", "w13"), - ("experts.routed_experts.w2_weight", "experts.fc2.weight", "w2"), - ] - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - for name, loaded_weight in weights: - if "rotary_emb.inv_freq" in name: - continue - if "rotary_emb.cos_cached" in name or "rotary_emb.sin_cached" in name: - # Models trained using ColossalAI may include these tensors in - # the checkpoint. Skip them. - continue - for param_name, weight_name, shard_id in stacked_params_mapping: - if weight_name not in name: - continue - name = name.replace(weight_name, param_name) - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue - - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - break - else: - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue - # Remapping the name of FP8 kv-scale. - name = maybe_remap_kv_scale_name(name, params_dict) - if name is None: - continue - - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params - class AriaProcessingInfo(BaseProcessingInfo): def get_hf_config(self): diff --git a/vllm/model_executor/models/bailing_moe.py b/vllm/model_executor/models/bailing_moe.py index a45d0ca81de..642d07ee659 100644 --- a/vllm/model_executor/models/bailing_moe.py +++ b/vllm/model_executor/models/bailing_moe.py @@ -41,10 +41,7 @@ from vllm.distributed import ( ) from vllm.model_executor.layers.activation import SiluAndMul from vllm.model_executor.layers.attention import Attention -from vllm.model_executor.layers.fused_moe import ( - FusedMoE, - fused_moe_make_expert_params_mapping, -) +from vllm.model_executor.layers.fused_moe import FusedMoE from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ( MergedColumnParallelLinear, @@ -58,14 +55,13 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.sequence import IntermediateTensors from .interfaces import SupportsLoRA, SupportsPP from .utils import ( AutoWeightsLoader, PPMissingLayer, - is_pp_missing_parameter, + WeightsMapper, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -379,6 +375,16 @@ class BailingMoeBlock(nn.Module): @support_torch_compile class BailingMoeModel(nn.Module): + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_stacked={ + # .experts.gate_up_proj must be handled by MoERunner.load_weights for EP + ".mlp.gate_proj": (".mlp.gate_up_proj", 0), + ".mlp.up_proj": (".mlp.gate_up_proj", 1), + ".shared_experts.gate_proj": (".shared_experts.gate_up_proj", 0), + ".shared_experts.up_proj": (".shared_experts.gate_up_proj", 1), + } + ) + def __init__( self, *, @@ -468,89 +474,9 @@ class BailingMoeModel(nn.Module): hidden_states, _ = self.norm(hidden_states, residual) return hidden_states - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - return fused_moe_make_expert_params_mapping( - self, - ckpt_gate_proj_name="gate_proj", - ckpt_down_proj_name="down_proj", - ckpt_up_proj_name="up_proj", - num_experts=self.config.num_experts, - ) - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] - - params_dict = dict(self.named_parameters(remove_duplicate=False)) - loaded_params: set[str] = set() - expert_params_mapping = self.get_expert_mapping() - for name, loaded_weight in weights: - if ( - hasattr(self.config, "norm_head") - and self.config.norm_head - and "lm_head.weight" in name - ): - loaded_weight = F.normalize(loaded_weight, dim=0, p=2, eps=1e-7) - - for param_name, weight_name, shard_id in stacked_params_mapping: - if weight_name not in name: - continue - if "mlp.experts" in name: - continue - name = name.replace(weight_name, param_name) - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue - if name not in params_dict: - continue - - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - break - else: - for mapping in expert_params_mapping: - param_name, weight_name, expert_id, shard_id = mapping - if weight_name not in name: - continue - name = name.replace(weight_name, param_name) - - if is_pp_missing_parameter(name, self): - continue - if name not in params_dict: - continue - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader( - param, - loaded_weight, - name, - shard_id=shard_id, - expert_id=expert_id, - ) - break - else: - if name.endswith(".bias") and name not in params_dict: - continue - if name not in params_dict: - continue - - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = getattr( - param, "weight_loader", default_weight_loader - ) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + loader = AutoWeightsLoader(self) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) class BailingMoeForCausalLM(nn.Module, SupportsPP, SupportsLoRA): @@ -622,15 +548,21 @@ class BailingMoeForCausalLM(nn.Module, SupportsPP, SupportsLoRA): logits = self.logits_processor(self.lm_head, hidden_states) return logits + def _normalize_lm_head( + self, weights: Iterable[tuple[str, torch.Tensor]] + ) -> Iterable[tuple[str, torch.Tensor]]: + norm_head = getattr(self.config, "norm_head", False) + for name, loaded_weight in weights: + if norm_head and "lm_head.weight" in name: + loaded_weight = F.normalize(loaded_weight, dim=0, p=2, eps=1e-7) + yield name, loaded_weight + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: loader = AutoWeightsLoader( self, skip_prefixes=(["lm_head."] if self.tie_word_embeddings else None), ) - return loader.load_weights(weights) - - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - return self.model.get_expert_mapping() + return loader.load_weights(self._normalize_lm_head(weights)) class BailingMoeV2ForCausalLM(BailingMoeForCausalLM): diff --git a/vllm/model_executor/models/bloom.py b/vllm/model_executor/models/bloom.py index 233028a905f..cdcf82f385f 100644 --- a/vllm/model_executor/models/bloom.py +++ b/vllm/model_executor/models/bloom.py @@ -47,13 +47,11 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.sequence import IntermediateTensors from .interfaces import SupportsPP, SupportsQuant from .utils import ( AutoWeightsLoader, - is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -297,36 +295,23 @@ class BloomModel(nn.Module): hidden_states = self.ln_f(hidden_states) return hidden_states - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - params_dict = dict(self.named_parameters(remove_duplicate=False)) - loaded_params: set[str] = set() + def _repack_qkv( + self, weights: Iterable[tuple[str, torch.Tensor]] + ) -> Iterable[tuple[str, torch.Tensor]]: + # BLOOM's fused QKV is laid out as (num_heads * 3 * head_size) on its + # output dim (0), while vLLM expects (3 * num_heads * head_size). + num_heads = self.config.num_attention_heads for name, loaded_weight in weights: - if is_pp_missing_parameter(name, self): - continue - param = params_dict[name] - if "query_key_value" in name: - # NOTE: BLOOM's fused QKV's output_dim has the shape of - # (num_heads * 3 * head_size), while the - # required shape is (3 * num_heads * head_size). - # Thus, we need weight conversion. - output_dim = getattr(param, "output_dim", None) - num_heads = self.config.num_attention_heads - if output_dim is not None: - loaded_weight_shape = loaded_weight.shape - loaded_weight = loaded_weight.view( - loaded_weight_shape[:output_dim] - + (num_heads, 3, -1) - + loaded_weight_shape[output_dim + 1 :] - ) - loaded_weight = loaded_weight.transpose(output_dim, output_dim + 1) - loaded_weight = loaded_weight.reshape(loaded_weight_shape) + shape = loaded_weight.shape + loaded_weight = loaded_weight.view((num_heads, 3, -1) + shape[1:]) + loaded_weight = loaded_weight.transpose(0, 1) + loaded_weight = loaded_weight.reshape(shape) + yield name, loaded_weight - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - - return loaded_params + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + loader = AutoWeightsLoader(self) + return loader.load_weights(self._repack_qkv(weights)) class BloomForCausalLM(nn.Module, SupportsPP, SupportsQuant): diff --git a/vllm/model_executor/models/cohere2_moe.py b/vllm/model_executor/models/cohere2_moe.py index 80dc6802060..b993247ec6f 100644 --- a/vllm/model_executor/models/cohere2_moe.py +++ b/vllm/model_executor/models/cohere2_moe.py @@ -16,9 +16,7 @@ from vllm.distributed import ( ) from vllm.model_executor.layers.activation import SiluAndMul from vllm.model_executor.layers.attention import Attention -from vllm.model_executor.layers.fused_moe import ( - FusedMoE, -) +from vllm.model_executor.layers.fused_moe import FusedMoE from vllm.model_executor.layers.linear import ( MergedColumnParallelLinear, QKVParallelLinear, diff --git a/vllm/model_executor/models/deepencoder.py b/vllm/model_executor/models/deepencoder.py index 68c101460d5..fffd9382dd7 100644 --- a/vllm/model_executor/models/deepencoder.py +++ b/vllm/model_executor/models/deepencoder.py @@ -22,9 +22,9 @@ from vllm.model_executor.custom_op import PluggableLayer from vllm.model_executor.layers.attention import MMEncoderAttention from vllm.model_executor.layers.conv import Conv2dLayer from vllm.model_executor.layers.quantization import QuantizationConfig -from vllm.model_executor.model_loader.weight_utils import default_weight_loader from .clip import CLIPEncoder, CLIPVisionEmbeddings +from .utils import AutoWeightsLoader class MLPBlock(nn.Module): @@ -671,12 +671,5 @@ class DeepCLIPVisionTransformer(nn.Module): return encoder_outputs def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - - for name, loaded_weight in weights: - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + loader = AutoWeightsLoader(self) + return loader.load_weights(weights) diff --git a/vllm/model_executor/models/deepseek_mtp.py b/vllm/model_executor/models/deepseek_mtp.py index a5dd6ae8f9f..3a0c21fe7d2 100644 --- a/vllm/model_executor/models/deepseek_mtp.py +++ b/vllm/model_executor/models/deepseek_mtp.py @@ -33,9 +33,12 @@ from .deepseek_v2 import ( DeepseekV2MixtureOfExperts, DeepseekV2MoE, _try_load_fp8_indexer_wk, - get_spec_layer_idx_from_weight_name, ) -from .utils import get_pp_missing_layer_names, maybe_prefix +from .utils import ( + get_pp_missing_layer_names, + get_spec_layer_idx_from_weight_name, + maybe_prefix, +) def _restore_full_token_layout_if_needed( diff --git a/vllm/model_executor/models/deepseek_v2.py b/vllm/model_executor/models/deepseek_v2.py index ac187151a0a..daf4ea3de9f 100644 --- a/vllm/model_executor/models/deepseek_v2.py +++ b/vllm/model_executor/models/deepseek_v2.py @@ -108,6 +108,7 @@ from .interfaces import ( from .utils import ( PPMissingLayer, get_pp_missing_layer_names, + get_spec_layer_idx_from_weight_name, is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, @@ -1918,21 +1919,3 @@ class DeepseekV3ForCausalLM(DeepseekV2ForCausalLM): class GlmMoeDsaForCausalLM(DeepseekV2ForCausalLM): pass - - -# Compatibility with -# https://huggingface.co/deepseek-ai/DeepSeek-V3-Base/blob/main/configuration_deepseek.py -def get_spec_layer_idx_from_weight_name( - config: DeepseekV2Config | DeepseekV3Config, weight_name: str -) -> int | None: - if ( - hasattr(config, "num_nextn_predict_layers") - and config.num_nextn_predict_layers > 0 - ): - layer_idx = config.num_hidden_layers - for i in range(config.num_nextn_predict_layers): - if weight_name.startswith( - f"model.layers.{layer_idx + i}." - ) or weight_name.startswith(f"layers.{layer_idx + i}."): - return layer_idx + i - return None diff --git a/vllm/model_executor/models/ernie45_moe.py b/vllm/model_executor/models/ernie45_moe.py index c2d9f92a666..fea390ca21c 100644 --- a/vllm/model_executor/models/ernie45_moe.py +++ b/vllm/model_executor/models/ernie45_moe.py @@ -23,8 +23,7 @@ # limitations under the License. """Inference-only ErineMoE model compatible with HuggingFace weights.""" -import typing -from collections.abc import Callable, Iterable +from collections.abc import Iterable from itertools import islice from typing import Any @@ -42,11 +41,7 @@ from vllm.distributed import ( from vllm.logger import init_logger from vllm.model_executor.layers.activation import SiluAndMul from vllm.model_executor.layers.attention import Attention -from vllm.model_executor.layers.fused_moe import ( - FusedMoE, - MoERunner, - fused_moe_make_expert_params_mapping, -) +from vllm.model_executor.layers.fused_moe import FusedMoE, MoERunner from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ( MergedColumnParallelLinear, @@ -61,10 +56,6 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import ( - default_weight_loader, - maybe_remap_kv_scale_name, -) from vllm.sequence import IntermediateTensors from vllm.transformers_utils.config import set_default_rope_theta @@ -72,8 +63,8 @@ from .interfaces import MixtureOfExperts, SupportsLoRA, SupportsPP from .utils import ( AutoWeightsLoader, PPMissingLayer, + WeightsMapper, extract_layer_index, - is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -407,6 +398,19 @@ class Ernie4_5_MoeDecoderLayer(nn.Module): @support_torch_compile class Ernie4_5_MoeModel(nn.Module): + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_stacked={ + ".q_proj": (".qkv_proj", "q"), + ".k_proj": (".qkv_proj", "k"), + ".v_proj": (".qkv_proj", "v"), + # .experts.gate_up_proj must be handled by MoERunner.load_weights for EP + ".mlp.gate_proj": (".mlp.gate_up_proj", 0), + ".mlp.up_proj": (".mlp.gate_up_proj", 1), + ".shared_experts.gate_proj": (".shared_experts.gate_up_proj", 0), + ".shared_experts.up_proj": (".shared_experts.gate_up_proj", 1), + } + ) + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): super().__init__() @@ -486,132 +490,26 @@ class Ernie4_5_MoeModel(nn.Module): return hidden_states - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - # Params for weights, fp8 weight scales, fp8 activation scales - # (param_name, weight_name, expert_id, shard_id) - return fused_moe_make_expert_params_mapping( - self, - ckpt_gate_proj_name="gate_proj", - ckpt_down_proj_name="down_proj", - ckpt_up_proj_name="up_proj", - num_experts=self.config.moe_num_experts, - num_redundant_experts=self.num_redundant_experts, - ) - - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] - - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - expert_params_mapping = self.get_expert_mapping() + def _preprocess( + self, weights: Iterable[tuple[str, torch.Tensor]] + ) -> Iterable[tuple[str, torch.Tensor]]: for name, loaded_weight in weights: - if self.config.tie_word_embeddings and name.endswith("lm_head.weight"): - continue - # MTP will be supported soon. - if "mtp" in name: - continue - + # moe_statics.e_score_correction_bias is stored with a leading + # singleton dim and under a different module name. if "e_score_correction_bias" in name: name = name.replace("moe_statics", "gate") loaded_weight = loaded_weight.squeeze(0) + yield name, loaded_weight - for param_name, weight_name, shard_id in stacked_params_mapping: - # Skip non-stacked layers and experts (experts handled below). - if weight_name not in name: - continue - - if ("mlp.experts." in name) and name not in params_dict: - continue - name = name.replace(weight_name, param_name) - # Skip loading extra bias for GPTQ models. - if ( - name.endswith(".bias") or name.endswith("_bias") - ) and name not in params_dict: - continue - # Skip layers on other devices. - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - break - else: - is_expert_weight = False - for mapping in expert_params_mapping: - param_name, weight_name, expert_id, shard_id = mapping - - if weight_name not in name: - continue - - # Anyway, this is an expert weight and should not be - # attempted to load as other weights later - is_expert_weight = True - - # Do not modify `name` since the loop may continue here - # Instead, create a new variable - name_mapped = name.replace(weight_name, param_name) - # Skip layers on other devices. - if is_pp_missing_parameter(name_mapped, self): - continue - - # Skip loading extra bias for GPTQ models. - if ( - name_mapped.endswith(".bias") or name_mapped.endswith("_bias") - ) and name_mapped not in params_dict: - continue - param = params_dict[name_mapped] - # We should ask the weight loader to return success or not - # here since otherwise we may skip experts with other - # available replicas. - weight_loader = typing.cast( - Callable[..., bool], param.weight_loader - ) - success = weight_loader( - param, - loaded_weight, - name_mapped, - shard_id=shard_id, - expert_id=expert_id, - return_success=True, - ) - if success: - name = name_mapped - break - else: - if is_expert_weight: - # We've checked that this is an expert weight - # However it's not mapped locally to this rank - # So we simply skip it - continue - - # Skip loading extra bias for GPTQ models. - if ( - name.endswith(".bias") or name.endswith("_bias") - ) and name not in params_dict: - continue - # Skip layers on other devices. - if is_pp_missing_parameter(name, self): - continue - # Remapping the name of FP8 kv-scale. - name = maybe_remap_kv_scale_name(name, params_dict) - if name is None: - continue - - param = params_dict[name] - weight_loader = getattr( - param, "weight_loader", default_weight_loader - ) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + loader = AutoWeightsLoader( + self, + skip_substrs=["mtp"], + ignore_unexpected_suffixes=[".bias", "_bias"], + ) + return loader.load_weights( + self._preprocess(weights), mapper=self.hf_to_vllm_mapper + ) class Ernie4_5_MoeForCausalLM(nn.Module, SupportsPP, SupportsLoRA, MixtureOfExperts): @@ -741,6 +639,3 @@ class Ernie4_5_MoeForCausalLM(nn.Module, SupportsPP, SupportsLoRA, MixtureOfExpe skip_prefixes=(["lm_head."] if self.config.tie_word_embeddings else None), ) return loader.load_weights(weights) - - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - return self.model.get_expert_mapping() diff --git a/vllm/model_executor/models/funaudiochat.py b/vllm/model_executor/models/funaudiochat.py index 7e7cdcd822c..72b12e26b5c 100644 --- a/vllm/model_executor/models/funaudiochat.py +++ b/vllm/model_executor/models/funaudiochat.py @@ -30,7 +30,6 @@ from vllm.config.multimodal import BaseDummyOptions from vllm.inputs import MultiModalDataDict from vllm.model_executor.layers.attention.mm_encoder_attention import MMEncoderAttention from vllm.model_executor.layers.linear import QKVParallelLinear, RowParallelLinear -from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import ( MultiModalFieldConfig, @@ -53,7 +52,12 @@ from vllm.sequence import IntermediateTensors from vllm.utils.import_utils import _has_module from .interfaces import MultiModalEmbeddings, SupportsMultiModal, SupportsPP -from .utils import AutoWeightsLoader, init_vllm_registered_model, maybe_prefix +from .utils import ( + AutoWeightsLoader, + WeightsMapper, + init_vllm_registered_model, + maybe_prefix, +) class _SinusoidsPositionEmbedding(nn.Module): @@ -79,6 +83,14 @@ class _SinusoidsPositionEmbedding(nn.Module): class FunAudioChatAudioAttention(nn.Module): """Multi-headed attention used inside the continuous audio tower.""" + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_stacked={ + ".q_proj": (".qkv_proj", "q"), + ".k_proj": (".qkv_proj", "k"), + ".v_proj": (".qkv_proj", "v"), + } + ) + def __init__(self, config: Any): super().__init__() self.embed_dim = int(config.d_model) @@ -123,42 +135,15 @@ class FunAudioChatAudioAttention(nn.Module): bias=True, ) - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ] - - params_dict = dict(self.named_parameters()) with torch.no_grad(): if self.qkv_proj.bias is not None: # HF FunAudioChat uses bias=False for k_proj. Ensure the missing # shard starts as zeros, while allowing q/v shards to load. self.qkv_proj.bias.zero_() - loaded_params: set[str] = set() - for name, loaded_weight in weights: - for param_name, shard_name, shard_id in stacked_params_mapping: - if shard_name not in name: - continue - name = name.replace(shard_name, param_name) - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight, shard_id) - break - else: - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue - - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - - loaded_params.add(name) - - return loaded_params + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + loader = AutoWeightsLoader(self) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) def forward( self, diff --git a/vllm/model_executor/models/gemma3n.py b/vllm/model_executor/models/gemma3n.py index ad8b21d86b4..4b06fd418f3 100644 --- a/vllm/model_executor/models/gemma3n.py +++ b/vllm/model_executor/models/gemma3n.py @@ -44,18 +44,14 @@ from vllm.model_executor.layers.logits_processor import LogitsProcessor from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.layers.rotary_embedding import get_rope from vllm.model_executor.layers.vocab_parallel_embedding import VocabParallelEmbedding -from vllm.model_executor.model_loader.weight_utils import ( - default_weight_loader, - maybe_remap_kv_scale_name, -) from vllm.sequence import IntermediateTensors from vllm.v1.attention.backends.utils import KVSharingFastPrefillMetadata from .interfaces import SupportsQuant from .utils import ( AutoWeightsLoader, + WeightsMapper, extract_layer_index, - is_pp_missing_parameter, make_layers, maybe_prefix, ) @@ -788,6 +784,25 @@ class Gemma3nCrossDecoder(nn.Module): enable_if=lambda vllm_config: not vllm_config.cache_config.kv_sharing_fast_prefill ) class Gemma3nTextModel(nn.Module, SupportsQuant): + # Decoder layers, altup_unembed_projections and norm live on the text + # model; every other submodule lives under self_decoder. + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_prefix={ + "embed_tokens.": "self_decoder.embed_tokens.", + "embed_tokens_per_layer.": "self_decoder.embed_tokens_per_layer.", + "per_layer_model_projection.": "self_decoder.per_layer_model_projection.", + "per_layer_projection_norm.": "self_decoder.per_layer_projection_norm.", + "altup_projections.": "self_decoder.altup_projections.", + }, + orig_to_new_stacked={ + ".q_proj": (".qkv_proj", "q"), + ".k_proj": (".qkv_proj", "k"), + ".v_proj": (".qkv_proj", "v"), + ".gate_proj": (".gate_up_proj", 0), + ".up_proj": (".gate_up_proj", 1), + }, + ) + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): super().__init__() config = vllm_config.model_config.hf_config @@ -1036,58 +1051,8 @@ class Gemma3nTextModel(nn.Module, SupportsQuant): return self.norm(hidden_states) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - for name, loaded_weight in weights: - # decoder layer weights, altup_unembed_projections and rmsnorm - # are initialized in text model, others are in self decoder - if ( - not name.startswith("layers") - and not name.startswith("altup_unembed_projections") - and not name.startswith("norm") - ): - name = f"self_decoder.{name}" - - for param_name, shard_name, shard_id in stacked_params_mapping: - if shard_name not in name: - continue - # Avoid spurious match with ".up_proj". - if "altup_projections" in name: - continue - name = name.replace(shard_name, param_name) - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue - if is_pp_missing_parameter(name, self): - continue - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - break - else: - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue - # Remapping the name of FP8 kv-scale. - name = maybe_remap_kv_scale_name(name, params_dict) - if name is None: - continue - if is_pp_missing_parameter(name, self): - continue - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - - return loaded_params + loader = AutoWeightsLoader(self) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) class Gemma3nForCausalLM(nn.Module): diff --git a/vllm/model_executor/models/glm4.py b/vllm/model_executor/models/glm4.py index 3a25f90ad2a..a1fb94fb26f 100644 --- a/vllm/model_executor/models/glm4.py +++ b/vllm/model_executor/models/glm4.py @@ -303,16 +303,3 @@ class Glm4ForCausalLM(nn.Module, SupportsLoRA, SupportsPP): ] loader = AutoWeightsLoader(self, skip_prefixes=skip_prefixes) return loader.load_weights(weights) - - -def get_spec_layer_idx_from_weight_name( - config: Glm4Config, weight_name: str -) -> int | None: - if hasattr(config, "num_nextn_predict_layers") and ( - config.num_nextn_predict_layers > 0 - ): - layer_idx = config.num_hidden_layers - for i in range(config.num_nextn_predict_layers): - if f"layers.{layer_idx + i}." in weight_name: - return layer_idx + i - return None diff --git a/vllm/model_executor/models/glm4_moe.py b/vllm/model_executor/models/glm4_moe.py index e3f94c673f4..abb9970c403 100644 --- a/vllm/model_executor/models/glm4_moe.py +++ b/vllm/model_executor/models/glm4_moe.py @@ -24,8 +24,7 @@ """Inference-only GLM-4.5, GLM-4.6, GLM-4.7 model compatible with HuggingFace weights.""" -import typing -from collections.abc import Callable, Iterable +from collections.abc import Iterable from itertools import islice import torch @@ -43,10 +42,7 @@ from vllm.distributed import ( from vllm.logger import init_logger from vllm.model_executor.layers.activation import SiluAndMul from vllm.model_executor.layers.attention import Attention -from vllm.model_executor.layers.fused_moe import ( - FusedMoE, - fused_moe_make_expert_params_mapping, -) +from vllm.model_executor.layers.fused_moe import FusedMoE from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ( MergedColumnParallelLinear, @@ -60,20 +56,18 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import ( - default_weight_loader, - maybe_remap_kv_scale_name, -) from vllm.sequence import IntermediateTensors from .interfaces import MixtureOfExperts, SupportsLoRA, SupportsPP from .utils import ( AutoWeightsLoader, PPMissingLayer, - is_pp_missing_parameter, + WeightsMapper, make_empty_intermediate_tensors_factory, make_layers, + maybe_fuse_shared_experts, maybe_prefix, + skip_spec_layers, ) logger = init_logger(__name__) @@ -410,6 +404,19 @@ class Glm4MoeDecoderLayer(nn.Module): } ) class Glm4MoeModel(nn.Module): + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_stacked={ + ".q_proj": (".qkv_proj", "q"), + ".k_proj": (".qkv_proj", "k"), + ".v_proj": (".qkv_proj", "v"), + # .experts.gate_up_proj must be handled by MoERunner.load_weights for EP + ".mlp.gate_proj": (".mlp.gate_up_proj", 0), + ".mlp.up_proj": (".mlp.gate_up_proj", 1), + ".shared_experts.gate_proj": (".shared_experts.gate_up_proj", 0), + ".shared_experts.up_proj": (".shared_experts.gate_up_proj", 1), + } + ) + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): super().__init__() @@ -480,189 +487,14 @@ class Glm4MoeModel(nn.Module): hidden_states, _ = self.norm(hidden_states, residual) return hidden_states - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - # Params for weights, fp8 weight scales, fp8 activation scales - # (param_name, weight_name, expert_id, shard_id) - # FSE widens the mapping by n_shared_experts slots; see deepseek_v2.py. - num_experts = self.config.n_routed_experts - if ( - rocm_aiter_ops.is_fusion_moe_shared_experts_enabled() - and self.config.n_shared_experts - ): - num_experts += self.config.n_shared_experts - return fused_moe_make_expert_params_mapping( - self, - ckpt_gate_proj_name="gate_proj", - ckpt_down_proj_name="down_proj", - ckpt_up_proj_name="up_proj", - num_experts=num_experts, - ) - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - rocm_aiter_moe_shared_expert_enabled = ( - rocm_aiter_ops.is_fusion_moe_shared_experts_enabled() + weights = maybe_fuse_shared_experts( + skip_spec_layers(weights, self.config), + n_routed_experts=self.config.n_routed_experts, + n_shared_experts=self.config.n_shared_experts or 1, ) - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] - - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - expert_params_mapping = self.get_expert_mapping() - for name, loaded_weight in weights: - spec_layer = get_spec_layer_idx_from_weight_name(self.config, name) - if spec_layer is not None: - continue - - is_fusion_moe_shared_experts_layer = ( - rocm_aiter_moe_shared_expert_enabled and ("mlp.shared_experts" in name) - ) - - for param_name, weight_name, shard_id in stacked_params_mapping: - # Skip non-stacked layers and experts (experts handled below). - if weight_name not in name: - continue - # We have mlp.experts[0].gate_proj in the checkpoint. - # Since we handle the experts below in expert_params_mapping, - # we need to skip here BEFORE we update the name, otherwise - # name will be updated to mlp.experts[0].gate_up_proj, which - # will then be updated below in expert_params_mapping - # for mlp.experts[0].gate_gate_up_proj, which breaks load. - if ("mlp.experts." in name) and name not in params_dict: - continue - if is_fusion_moe_shared_experts_layer: - continue - - name = name.replace(weight_name, param_name) - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue - - name = maybe_remap_kv_scale_name(name, params_dict) - if name is None: - continue - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - if weight_loader == default_weight_loader: - weight_loader(param, loaded_weight) - else: - weight_loader(param, loaded_weight, shard_id) - break - else: - is_expert_weight = False - - # FSE: split a widened mlp.shared_experts tensor into - # n_shared_experts chunks; see deepseek_v2.py for details. - num_chunks = 1 - split_dim = 0 - chunk_size = 0 - if is_fusion_moe_shared_experts_layer: - num_chunks = getattr(self.config, "n_shared_experts", 1) or 1 - split_dim = ( - 1 - if ("down_proj.weight" in name and loaded_weight.ndim > 1) - else 0 - ) - total = loaded_weight.shape[split_dim] - if total % num_chunks != 0: - raise ValueError( - f"FSE shared-expert weight {name} has dim " - f"{total} along axis {split_dim} which is not " - f"divisible by n_shared_experts={num_chunks}." - ) - chunk_size = total // num_chunks - - for j in range(num_chunks): - chunk_name = name - weight_to_load = loaded_weight - - if is_fusion_moe_shared_experts_layer: - chunk_slice = slice(j * chunk_size, (j + 1) * chunk_size) - if loaded_weight.ndim == 1: - weight_to_load = loaded_weight[chunk_slice] - elif split_dim == 0: - weight_to_load = loaded_weight[chunk_slice, :] - else: - weight_to_load = loaded_weight[:, chunk_slice] - # Synthesize an expert-style name for expert mapping. - chunk_name = name.replace( - "mlp.shared_experts", - f"mlp.experts.{self.config.n_routed_experts + j}", - ) - - for mapping in expert_params_mapping: - param_name, weight_name, expert_id, shard_id = mapping - if weight_name not in chunk_name: - continue - - # Anyway, this is an expert weight and should not be - # attempted to load as other weights later - is_expert_weight = True - - # Do not modify `name` since the loop may continue here - # Instead, create a new variable - name_mapped = chunk_name.replace(weight_name, param_name) - - if is_pp_missing_parameter(name_mapped, self): - continue - - param = params_dict[name_mapped] - # We should ask the weight loader to return success - # or not here since otherwise we may skip experts - # with other available replicas. - weight_loader = typing.cast( - Callable[..., bool], param.weight_loader - ) - success = weight_loader( - param, - weight_to_load, - name_mapped, - shard_id=shard_id, - expert_id=expert_id, - return_success=True, - ) - if success: - if not is_fusion_moe_shared_experts_layer: - name = name_mapped - else: - loaded_params.add(name_mapped) - break - else: - if is_expert_weight: - # We've checked that this is an expert weight - # However it's not mapped locally to this rank - # So we simply skip it - continue - - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue - - # Remapping the name of FP8 kv-scale. - name = maybe_remap_kv_scale_name(name, params_dict) - if name is None: - continue - - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = getattr( - param, "weight_loader", default_weight_loader - ) - weight_loader(param, loaded_weight) - if name is not None and not is_fusion_moe_shared_experts_layer: - loaded_params.add(name) - - return loaded_params + loader = AutoWeightsLoader(self) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) class Glm4MixtureOfExperts(MixtureOfExperts): @@ -776,19 +608,3 @@ class Glm4MoeForCausalLM(nn.Module, SupportsPP, SupportsLoRA, Glm4MixtureOfExper def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: loader = AutoWeightsLoader(self) return loader.load_weights(weights) - - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - return self.model.get_expert_mapping() - - -def get_spec_layer_idx_from_weight_name( - config: Glm4MoeConfig, weight_name: str -) -> int | None: - if hasattr(config, "num_nextn_predict_layers") and ( - config.num_nextn_predict_layers > 0 - ): - layer_idx = config.num_hidden_layers - for i in range(config.num_nextn_predict_layers): - if f"layers.{layer_idx + i}." in weight_name: - return layer_idx + i - return None diff --git a/vllm/model_executor/models/glm4_moe_lite.py b/vllm/model_executor/models/glm4_moe_lite.py index 432fa5e6fa0..7f63130d883 100644 --- a/vllm/model_executor/models/glm4_moe_lite.py +++ b/vllm/model_executor/models/glm4_moe_lite.py @@ -70,6 +70,7 @@ from .interfaces import SupportsLoRA, SupportsPP from .utils import ( AutoWeightsLoader, PPMissingLayer, + get_spec_layer_idx_from_weight_name, is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, @@ -628,16 +629,3 @@ class Glm4MoeLiteForCausalLM( def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: loader = AutoWeightsLoader(self) return loader.load_weights(weights) - - -def get_spec_layer_idx_from_weight_name( - config: "Glm4MoeLiteConfig", weight_name: str -) -> int | None: - if hasattr(config, "num_nextn_predict_layers") and ( - config.num_nextn_predict_layers > 0 - ): - layer_idx = config.num_hidden_layers - for i in range(config.num_nextn_predict_layers): - if f"layers.{layer_idx + i}." in weight_name: - return layer_idx + i - return None diff --git a/vllm/model_executor/models/glm4_moe_lite_mtp.py b/vllm/model_executor/models/glm4_moe_lite_mtp.py index 222705c14ee..bbede18c40a 100644 --- a/vllm/model_executor/models/glm4_moe_lite_mtp.py +++ b/vllm/model_executor/models/glm4_moe_lite_mtp.py @@ -54,10 +54,9 @@ from .glm4_moe_lite import ( Glm4MixtureOfExperts, Glm4MoeLite, Glm4MoeLiteDecoderLayer, - get_spec_layer_idx_from_weight_name, ) from .interfaces import SupportsPP -from .utils import maybe_prefix +from .utils import get_spec_layer_idx_from_weight_name, maybe_prefix class SharedHead(nn.Module): diff --git a/vllm/model_executor/models/glm4_moe_mtp.py b/vllm/model_executor/models/glm4_moe_mtp.py index 4d7b291df12..6708ba8139a 100644 --- a/vllm/model_executor/models/glm4_moe_mtp.py +++ b/vllm/model_executor/models/glm4_moe_mtp.py @@ -51,9 +51,8 @@ from .glm4_moe import ( Glm4MixtureOfExperts, Glm4MoE, Glm4MoeDecoderLayer, - get_spec_layer_idx_from_weight_name, ) -from .utils import maybe_prefix +from .utils import get_spec_layer_idx_from_weight_name, maybe_prefix class SharedHead(nn.Module): diff --git a/vllm/model_executor/models/glm_ocr_mtp.py b/vllm/model_executor/models/glm_ocr_mtp.py index 9b2369f93d3..0904b44eb28 100644 --- a/vllm/model_executor/models/glm_ocr_mtp.py +++ b/vllm/model_executor/models/glm_ocr_mtp.py @@ -41,13 +41,14 @@ from vllm.model_executor.model_loader.weight_utils import ( from vllm.platforms import current_platform from vllm.sequence import IntermediateTensors -from .glm4 import Glm4DecoderLayer, get_spec_layer_idx_from_weight_name +from .glm4 import Glm4DecoderLayer from .glm4_moe_lite_mtp import ( Glm4MoeLiteMultiTokenPredictor, SharedHead, ) from .interfaces import SupportsPP from .utils import ( + get_spec_layer_idx_from_weight_name, is_pp_missing_parameter, maybe_prefix, ) diff --git a/vllm/model_executor/models/glmasr.py b/vllm/model_executor/models/glmasr.py index cd168b6b461..1ff7224d7e4 100644 --- a/vllm/model_executor/models/glmasr.py +++ b/vllm/model_executor/models/glmasr.py @@ -65,7 +65,12 @@ from .interfaces import ( SupportsPP, SupportsTranscription, ) -from .utils import AutoWeightsLoader, init_vllm_registered_model, maybe_prefix +from .utils import ( + AutoWeightsLoader, + WeightsMapper, + init_vllm_registered_model, + maybe_prefix, +) from .whisper import ISO639_1_SUPPORTED_LANGS, _create_fake_bias_for_k_proj @@ -395,6 +400,14 @@ class GlmAsrEncoder(nn.Module): "qkv_proj": ["q_proj", "k_proj", "v_proj"], } + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_stacked={ + ".q_proj": (".qkv_proj", "q"), + ".k_proj": (".qkv_proj", "k"), + ".v_proj": (".qkv_proj", "v"), + } + ) + def __init__( self, config, @@ -496,44 +509,9 @@ class GlmAsrEncoder(nn.Module): return _GlmAsrEncoderOutput(last_hidden_state=hidden_states) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - """Custom weight loading to handle q_proj/k_proj/v_proj -> qkv_proj mapping.""" - from vllm.model_executor.model_loader.weight_utils import default_weight_loader - weights = _create_fake_bias_for_k_proj(weights, ".k_proj.weight") - - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ] - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - - for name, loaded_weight in weights: - for param_name, weight_name, shard_id in stacked_params_mapping: - if weight_name not in name: - continue - name = name.replace(weight_name, param_name) - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue - - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - break - else: - # Default weight loading for non-stacked params - if name.endswith(".bias") and name not in params_dict: - continue - if name not in params_dict: - continue - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + loader = AutoWeightsLoader(self) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) class GlmAsrFeatureInputs(TensorSchema): diff --git a/vllm/model_executor/models/gpt2.py b/vllm/model_executor/models/gpt2.py index 31fc3946536..01dc119f850 100644 --- a/vllm/model_executor/models/gpt2.py +++ b/vllm/model_executor/models/gpt2.py @@ -47,13 +47,11 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.sequence import IntermediateTensors from .interfaces import SupportsCrossEncoding, SupportsPP from .utils import ( AutoWeightsLoader, - is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -241,32 +239,25 @@ class GPT2Model(nn.Module): hidden_states = self.ln_f(hidden_states) return hidden_states - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - params_dict = dict(self.named_parameters(remove_duplicate=False)) - loaded_params: set[str] = set() + def _transpose_conv1d( + self, weights: Iterable[tuple[str, torch.Tensor]] + ) -> Iterable[tuple[str, torch.Tensor]]: + # HF's GPT-2 uses Conv1D instead of Linear, so its 2D weights are + # stored transposed relative to what vLLM expects. + # Note(zhuohan): the logic below might break quantized models. for name, loaded_weight in weights: - if ".attn.bias" in name or ".attn.masked_bias" in name: - # Skip attention mask. - # NOTE: "c_attn.bias" should not be skipped. - continue - - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - # The HF's GPT-2 implementation uses Conv1D instead of Linear. - # Because of this, we need to transpose the weights. - # Note(zhuohan): the logic below might break quantized models. - for conv1d_weight_name in ["c_attn", "c_proj", "c_fc"]: - if conv1d_weight_name not in name: - continue - if not name.endswith(".weight"): - continue + if name.endswith(".weight") and any( + proj in name for proj in ("c_attn", "c_proj", "c_fc") + ): loaded_weight = loaded_weight.t() - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + yield name, loaded_weight + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + # Skip attention mask buffers; NOTE: "c_attn.bias" must not be skipped. + loader = AutoWeightsLoader( + self, skip_substrs=[".attn.bias", ".attn.masked_bias"] + ) + return loader.load_weights(self._transpose_conv1d(weights)) class GPT2LMHeadModel(nn.Module, SupportsPP): diff --git a/vllm/model_executor/models/gpt_neox.py b/vllm/model_executor/models/gpt_neox.py index 8d44d12fc21..907ab777601 100644 --- a/vllm/model_executor/models/gpt_neox.py +++ b/vllm/model_executor/models/gpt_neox.py @@ -43,13 +43,11 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.sequence import IntermediateTensors from .interfaces import SupportsPP from .utils import ( AutoWeightsLoader, - is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -249,45 +247,25 @@ class GPTNeoXModel(nn.Module): hidden_states = self.final_layer_norm(hidden_states) return hidden_states - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() + def _repack_qkv( + self, weights: Iterable[tuple[str, torch.Tensor]] + ) -> Iterable[tuple[str, torch.Tensor]]: + # GPT-NeoX's fused QKV is laid out as (num_heads * 3 * head_size) on + # its output dim (0), while vLLM expects (3 * num_heads * head_size). + num_heads = self.config.num_attention_heads for name, loaded_weight in weights: - if ( - "attention.bias" in name - or "attention.masked_bias" in name - or "rotary_emb.inv_freq" in name - ): - continue - if "rotary_emb.cos_cached" in name or "rotary_emb.sin_cached" in name: - # Models trained using OpenRLHF may include - # these tensors in the checkpoint. Skip them. - continue - if is_pp_missing_parameter(name, self): - continue - param = params_dict[name] - if "query_key_value" in name: - # NOTE: GPT-NeoX's fused QKV's output_dim has the shape of - # (num_heads * 3 * head_size), while the - # required shape is (3 * num_heads * head_size). - # Thus, we need weight conversion. - output_dim = getattr(param, "output_dim", None) - num_heads = self.config.num_attention_heads - if output_dim is not None: - loaded_weight_shape = loaded_weight.shape - loaded_weight = loaded_weight.view( - loaded_weight_shape[:output_dim] - + (num_heads, 3, -1) - + loaded_weight_shape[output_dim + 1 :] - ) - loaded_weight = loaded_weight.transpose(output_dim, output_dim + 1) - loaded_weight = loaded_weight.reshape(loaded_weight_shape) + shape = loaded_weight.shape + loaded_weight = loaded_weight.view((num_heads, 3, -1) + shape[1:]) + loaded_weight = loaded_weight.transpose(0, 1) + loaded_weight = loaded_weight.reshape(shape) + yield name, loaded_weight - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + loader = AutoWeightsLoader( + self, skip_substrs=["attention.bias", "attention.masked_bias"] + ) + return loader.load_weights(self._repack_qkv(weights)) class GPTNeoXForCausalLM(nn.Module, SupportsPP): diff --git a/vllm/model_executor/models/hy_v3.py b/vllm/model_executor/models/hy_v3.py index 62bdd028bb5..68c6f238279 100644 --- a/vllm/model_executor/models/hy_v3.py +++ b/vllm/model_executor/models/hy_v3.py @@ -73,6 +73,7 @@ from .interfaces import MixtureOfExperts, SupportsLoRA, SupportsPP from .utils import ( AutoWeightsLoader, PPMissingLayer, + get_spec_layer_idx_from_weight_name, is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, @@ -650,21 +651,6 @@ class HYV3Model(nn.Module, MixtureOfExperts): return loaded_params -def get_spec_layer_idx_from_weight_name( - config: PretrainedConfig, weight_name: str -) -> int | None: - # HYV3MTP is enabled only when num_nextn_predict_layers is greater than 1 - if ( - hasattr(config, "num_nextn_predict_layers") - and config.num_nextn_predict_layers > 0 - ): - layer_idx = config.num_hidden_layers - for i in range(config.num_nextn_predict_layers): - if weight_name.startswith(f"model.layers.{layer_idx + i}."): - return layer_idx + i - return None - - class HYV3ForCausalLM(nn.Module, SupportsPP, SupportsLoRA): packed_modules_mapping = { "qkv_proj": ["q_proj", "k_proj", "v_proj"], diff --git a/vllm/model_executor/models/hy_v3_mtp.py b/vllm/model_executor/models/hy_v3_mtp.py index 77f323aae42..193280d7c85 100644 --- a/vllm/model_executor/models/hy_v3_mtp.py +++ b/vllm/model_executor/models/hy_v3_mtp.py @@ -51,8 +51,12 @@ from vllm.v1.outputs import SamplerOutput from vllm.v1.sample.metadata import SamplingMetadata from vllm.v1.sample.sampler import Sampler -from .hy_v3 import HYV3DecoderLayer, get_spec_layer_idx_from_weight_name -from .utils import is_pp_missing_parameter, maybe_prefix +from .hy_v3 import HYV3DecoderLayer +from .utils import ( + get_spec_layer_idx_from_weight_name, + is_pp_missing_parameter, + maybe_prefix, +) def _is_moe(config: PretrainedConfig) -> bool: diff --git a/vllm/model_executor/models/intern_vit.py b/vllm/model_executor/models/intern_vit.py index 7da498bf6de..44ed3ec0fed 100644 --- a/vllm/model_executor/models/intern_vit.py +++ b/vllm/model_executor/models/intern_vit.py @@ -36,8 +36,8 @@ from vllm.model_executor.layers.linear import ( RowParallelLinear, ) from vllm.model_executor.layers.quantization import QuantizationConfig -from vllm.model_executor.model_loader.weight_utils import default_weight_loader +from .utils import AutoWeightsLoader from .vision import is_vit_use_data_parallel, run_dp_sharded_vision_model NORM2FN = { @@ -445,11 +445,5 @@ class InternVisionModel(nn.Module): return encoder_outputs def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - for name, loaded_weight in weights: - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + loader = AutoWeightsLoader(self) + return loader.load_weights(weights) diff --git a/vllm/model_executor/models/interns1_vit.py b/vllm/model_executor/models/interns1_vit.py index 533f0681c1d..a5bc781561f 100644 --- a/vllm/model_executor/models/interns1_vit.py +++ b/vllm/model_executor/models/interns1_vit.py @@ -20,7 +20,8 @@ from vllm.model_executor.layers.conv import Conv2dLayer from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ColumnParallelLinear, RowParallelLinear from vllm.model_executor.layers.quantization import QuantizationConfig -from vllm.model_executor.model_loader.weight_utils import default_weight_loader + +from .utils import AutoWeightsLoader NORM2FN = { "rms_norm": RMSNorm, @@ -433,11 +434,5 @@ class InternS1VisionModel(nn.Module): return encoder_outputs def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - for name, loaded_weight in weights: - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + loader = AutoWeightsLoader(self) + return loader.load_weights(weights) diff --git a/vllm/model_executor/models/jamba.py b/vllm/model_executor/models/jamba.py index 84e96def6c1..33a8c636417 100644 --- a/vllm/model_executor/models/jamba.py +++ b/vllm/model_executor/models/jamba.py @@ -14,10 +14,7 @@ from vllm.config import CacheConfig, ModelConfig, VllmConfig from vllm.distributed import get_tensor_model_parallel_world_size from vllm.distributed.parallel_state import get_pp_group from vllm.model_executor.layers.attention import Attention -from vllm.model_executor.layers.fused_moe import ( - FusedMoE, - fused_moe_make_expert_params_mapping, -) +from vllm.model_executor.layers.fused_moe import FusedMoE from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ( QKVParallelLinear, @@ -38,7 +35,6 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.model_executor.models.llama import LlamaMLP as JambaMLP from vllm.sequence import IntermediateTensors @@ -52,7 +48,6 @@ from .interfaces import ( from .utils import ( AutoWeightsLoader, WeightsMapper, - is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -378,87 +373,6 @@ class JambaModel(nn.Module): hidden_states, _ = self.final_layernorm(hidden_states, residual) return hidden_states - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - # Params for weights, fp8 weight scales, fp8 activation scales - # (param_name, weight_name, expert_id, shard_id) - return fused_moe_make_expert_params_mapping( - self, - ckpt_gate_proj_name="gate_proj", - ckpt_down_proj_name="down_proj", - ckpt_up_proj_name="up_proj", - num_experts=self.config.num_experts, - ) - - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - (".gate_up_proj", ".gate_proj", 0), - (".gate_up_proj", ".up_proj", 1), - ] - - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - expert_params_mapping = self.get_expert_mapping() - for name, loaded_weight in weights: - if "rotary_emb.inv_freq" in name: - continue - for param_name, weight_name, shard_id in stacked_params_mapping: - if weight_name not in name: - continue - if "experts" in name: - continue - name = name.replace(weight_name, param_name) - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue - # Skip layers on other devices. - if is_pp_missing_parameter(name, self): - continue - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - break - else: - for ( - param_name, - weight_name, - expert_id, - shard_id, - ) in expert_params_mapping: - if weight_name not in name: - continue - - if is_pp_missing_parameter(name, self): - continue - name = name.replace(weight_name, param_name) - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader( - param, - loaded_weight, - name, - shard_id=shard_id, - expert_id=expert_id, - ) - break - else: - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = getattr( - param, "weight_loader", default_weight_loader - ) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params - class JambaForCausalLM( nn.Module, @@ -470,6 +384,14 @@ class JambaForCausalLM( ): hf_to_vllm_mapper = WeightsMapper( orig_to_new_substr={".self_attn.": ".", ".A_log": ".A"}, + orig_to_new_stacked={ + ".q_proj": (".qkv_proj", "q"), + ".k_proj": (".qkv_proj", "k"), + ".v_proj": (".qkv_proj", "v"), + # .experts.gate_up_proj must be handled by MoERunner.load_weights for EP + ".feed_forward.gate_proj": (".feed_forward.gate_up_proj", 0), + ".feed_forward.up_proj": (".feed_forward.gate_up_proj", 1), + }, ) packed_modules_mapping = { "qkv_proj": [ @@ -577,9 +499,6 @@ class JambaForCausalLM( loader = AutoWeightsLoader(self) return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - return self.model.get_expert_mapping() - class JambaForSequenceClassification(JambaForCausalLM): is_pooling_model = True diff --git a/vllm/model_executor/models/kimi_linear.py b/vllm/model_executor/models/kimi_linear.py index 307b24ac112..057d4d01cb0 100644 --- a/vllm/model_executor/models/kimi_linear.py +++ b/vllm/model_executor/models/kimi_linear.py @@ -52,6 +52,7 @@ from .interfaces import HasInnerState, IsHybrid, MixtureOfExperts, SupportsPP from .utils import ( AutoWeightsLoader, PPMissingLayer, + get_spec_layer_idx_from_weight_name, is_pp_missing_parameter, make_layers, maybe_prefix, @@ -643,16 +644,3 @@ class KimiLinearForCausalLM( skip_prefixes=(["lm_head."] if self.config.tie_word_embeddings else None), ) return loader.load_weights(weights) - - -def get_spec_layer_idx_from_weight_name( - config: KimiLinearConfig, weight_name: str -) -> int | None: - if hasattr(config, "num_nextn_predict_layers") and ( - config.num_nextn_predict_layers > 0 - ): - layer_idx = config.num_hidden_layers - for i in range(config.num_nextn_predict_layers): - if weight_name.startswith(f"model.layers.{layer_idx + i}."): - return layer_idx + i - return None diff --git a/vllm/model_executor/models/laguna.py b/vllm/model_executor/models/laguna.py index c79e8d48cbe..e71054f4da3 100644 --- a/vllm/model_executor/models/laguna.py +++ b/vllm/model_executor/models/laguna.py @@ -2,8 +2,7 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project """Inference-only Laguna model compatible with HuggingFace weights.""" -import typing -from collections.abc import Callable, Iterable +from collections.abc import Iterable from itertools import islice import torch @@ -20,10 +19,7 @@ from vllm.distributed import ( ) from vllm.logger import init_logger from vllm.model_executor.layers.attention import Attention -from vllm.model_executor.layers.fused_moe import ( - FusedMoE, - fused_moe_make_expert_params_mapping, -) +from vllm.model_executor.layers.fused_moe import FusedMoE from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ( ColumnParallelLinear, @@ -38,10 +34,6 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import ( - default_weight_loader, - maybe_remap_kv_scale_name, -) from vllm.model_executor.models.interfaces import ( EagleModelMixin, SupportsEagle3, @@ -51,8 +43,8 @@ from vllm.model_executor.models.interfaces import ( from vllm.model_executor.models.utils import ( AutoWeightsLoader, PPMissingLayer, + WeightsMapper, extract_layer_index, - is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -575,6 +567,14 @@ class LagunaDecoderLayer(nn.Module): @support_torch_compile class LagunaModel(nn.Module, EagleModelMixin): + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_stacked={ + ".q_proj": (".qkv_proj", "q"), + ".k_proj": (".qkv_proj", "k"), + ".v_proj": (".qkv_proj", "v"), + } + ) + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): super().__init__() @@ -675,167 +675,41 @@ class LagunaModel(nn.Module, EagleModelMixin): return hidden_states, aux_hidden_states return hidden_states - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - """Get expert parameter mapping for weight loading. - - Returns mapping tuples of (param_name, weight_name, expert_id, shard_id) - that handle both weights and quantization scales. - """ - return fused_moe_make_expert_params_mapping( - self, - ckpt_gate_proj_name="gate_proj", - ckpt_down_proj_name="down_proj", - ckpt_up_proj_name="up_proj", - num_experts=self.config.num_experts, - num_redundant_experts=self.num_redundant_experts, - ) - - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - # gate_proj and up_proj are loaded as separate Linears (see - # LagunaMLP) so no merge entry is needed here. - ] - - # Suffixes to skip for GPTQ/modelopt models if param doesn't exist - ignore_suffixes = ( - ".bias", - "_bias", - ".k_scale", - "_k_scale", - ".v_scale", - "_v_scale", - ".weight_scale", - "_weight_scale", - ".input_scale", - "_input_scale", - ) - - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - expert_params_mapping = self.get_expert_mapping() - + def _slice_sink( + self, weights: Iterable[tuple[str, torch.Tensor]] + ) -> Iterable[tuple[str, torch.Tensor]]: + tp_size = get_tensor_model_parallel_world_size() tp_rank = get_tensor_model_parallel_rank() - for name, loaded_weight in weights: # Handle attention sinks (distributed across ranks). Derive the # per-rank slice from the parameter's own shape so per-layer # variations in head count are handled correctly. if "sink" in name: - param = params_dict.get(name) - if param is not None: - layer_heads_per_rank = param.shape[0] - layer_head_start = tp_rank * layer_heads_per_rank - narrow_weight = loaded_weight.narrow( - 0, layer_head_start, layer_heads_per_rank - ) - param.data.copy_(narrow_weight) - loaded_params.add(name) - continue + heads_per_rank = loaded_weight.shape[0] // tp_size + loaded_weight = loaded_weight.narrow( + 0, tp_rank * heads_per_rank, heads_per_rank + ) + yield name, loaded_weight - # Handle stacked params (QKV, gate_up for - # non-expert layers and shared_expert) - for param_name, weight_name, shard_id in stacked_params_mapping: - if weight_name not in name: - continue - # Skip expert weights - handled below via expert_params_mapping - if "mlp.experts" in name and "shared_expert" not in name: - continue - name = name.replace(weight_name, param_name) - - if name.endswith(ignore_suffixes) and name not in params_dict: - continue - if is_pp_missing_parameter(name, self): - continue - # Remap FP8 kv_scale names for backwards compatibility - if name.endswith("scale"): - name = maybe_remap_kv_scale_name(name, params_dict) - if name is None: - continue - if name not in params_dict: - continue - - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - if weight_loader == default_weight_loader: - weight_loader(param, loaded_weight) - else: - weight_loader(param, loaded_weight, shard_id) - loaded_params.add(name) - break - else: - # Try expert params mapping (handles weights + quantization scales) - is_expert_weight = False - for mapping in expert_params_mapping: - param_name, weight_name, expert_id, shard_id = mapping - if weight_name not in name: - continue - - # Mark as expert weight so we skip regular loading below - is_expert_weight = True - - # Create mapped name without modifying original - name_mapped = name.replace(weight_name, param_name) - - if is_pp_missing_parameter(name_mapped, self): - continue - if ( - name_mapped.endswith(ignore_suffixes) - and name_mapped not in params_dict - ): - continue - if name_mapped not in params_dict: - continue - - param = params_dict[name_mapped] - # Use return_success to handle expert parallelism correctly - weight_loader = typing.cast( - Callable[..., bool], param.weight_loader - ) - success = weight_loader( - param, - loaded_weight, - name_mapped, - shard_id=shard_id, - expert_id=expert_id, - return_success=True, - ) - if success: - loaded_params.add(name_mapped) - break - else: - # Expert weight not mapped to this rank - skip - if is_expert_weight: - continue - - # Remap kv_scale names before the ignore_suffixes filter: - # the suffix list includes .k_scale/.v_scale, so filtering - # first drops the checkpoint key before remap can rewrite - # it to the .attn.* name that exists in params_dict. - name = maybe_remap_kv_scale_name(name, params_dict) - if name is None: - continue - - if name.endswith(ignore_suffixes) and name not in params_dict: - continue - - if is_pp_missing_parameter(name, self): - continue - - if name not in params_dict: - continue - - param = params_dict[name] - weight_loader = getattr( - param, "weight_loader", default_weight_loader - ) - weight_loader(param, loaded_weight) - loaded_params.add(name) - - return loaded_params + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + loader = AutoWeightsLoader( + self, + ignore_unexpected_suffixes=[ + ".bias", + "_bias", + ".k_scale", + "_k_scale", + ".v_scale", + "_v_scale", + ".weight_scale", + "_weight_scale", + ".input_scale", + "_input_scale", + ], + ) + return loader.load_weights( + self._slice_sink(weights), mapper=self.hf_to_vllm_mapper + ) class LagunaForCausalLM(nn.Module, SupportsPP, SupportsLoRA, SupportsEagle3): @@ -892,9 +766,6 @@ class LagunaForCausalLM(nn.Module, SupportsPP, SupportsLoRA, SupportsEagle3): logits = self.logits_processor(self.lm_head, hidden_states) return logits - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - return self.model.get_expert_mapping() - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: loader = AutoWeightsLoader( self, diff --git a/vllm/model_executor/models/lfm2_moe.py b/vllm/model_executor/models/lfm2_moe.py index 94f7f4e2890..698f1fb72ef 100644 --- a/vllm/model_executor/models/lfm2_moe.py +++ b/vllm/model_executor/models/lfm2_moe.py @@ -15,10 +15,7 @@ from vllm.distributed import ( ) from vllm.model_executor.layers.activation import SiluAndMul from vllm.model_executor.layers.attention import Attention -from vllm.model_executor.layers.fused_moe import ( - FusedMoE, - fused_moe_make_expert_params_mapping, -) +from vllm.model_executor.layers.fused_moe import FusedMoE from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ( MergedColumnParallelLinear, @@ -40,10 +37,6 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import ( - default_weight_loader, - maybe_remap_moe_expert_param_name, -) from vllm.sequence import IntermediateTensors from vllm.transformers_utils.configs.lfm2_moe import Lfm2MoeConfig @@ -60,7 +53,6 @@ from .utils import ( PPMissingLayer, WeightsMapper, extract_layer_index, - is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -167,6 +159,7 @@ class Lfm2MoeSparseMoeBlock(nn.Module): scoring_func="sigmoid", e_score_correction_bias=self.gate.e_score_correction_bias, routed_scaling_factor=self.routed_scaling_factor, + ckpt_names=("w1", "w2", "w3"), ) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: @@ -402,6 +395,22 @@ class Lfm2MoeShortConvDecoderLayer(nn.Module): @support_torch_compile class Lfm2MoeModel(nn.Module): + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_substr={ + ".conv.": ".short_conv.", + "expert_bias": "gate.e_score_correction_bias", + }, + orig_to_new_stacked={ + ".q_proj": (".qkv_proj", "q"), + ".k_proj": (".qkv_proj", "k"), + ".v_proj": (".qkv_proj", "v"), + # Scoped (with trailing dots) to the dense MLP so routed experts.*.w1/w3 + # (loaded by FusedMoE) are left untouched and .w1 does not match inside .w13 + ".feed_forward.w1.": (".feed_forward.w13.", 0), + ".feed_forward.w3.": (".feed_forward.w13.", 1), + }, + ) + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): super().__init__() @@ -487,102 +496,9 @@ class Lfm2MoeModel(nn.Module): hidden_states, _ = self.embedding_norm(hidden_states, residual) return hidden_states - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - return fused_moe_make_expert_params_mapping( - self, - ckpt_gate_proj_name="w1", - ckpt_down_proj_name="w2", - ckpt_up_proj_name="w3", - num_experts=self.config.num_experts, - num_redundant_experts=self.num_redundant_experts, - ) - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - (".qkv_proj", ".q_proj", "q"), - (".qkv_proj", ".k_proj", "k"), - (".qkv_proj", ".v_proj", "v"), - (".w13", ".w1", 0), - (".w13", ".w3", 1), - ] - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - expert_params_mapping = self.get_expert_mapping() - for name, loaded_weight in weights: - if "expert_bias" in name: - name = name.replace("expert_bias", "gate.e_score_correction_bias") - - if ".conv." in name: - name = name.replace(".conv.", ".short_conv.", 1) - - for param_name, weight_name, shard_id in stacked_params_mapping: - # Skip non-stacked layers and experts (experts handled below). - # Use segment-boundary matching (trailing dot) to prevent - # e.g. ".w1" from matching inside ".w13" in pre-fused keys. - if weight_name + "." not in name: - continue - - if ("feed_forward.experts." in name) and name not in params_dict: - continue - name = name.replace(weight_name + ".", param_name + ".") - # Skip loading extra bias for GPTQ models. - if ( - name.endswith(".bias") or name.endswith("_bias") - ) and name not in params_dict: - continue - # Skip layers on other devices. - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - break - else: - for mapping in expert_params_mapping: - param_name, weight_name, expert_id, shard_id = mapping - - if weight_name not in name: - continue - - name = name.replace(weight_name, param_name) - # Skip layers on other devices. - if is_pp_missing_parameter(name, self): - continue - - # Skip loading extra bias for GPTQ models. - if ( - name.endswith(".bias") or name.endswith("_bias") - ) and name not in params_dict: - continue - param = params_dict[name] - - weight_loader = param.weight_loader - weight_loader( - param, - loaded_weight, - name, - shard_id=shard_id, - expert_id=expert_id, - ) - break - else: - # Skip loading extra bias for GPTQ models. - if ( - name.endswith(".bias") or name.endswith("_bias") - ) and name not in params_dict: - continue - # Skip layers on other devices. - if is_pp_missing_parameter(name, self): - continue - name = maybe_remap_moe_expert_param_name(name, params_dict) - param = params_dict[name] - weight_loader = getattr( - param, "weight_loader", default_weight_loader - ) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + loader = AutoWeightsLoader(self, ignore_unexpected_suffixes=[".bias", "_bias"]) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) class Lfm2MoeForCausalLM( @@ -763,6 +679,3 @@ class Lfm2MoeForCausalLM( skip_prefixes=(["lm_head."] if self.config.tie_word_embeddings else None), ) return loader.load_weights(weights) - - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - return self.model.get_expert_mapping() diff --git a/vllm/model_executor/models/molmo.py b/vllm/model_executor/models/molmo.py index b3279e7dbd2..cf4550a1076 100644 --- a/vllm/model_executor/models/molmo.py +++ b/vllm/model_executor/models/molmo.py @@ -46,7 +46,6 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.model_executor.models.module_mapping import MultiModelKeys from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import ( @@ -76,7 +75,6 @@ from .interfaces import ( from .utils import ( AutoWeightsLoader, WeightsMapper, - is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -888,20 +886,7 @@ class MolmoModel(nn.Module, SupportsQuant): return hidden_states def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - - for name, loaded_weight in weights: - if name.endswith(".bias") and name not in params_dict: - continue - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + return AutoWeightsLoader(self).load_weights(weights) def _lowest_multiple(x: int, k: int) -> int: diff --git a/vllm/model_executor/models/molmo2.py b/vllm/model_executor/models/molmo2.py index 22a2b0cf328..d8de4f5ed17 100644 --- a/vllm/model_executor/models/molmo2.py +++ b/vllm/model_executor/models/molmo2.py @@ -51,7 +51,6 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.model_executor.models.module_mapping import MultiModelKeys from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import ( @@ -89,7 +88,6 @@ from .utils import ( WeightsMapper, _merge_multimodal_embeddings, extract_layer_index, - is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -1221,20 +1219,7 @@ class Molmo2TextModel(nn.Module, SupportsQuant): return hidden_states def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - - for name, loaded_weight in weights: - if name.endswith(".bias") and name not in params_dict: - continue - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + return AutoWeightsLoader(self).load_weights(weights) def get_patches_grid_size( diff --git a/vllm/model_executor/models/nemotron_h.py b/vllm/model_executor/models/nemotron_h.py index 016bff18ec2..84f24094a3d 100644 --- a/vllm/model_executor/models/nemotron_h.py +++ b/vllm/model_executor/models/nemotron_h.py @@ -18,8 +18,7 @@ # limitations under the License. """Inference-only NemotronH model.""" -import typing -from collections.abc import Callable, Iterable, Mapping +from collections.abc import Iterable, Mapping from itertools import islice import torch @@ -59,10 +58,6 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import ( - default_weight_loader, - maybe_remap_kv_scale_name, -) from vllm.model_executor.models.interfaces import ( EagleModelMixin, HasInnerState, @@ -78,7 +73,6 @@ from vllm.model_executor.models.interfaces import ( from vllm.model_executor.models.utils import ( AutoWeightsLoader, WeightsMapper, - is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -222,6 +216,7 @@ class NemotronHMoE(nn.Module): intermediate_size=config.moe_intermediate_size, renormalize=config.norm_topk_prob, quant_config=quant_config, + ckpt_names=("up_proj", "down_proj", ""), use_grouped_topk=True, num_expert_group=config.n_group, topk_group=config.topk_group, @@ -680,9 +675,12 @@ class NemotronHModel(nn.Module, EagleModelMixin): return max_experts def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: + # Consumed by `get_moe_expert_mapping` (bitsandbytes / LoRA); the main + # weight load is self-served by each `RoutedExperts` layer. Sized to the + # MAX expert count so heterogeneous puzzle models load every expert. if self.has_moe: # (param_name, weight_name, expert_id, shard_id) - expert_params_mapping = fused_moe_make_expert_params_mapping( + return fused_moe_make_expert_params_mapping( # - FusedMoe.w1 (aka gate_proj) should be up_proj since that's # what the activation is applied to # - FusedMoe.w3 (aka up_proj) should be ignored since we're @@ -694,102 +692,9 @@ class NemotronHModel(nn.Module, EagleModelMixin): num_experts=self._get_max_n_routed_experts(), num_redundant_experts=getattr(self, "num_redundant_experts", 0), ) - return expert_params_mapping return [] - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ] - - expert_params_mapping = self.get_expert_mapping() - - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - for name, loaded_weight in weights: - if "scale" in name or "zero_point" in name: - # Remapping the name of FP8 kv-scale. - name = maybe_remap_kv_scale_name(name, params_dict) - if name is None: - continue - - # Skip MTP/spec decode layers early (before stacked params mapping) - if name.startswith("mtp."): - continue - - # load stacked params - for param_name, weight_name, shard_id in stacked_params_mapping: - if weight_name not in name: - continue - name = name.replace(weight_name, param_name) - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue - - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - break - - # load other params - else: - is_expert_weight = False - for mapping in expert_params_mapping: - param_name, weight_name, expert_id, shard_id = mapping - if weight_name not in name: - continue - - # Anyway, this is an expert weight and should not be - # attempted to load as other weights later - is_expert_weight = True - - # Do not modify `name` since the loop may continue here - # Instead, create a new variable - name_mapped = name.replace(weight_name, param_name) - - if is_pp_missing_parameter(name_mapped, self): - continue - param = params_dict[name_mapped] - # We should ask the weight loader to return success or not - # here since otherwise we may skip experts with other - # available replicas. - weight_loader = typing.cast( - Callable[..., bool], param.weight_loader - ) - success = weight_loader( - param, - loaded_weight, - name_mapped, - shard_id=shard_id, - expert_id=expert_id, - return_success=True, - ) - if success: - name = name_mapped - break - else: - if is_expert_weight: - continue - - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = getattr( - param, "weight_loader", default_weight_loader - ) - weight_loader(param, loaded_weight) - - loaded_params.add(name) - return loaded_params - class NemotronHForCausalLM( nn.Module, @@ -809,6 +714,11 @@ class NemotronHForCausalLM( hf_to_vllm_mapper = WeightsMapper( orig_to_new_prefix={"backbone": "model"}, orig_to_new_substr={"A_log": "A", "embeddings": "embed_tokens"}, + orig_to_new_stacked={ + ".q_proj": (".qkv_proj", "q"), + ".k_proj": (".qkv_proj", "k"), + ".v_proj": (".qkv_proj", "v"), + }, ) packed_modules_mapping = { diff --git a/vllm/model_executor/models/olmo_hybrid.py b/vllm/model_executor/models/olmo_hybrid.py index 49969cdc586..51bc410363c 100644 --- a/vllm/model_executor/models/olmo_hybrid.py +++ b/vllm/model_executor/models/olmo_hybrid.py @@ -63,16 +63,13 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import ( - default_weight_loader, -) from vllm.sequence import IntermediateTensors from .interfaces import HasInnerState, IsHybrid, SupportsLoRA, SupportsPP from .utils import ( AutoWeightsLoader, + WeightsMapper, extract_layer_index, - is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -302,6 +299,23 @@ class OlmoHybridDecoderLayer(nn.Module): @support_torch_compile class OlmoHybridModel(nn.Module): + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_stacked={ + ".self_attn.q_proj": (".self_attn.qkv_proj", "q"), + ".self_attn.k_proj": (".self_attn.qkv_proj", "k"), + ".self_attn.v_proj": (".self_attn.qkv_proj", "v"), + ".linear_attn.q_proj": (".linear_attn.in_proj_qkvg", 0), + ".linear_attn.k_proj": (".linear_attn.in_proj_qkvg", 1), + ".linear_attn.v_proj": (".linear_attn.in_proj_qkvg", 2), + ".linear_attn.g_proj": (".linear_attn.in_proj_qkvg", 3), + ".q_conv1d": (".conv1d", 0), + ".k_conv1d": (".conv1d", 1), + ".v_conv1d": (".conv1d", 2), + ".gate_proj": (".gate_up_proj", 0), + ".up_proj": (".gate_up_proj", 1), + } + ) + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): super().__init__() self.config = vllm_config.model_config.hf_config @@ -359,79 +373,8 @@ class OlmoHybridModel(nn.Module): return hidden_states def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] - - linear_attn_stacked_params_mapping = [ - ("in_proj_qkvg", "q_proj", 0), - ("in_proj_qkvg", "k_proj", 1), - ("in_proj_qkvg", "v_proj", 2), - ("in_proj_qkvg", "g_proj", 3), - ("conv1d", "q_conv1d", 0), - ("conv1d", "k_conv1d", 1), - ("conv1d", "v_conv1d", 2), - ] - - params_dict = dict(self.named_parameters(remove_duplicate=False)) - loaded_params: set[str] = set() - - for name, loaded_weight in weights: - if is_pp_missing_parameter(name, self): - continue - - handled = False - - if "linear_attn" in name: - for ( - param_name, - weight_name, - shard_id, - ) in linear_attn_stacked_params_mapping: - if weight_name not in name: - continue - mapped_name = name.replace(weight_name, param_name) - if mapped_name.endswith(".bias") and ( - mapped_name not in params_dict - ): - continue - if mapped_name not in params_dict: - continue - param = params_dict[mapped_name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - name = mapped_name - handled = True - break - else: - for param_name, weight_name, shard_id in stacked_params_mapping: - if weight_name not in name: - continue - name = name.replace(weight_name, param_name) - if name.endswith(".bias") and name not in params_dict: - continue - if name not in params_dict: - continue - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - handled = True - break - - if not handled: - if name.endswith(".bias") and name not in params_dict: - continue - if name not in params_dict: - continue - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + loader = AutoWeightsLoader(self) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) class OlmoHybridForCausalLM( diff --git a/vllm/model_executor/models/openpangu.py b/vllm/model_executor/models/openpangu.py index 91120840bdf..fd40033e715 100644 --- a/vllm/model_executor/models/openpangu.py +++ b/vllm/model_executor/models/openpangu.py @@ -21,8 +21,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -import typing -from collections.abc import Callable, Iterable +from collections.abc import Iterable from typing import Any import torch @@ -44,10 +43,7 @@ from vllm.model_executor.layers.attention import ( Attention, StaticSinkAttention, ) -from vllm.model_executor.layers.fused_moe import ( - FusedMoE, - fused_moe_make_expert_params_mapping, -) +from vllm.model_executor.layers.fused_moe import FusedMoE from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ( ColumnParallelLinear, @@ -64,10 +60,6 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import ( - default_weight_loader, - maybe_remap_kv_scale_name, -) from vllm.model_executor.models.interfaces import ( MixtureOfExperts, SupportsLoRA, @@ -76,8 +68,8 @@ from vllm.model_executor.models.interfaces import ( from vllm.model_executor.models.utils import ( AutoWeightsLoader, PPMissingLayer, + WeightsMapper, extract_layer_index, - is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -1056,153 +1048,48 @@ class OpenPanguModel(nn.Module): hidden_states, _ = self.norm(hidden_states, residual) return hidden_states - def load_attn_mlp_weight( - self, - attn_mlp_replace_mapping: list[tuple[str, str, int]], - params_dict: dict[str, Any], - weight_name: str, - loaded_weight: torch.Tensor, - loaded_params: set[str], - ) -> bool: - for param_name, origin_name, shard_id in attn_mlp_replace_mapping: - if origin_name not in weight_name or ( - ("mlp.experts." in weight_name) and weight_name not in params_dict - ): - continue - weight_name_mapped = weight_name.replace(origin_name, param_name) - if ( - param_name == "fused_qkv_a_proj" - and weight_name_mapped not in params_dict - ): - continue - else: - weight_name = weight_name_mapped - if weight_name.endswith(".bias") and weight_name not in params_dict: - continue - if is_pp_missing_parameter(weight_name, self): - continue - - param = params_dict[weight_name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - loaded_params.add(weight_name) - return True - return False - - def load_expert_weight( - self, - expert_merge_mapping: list[tuple[str, str, int, str]], - params_dict: dict[str, Any], - weight_name: str, - loaded_weight: torch.Tensor, - loaded_params: set[str], - flag_dict: dict[str, bool], - ) -> bool: - for mapping in expert_merge_mapping: - param_name, origin_name, expert_id, shard_id = mapping - if origin_name not in weight_name: - continue - flag_dict["is_expert_weight"] = True - weight_name_mapped = weight_name.replace(origin_name, param_name) - if is_pp_missing_parameter(weight_name_mapped, self): - continue - param = params_dict[weight_name_mapped] - weight_loader = typing.cast(Callable[..., bool], param.weight_loader) - success = weight_loader( - param, - loaded_weight, - weight_name_mapped, - shard_id=shard_id, - expert_id=expert_id, - return_success=True, - ) - if success: - weight_name = weight_name_mapped - loaded_params.add(weight_name_mapped) - return True - return False - - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - attn_mlp_replace_mapping = [ - (".qkv_proj", ".q_proj", "q"), - (".qkv_proj", ".k_proj", "k"), - (".qkv_proj", ".v_proj", "v"), - (".fused_qkv_a_proj", ".q_a_proj", 0), - (".fused_qkv_a_proj", ".kv_a_proj_with_mqa", 1), - (".gate_up_proj", ".gate_proj", 0), - (".gate_up_proj", ".up_proj", 1), - ] - has_experts = hasattr(self.config, "n_routed_experts") - if has_experts: - expert_merge_mapping = fused_moe_make_expert_params_mapping( - self, - ckpt_gate_proj_name="gate_proj", - ckpt_down_proj_name="down_proj", - ckpt_up_proj_name="up_proj", - num_experts=self.config.n_routed_experts, - num_redundant_experts=self.num_redundant_experts, - ) - - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() + def _filter_spec_layers( + self, weights: Iterable[tuple[str, torch.Tensor]] + ) -> Iterable[tuple[str, torch.Tensor]]: + # Skip the MTP (spec-decode) layers that follow the main model's layers. + num_mtp = getattr(self.config, "num_nextn_predict_layers", 0) for name, loaded_weight in weights: - if "rotary_emb.inv_freq" in name: - continue - if self.config.tie_word_embeddings and "lm_head.weight" in name: - continue - - if ( - "layers" in name - and hasattr(self.config, "num_nextn_predict_layers") - and (self.config.num_nextn_predict_layers > 0) - ): + if "layers" in name and num_mtp > 0: layer_idx = int(name.split("layers.")[-1].split(".")[0]) mtp_idx = layer_idx - self.config.num_hidden_layers - if mtp_idx >= 0 and mtp_idx < self.config.num_nextn_predict_layers: - continue # skip spec decode layers for main model - - flag_dict = {"is_expert_weight": False} - if ( - self.load_attn_mlp_weight( - attn_mlp_replace_mapping, - params_dict, - name, - loaded_weight, - loaded_params, - ) - or has_experts - and self.load_expert_weight( - expert_merge_mapping, - params_dict, - name, - loaded_weight, - loaded_params, - flag_dict, - ) - ): - continue - else: - if flag_dict["is_expert_weight"]: + if 0 <= mtp_idx < num_mtp: continue - if name.endswith(".bias") and name not in params_dict: - continue - name = maybe_remap_kv_scale_name(name, params_dict) - if name.endswith("e_score_correction_bias"): - name = name.replace( - "e_score_correction_bias", "gate.e_score_correction_bias" - ) - if name is None: - continue - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) + yield name, loaded_weight + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + # .experts.gate_up_proj must be handled by MoERunner.load_weights for EP + stacked: dict[str, tuple[str, int | str]] = { + ".mlp.gate_proj": (".mlp.gate_up_proj", 0), + ".mlp.up_proj": (".mlp.gate_up_proj", 1), + ".shared_experts.gate_proj": (".shared_experts.gate_up_proj", 0), + ".shared_experts.up_proj": (".shared_experts.gate_up_proj", 1), + } + # The attention layout is config-dependent: MLA with a fused low-rank + # projection, MLA without it, or standard qkv. Only add the stacked + # entries whose fused target actually exists as a parameter. + param_names = list(dict(self.named_parameters())) + if any(".fused_qkv_a_proj." in n for n in param_names): + stacked[".q_a_proj"] = (".fused_qkv_a_proj", 0) + stacked[".kv_a_proj_with_mqa"] = (".fused_qkv_a_proj", 1) + if any(".qkv_proj." in n for n in param_names): + stacked[".q_proj"] = (".qkv_proj", "q") + stacked[".k_proj"] = (".qkv_proj", "k") + stacked[".v_proj"] = (".qkv_proj", "v") + mapper = WeightsMapper( + orig_to_new_substr={ + "e_score_correction_bias": "gate.e_score_correction_bias", + }, + orig_to_new_stacked=stacked, + ) + loader = AutoWeightsLoader(self) + loaded = loader.load_weights(self._filter_spec_layers(weights), mapper=mapper) self.post_weight_load() - return loaded_params + return loaded def post_weight_load(self) -> None: for name, module in self.named_modules(): diff --git a/vllm/model_executor/models/param2moe.py b/vllm/model_executor/models/param2moe.py index ff56cf505f0..48c0c714ed8 100644 --- a/vllm/model_executor/models/param2moe.py +++ b/vllm/model_executor/models/param2moe.py @@ -32,10 +32,7 @@ from vllm.distributed import ( ) from vllm.model_executor.layers.activation import SiluAndMul from vllm.model_executor.layers.attention import Attention -from vllm.model_executor.layers.fused_moe import ( - FusedMoE, - fused_moe_make_expert_params_mapping, -) +from vllm.model_executor.layers.fused_moe import FusedMoE from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ( MergedColumnParallelLinear, @@ -49,14 +46,13 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.sequence import IntermediateTensors from .interfaces import MixtureOfExperts, SupportsLoRA, SupportsPP from .utils import ( AutoWeightsLoader, PPMissingLayer, - is_pp_missing_parameter, + WeightsMapper, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -74,45 +70,17 @@ def _zero_mean_tensor(t: torch.Tensor) -> torch.Tensor: return t - t.mean() -def _rename_and_normalize_weights( +def _normalize_expert_bias( weights: Iterable[tuple[str, torch.Tensor]], ) -> Iterator[tuple[str, torch.Tensor]]: - """ - Translate HuggingFace Param2MoE weight names to vLLM internal names - and zero-mean the expert-bias tensor so the router stays balanced. + """Zero-mean the MoE router's per-expert score bias for load balance. - Mapping table (HF → vLLM): - model.word_embeddings.* → model.embed_tokens.* - *.attention.query_key_value.* → *.self_attn.qkv_proj.* - *.attention.dense.* → *.self_attn.o_proj.* - *.attention.query_layernorm.* → *.self_attn.q_layernorm.* - *.attention.key_layernorm.* → *.self_attn.k_layernorm.* - *.mlp.gate.expert_bias → *.mlp.gate.e_score_correction_bias - (also zero-meant for load balance) + The rename to ``e_score_correction_bias`` is done by the mapper; only the + tensor adjustment lives here, since a WeightsMapper cannot transform data. """ for name, w in weights: - # Embedding table - name = name.replace("model.word_embeddings.", "model.embed_tokens.") - # Fused QKV projection (HF: query_key_value → vLLM: qkv_proj) - name = name.replace(".attention.query_key_value.", ".self_attn.qkv_proj.") - # Output projection (HF: dense → vLLM: o_proj) - name = name.replace(".attention.dense.", ".self_attn.o_proj.") - # Per-head query norm - name = name.replace(".attention.query_layernorm.", ".self_attn.q_layernorm.") - # Per-head key norm - name = name.replace(".attention.key_layernorm.", ".self_attn.k_layernorm.") - # Catch any remaining .attention. → .self_attn. prefixes - # (e.g. future bias params on the projection layers) - name = name.replace(".attention.", ".self_attn.") - - # Expert-score bias: rename + zero-mean - if name.endswith(".mlp.gate.expert_bias"): - name = name.replace( - ".mlp.gate.expert_bias", - ".mlp.gate.e_score_correction_bias", - ) + if _is_expert_bias_name(name): w = _zero_mean_tensor(w) - yield name, w @@ -123,7 +91,8 @@ class Param2MoEAttention(nn.Module): Notable differences from a vanilla GQA layer: * The checkpoint fuses Q, K, V into a single ``query_key_value`` weight. vLLM receives it already renamed to ``qkv_proj`` by the weight-name - translator and splits it during ``load_weights``. + translator and loads it directly; ``QKVParallelLinear`` splits the + fused ``[Q|K|V]`` tensor internally. * Optional per-head RMS norms on Q and K (``use_qk_norm=True``). """ @@ -472,6 +441,28 @@ class Param2MoEDecoderLayer(nn.Module): class Param2MoEModel(nn.Module): + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_substr={ + "word_embeddings.": "embed_tokens.", + # Fused query_key_value is renamed to qkv_proj and loaded directly; + # QKVParallelLinear splits the fused [Q|K|V] tensor internally. + ".attention.query_key_value.": ".self_attn.qkv_proj.", + ".attention.dense.": ".self_attn.o_proj.", + ".attention.query_layernorm.": ".self_attn.q_layernorm.", + ".attention.key_layernorm.": ".self_attn.k_layernorm.", + # Catch-all for any remaining .attention. names (must come last). + ".attention.": ".self_attn.", + ".mlp.gate.expert_bias": ".mlp.gate.e_score_correction_bias", + }, + orig_to_new_stacked={ + # .experts.gate_up_proj must be handled by MoERunner.load_weights for EP + ".mlp.gate_proj": (".mlp.gate_up_proj", 0), + ".mlp.up_proj": (".mlp.gate_up_proj", 1), + ".shared_experts.gate_proj": (".shared_experts.gate_up_proj", 0), + ".shared_experts.up_proj": (".shared_experts.gate_up_proj", 1), + }, + ) + def __init__( self, *, @@ -558,147 +549,9 @@ class Param2MoEModel(nn.Module): self, weights: Iterable[tuple[str, torch.Tensor]], ) -> set[str]: - """ - Custom weight loader for the inner Param2MoEModel. - - Receives weights that have already been renamed/normalised by the - outer model and whose ``model.`` prefix has been stripped by - ``AutoWeightsLoader``. Handles: - 1. Fused QKV split (query_key_value → qkv_proj q/k/v shards). - 2. gate_proj + up_proj → gate_up_proj stacking (dense + shared-exp). - 3. Routed-expert weights via the fused-MoE mapping. - 4. All remaining weights via their default loader. - """ - config = self.config - num_heads: int = config.num_attention_heads - num_kv_heads: int = config.num_key_value_heads - head_dim: int = config.head_dim or (config.hidden_size // num_heads) - q_split = num_heads * head_dim - kv_split = num_kv_heads * head_dim - - stacked_params_mapping = [ - # (vllm_param_name, ckpt_weight_name, shard_id) - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] - - params_dict = dict(self.named_parameters(remove_duplicate=False)) - loaded_params: set[str] = set() - expert_params_mapping = self.get_expert_mapping() - - for name, loaded_weight in weights: - # ------------------------------------------------------------------ - # 1. Fused QKV: split into q / k / v shards for QKVParallelLinear - # ------------------------------------------------------------------ - if name.endswith(".self_attn.qkv_proj.weight"): - if name not in params_dict: - continue - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - q_w = loaded_weight[:q_split, :] - k_w = loaded_weight[q_split : q_split + kv_split, :] - v_w = loaded_weight[q_split + kv_split :, :] - weight_loader(param, q_w, "q") - weight_loader(param, k_w, "k") - weight_loader(param, v_w, "v") - loaded_params.add(name) - continue - - # ------------------------------------------------------------------ - # 2. gate_proj / up_proj → gate_up_proj (dense MLP + shared-exp.) - # ------------------------------------------------------------------ - matched_stacked = False - for param_name, weight_name, shard_id in stacked_params_mapping: - if weight_name not in name: - continue - if "mlp.experts" in name: # routed experts handled below - continue - new_name = name.replace(weight_name, param_name) - if new_name.endswith(".bias") and new_name not in params_dict: - continue - if new_name not in params_dict: - continue - if is_pp_missing_parameter(new_name, self): - continue - - param = params_dict[new_name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight, shard_id) - loaded_params.add(new_name) - matched_stacked = True - break - - if matched_stacked: - continue - - # ------------------------------------------------------------------ - # 3. Routed expert weights → fused-MoE kernel layout - # ------------------------------------------------------------------ - matched_expert = False - for ( - param_name, - weight_name, - expert_id, - shard_id, - ) in expert_params_mapping: - if weight_name not in name: - continue - new_name = name.replace(weight_name, param_name) - if is_pp_missing_parameter(new_name, self): - continue - if new_name not in params_dict: - continue - - param = params_dict[new_name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader( - param, - loaded_weight, - name, - shard_id=shard_id, - expert_id=expert_id, - ) - loaded_params.add(new_name) - matched_expert = True - break - - if matched_expert: - continue - - # ------------------------------------------------------------------ - # 4. All other weights: direct load (layernorms, embed_tokens, …) - # ------------------------------------------------------------------ - if name.endswith(".bias") and name not in params_dict: - continue - if name not in params_dict: - continue - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - try: - weight_loader(param, loaded_weight) - except Exception as e: - raise RuntimeError( - f"[param2moe] Failed to load weight '{name}' " - f"with shape {tuple(loaded_weight.shape)} " - f"into param type {type(param).__name__}: {e}" - ) from e - loaded_params.add(name) - - return loaded_params - - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - return fused_moe_make_expert_params_mapping( - self, - ckpt_gate_proj_name="gate_proj", - ckpt_down_proj_name="down_proj", - ckpt_up_proj_name="up_proj", - num_experts=self.config.num_experts, + loader = AutoWeightsLoader(self) + return loader.load_weights( + _normalize_expert_bias(weights), mapper=self.hf_to_vllm_mapper ) @@ -859,4 +712,4 @@ class Param2MoEForCausalLM( weights: Iterable[tuple[str, torch.Tensor]], ) -> set[str]: loader = AutoWeightsLoader(self) - return loader.load_weights(_rename_and_normalize_weights(weights)) + return loader.load_weights(weights) diff --git a/vllm/model_executor/models/qwen3_5.py b/vllm/model_executor/models/qwen3_5.py index a58c3c4dd71..bcd2576b74f 100644 --- a/vllm/model_executor/models/qwen3_5.py +++ b/vllm/model_executor/models/qwen3_5.py @@ -93,6 +93,7 @@ from .utils import ( extract_layer_index, make_empty_intermediate_tensors_factory, make_layers, + maybe_fuse_shared_experts, maybe_prefix, ) @@ -261,20 +262,20 @@ class Qwen3_5Model(Qwen3NextModel): self.aux_hidden_state_layers: tuple[int, ...] = () def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - mapper = self.hf_to_vllm_mapper # FSE must match construction (Qwen3NextSparseMoeBlock): reroute the # shared expert into the extra fused slot only when AITER FSE is both # requested and compatible with the quant spec. - is_fse = rocm_aiter_ops.is_fusion_moe_shared_experts_enabled() and ( - _is_shared_expert_fse_compatible(self.quant_config) - ) - if is_fse: - num_routed = self.config.num_experts - mapper = mapper | WeightsMapper( - orig_to_new_substr={"mlp.shared_expert.": f"mlp.experts.{num_routed}."} + if "moe" in self.config.model_type: + weights = maybe_fuse_shared_experts( + weights, + enabled=rocm_aiter_ops.is_fusion_moe_shared_experts_enabled() + and _is_shared_expert_fse_compatible(self.quant_config), + n_routed_experts=self.config.num_experts, + n_shared_experts=1, + ckpt_prefix="mlp.shared_expert", ) loader = AutoWeightsLoader(self) - return loader.load_weights(weights, mapper=mapper) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) class Qwen3_5ForCausalLMBase( diff --git a/vllm/model_executor/models/qwen3_5_mtp.py b/vllm/model_executor/models/qwen3_5_mtp.py index b1580bfee71..0620509f2db 100644 --- a/vllm/model_executor/models/qwen3_5_mtp.py +++ b/vllm/model_executor/models/qwen3_5_mtp.py @@ -41,9 +41,9 @@ from .interfaces import ( from .utils import ( AutoWeightsLoader, PPMissingLayer, - WeightsMapper, _merge_multimodal_embeddings, make_empty_intermediate_tensors_factory, + maybe_fuse_shared_experts, maybe_prefix, ) @@ -174,17 +174,18 @@ class Qwen3_5MultiTokenPredictor(nn.Module): return hidden_states def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - mapper = self.hf_to_vllm_mapper - is_fse = rocm_aiter_ops.is_fusion_moe_shared_experts_enabled() and ( - _is_shared_expert_fse_compatible(get_current_vllm_config().quant_config) + weights = maybe_fuse_shared_experts( + weights, + enabled=rocm_aiter_ops.is_fusion_moe_shared_experts_enabled() + and _is_shared_expert_fse_compatible( + get_current_vllm_config().quant_config + ), + n_routed_experts=getattr(self.config, "num_experts", 0), + n_shared_experts=1, + ckpt_prefix="mlp.shared_expert", ) - if is_fse: - num_routed = getattr(self.config, "num_experts", 0) - mapper = mapper | WeightsMapper( - orig_to_new_substr={"mlp.shared_expert.": f"mlp.experts.{num_routed}."} - ) loader = AutoWeightsLoader(self) - return loader.load_weights(weights, mapper=mapper) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) @support_torch_compile( diff --git a/vllm/model_executor/models/qwen3_dflash.py b/vllm/model_executor/models/qwen3_dflash.py index 7843ce72bd4..bf5ea501cc4 100644 --- a/vllm/model_executor/models/qwen3_dflash.py +++ b/vllm/model_executor/models/qwen3_dflash.py @@ -31,10 +31,6 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import ( - default_weight_loader, - maybe_remap_kv_scale_name, -) from vllm.multimodal.inputs import NestedTensors from vllm.transformers_utils.config import set_default_rope_theta from vllm.transformers_utils.repo_utils import get_hf_file_bytes @@ -44,6 +40,7 @@ from .qwen2 import Qwen2MLP as Qwen3MLP from .qwen3 import Qwen3ForCausalLM from .utils import ( AutoWeightsLoader, + WeightsMapper, get_draft_quant_config, maybe_prefix, process_eagle_weight, @@ -333,6 +330,17 @@ class DFlashQwen3DecoderLayer(nn.Module): @support_torch_compile class DFlashQwen3Model(nn.Module): + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_substr={"midlayer.": "layers.0."}, + orig_to_new_stacked={ + ".q_proj": (".qkv_proj", "q"), + ".k_proj": (".qkv_proj", "k"), + ".v_proj": (".qkv_proj", "v"), + ".gate_proj": (".gate_up_proj", 0), + ".up_proj": (".gate_up_proj", 1), + }, + ) + def __init__( self, *, @@ -624,51 +632,26 @@ class DFlashQwen3Model(nn.Module): hidden_states, _ = self.norm(hidden_states, residual) return hidden_states - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - (".qkv_proj", ".q_proj", "q"), - (".qkv_proj", ".k_proj", "k"), - (".qkv_proj", ".v_proj", "v"), - (".gate_up_proj", ".gate_proj", 0), - (".gate_up_proj", ".up_proj", 1), - ] - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() - tp_rank = get_tensor_model_parallel_rank() + def _preprocess( + self, weights: Iterable[tuple[str, torch.Tensor]] + ) -> Iterable[tuple[str, torch.Tensor]]: tp_size = get_tensor_model_parallel_world_size() + tp_rank = get_tensor_model_parallel_rank() for name, loaded_weight in weights: - if "midlayer." in name: - name = name.replace("midlayer.", "layers.0.") - if "scale" in name: - name = maybe_remap_kv_scale_name(name, params_dict) - if name is None: - continue if "attention_sink_bias" in name: - if name not in params_dict: - continue # Sink bias is per-head; shard it across TP ranks like the # attention heads themselves. - param = params_dict[name] heads_per_rank = loaded_weight.shape[0] // tp_size - head_start = tp_rank * heads_per_rank - narrow_weight = loaded_weight.narrow(0, head_start, heads_per_rank) - param.data.copy_(narrow_weight) - loaded_params.add(name) - continue - for param_name, weight_name, shard_id in stacked_params_mapping: - if weight_name not in name: - continue - name = name.replace(weight_name, param_name) - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - break - else: - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + loaded_weight = loaded_weight.narrow( + 0, tp_rank * heads_per_rank, heads_per_rank + ) + yield name, loaded_weight + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + loader = AutoWeightsLoader(self) + return loader.load_weights( + self._preprocess(weights), mapper=self.hf_to_vllm_mapper + ) class DFlashQwen3ForCausalLM(Qwen3ForCausalLM): diff --git a/vllm/model_executor/models/qwen3_next.py b/vllm/model_executor/models/qwen3_next.py index d87d19f02a6..50bc4cce187 100644 --- a/vllm/model_executor/models/qwen3_next.py +++ b/vllm/model_executor/models/qwen3_next.py @@ -71,6 +71,7 @@ from .utils import ( extract_layer_index, make_empty_intermediate_tensors_factory, make_layers, + maybe_fuse_shared_experts, maybe_prefix, ) @@ -713,17 +714,14 @@ class Qwen3NextModel(nn.Module, EagleModelMixin): return hidden_states def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - mapper = self.hf_to_vllm_mapper - if rocm_aiter_ops.is_fusion_moe_shared_experts_enabled(): - # AITER fused-shared-experts: route the shared_expert checkpoint - # weights into the extra fused expert slot. Merge (not mutate) so the - # shared class mapper isn't permanently altered. - num_routed = getattr(self.config, "num_experts", 0) - mapper = mapper | WeightsMapper( - orig_to_new_substr={"mlp.shared_expert.": f"mlp.experts.{num_routed}."} - ) + weights = maybe_fuse_shared_experts( + weights, + n_routed_experts=getattr(self.config, "num_experts", 0), + n_shared_experts=1, + ckpt_prefix="mlp.shared_expert", + ) loader = AutoWeightsLoader(self) - return loader.load_weights(weights, mapper=mapper) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) class QwenNextMixtureOfExperts(MixtureOfExperts): diff --git a/vllm/model_executor/models/qwen3_next_mtp.py b/vllm/model_executor/models/qwen3_next_mtp.py index 1704f30a84a..c832e955b24 100644 --- a/vllm/model_executor/models/qwen3_next_mtp.py +++ b/vllm/model_executor/models/qwen3_next_mtp.py @@ -7,7 +7,6 @@ from collections.abc import Iterable import torch from torch import nn -from vllm._aiter_ops import rocm_aiter_ops from vllm.compilation.decorators import support_torch_compile from vllm.config import VllmConfig from vllm.distributed.parallel_state import get_pp_group @@ -30,8 +29,8 @@ from vllm.transformers_utils.configs.qwen3_next import Qwen3NextConfig from .utils import ( AutoWeightsLoader, - WeightsMapper, make_empty_intermediate_tensors_factory, + maybe_fuse_shared_experts, maybe_prefix, ) @@ -154,16 +153,14 @@ class Qwen3NextMultiTokenPredictor(nn.Module): return hidden_states def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - mapper = self.hf_to_vllm_mapper - if rocm_aiter_ops.is_fusion_moe_shared_experts_enabled(): - # AITER fused-shared-experts: route the shared_expert checkpoint - # weights into the extra fused expert slot. - num_routed = self.config.num_experts - mapper = mapper | WeightsMapper( - orig_to_new_substr={"mlp.shared_expert.": f"mlp.experts.{num_routed}."} - ) + weights = maybe_fuse_shared_experts( + weights, + n_routed_experts=self.config.num_experts, + n_shared_experts=1, + ckpt_prefix="mlp.shared_expert", + ) loader = AutoWeightsLoader(self) - return loader.load_weights(weights, mapper=mapper) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) @support_torch_compile diff --git a/vllm/model_executor/models/sarvam.py b/vllm/model_executor/models/sarvam.py index f59579b1bcc..b9f9532ae1b 100644 --- a/vllm/model_executor/models/sarvam.py +++ b/vllm/model_executor/models/sarvam.py @@ -35,11 +35,7 @@ from vllm.distributed import ( get_tensor_model_parallel_world_size, ) from vllm.model_executor.layers.activation import SiluAndMul -from vllm.model_executor.layers.fused_moe import ( - FusedMoE, - MoERunner, - fused_moe_make_expert_params_mapping, -) +from vllm.model_executor.layers.fused_moe import FusedMoE, MoERunner from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ( ColumnParallelLinear, @@ -55,7 +51,6 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.sequence import IntermediateTensors from .bailing_moe import BailingMoeForCausalLM @@ -63,7 +58,7 @@ from .interfaces import MixtureOfExperts, SupportsLoRA, SupportsPP from .utils import ( AutoWeightsLoader, PPMissingLayer, - is_pp_missing_parameter, + WeightsMapper, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, @@ -448,6 +443,16 @@ class SarvamMLABlock(nn.Module): class SarvamMLAModel(nn.Module): + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_stacked={ + # .experts.gate_up_proj must be handled by MoERunner.load_weights for EP + ".mlp.gate_proj": (".mlp.gate_up_proj", 0), + ".mlp.up_proj": (".mlp.gate_up_proj", 1), + ".shared_experts.gate_proj": (".shared_experts.gate_up_proj", 0), + ".shared_experts.up_proj": (".shared_experts.gate_up_proj", 1), + } + ) + def __init__( self, *, @@ -532,97 +537,14 @@ class SarvamMLAModel(nn.Module): hidden_states, _ = self.norm(hidden_states, residual) return hidden_states - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - return fused_moe_make_expert_params_mapping( - self, - ckpt_gate_proj_name="gate_proj", - ckpt_down_proj_name="down_proj", - ckpt_up_proj_name="up_proj", - num_experts=self.config.num_experts, - ) - def load_weights( self, weights: Iterable[tuple[str, torch.Tensor]], ) -> set[str]: - """Load weights with stacked gate+up and MoE expert remapping.""" - weights = _normalized_weights(weights) - stacked_params_mapping = [ - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] - - params_dict = dict(self.named_parameters(remove_duplicate=False)) - loaded_params: set[str] = set() - expert_params_mapping = self.get_expert_mapping() - - for name, loaded_weight in weights: - for param_name, weight_name, shard_id in stacked_params_mapping: - if weight_name not in name: - continue - if "mlp.experts" in name: - continue - new_name = name.replace(weight_name, param_name) - if new_name.endswith(".bias") and new_name not in params_dict: - continue - if new_name not in params_dict: - continue - if is_pp_missing_parameter(new_name, self): - continue - - param = params_dict[new_name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight, shard_id) - loaded_params.add(new_name) - break - else: - mapped = False - for ( - param_name, - weight_name, - expert_id, - shard_id, - ) in expert_params_mapping: - if weight_name not in name: - continue - - new_name = name.replace(weight_name, param_name) - if is_pp_missing_parameter(new_name, self): - continue - if new_name not in params_dict: - continue - - param = params_dict[new_name] - weight_loader = getattr( - param, "weight_loader", default_weight_loader - ) - weight_loader( - param, - loaded_weight, - name, - shard_id=shard_id, - expert_id=expert_id, - ) - loaded_params.add(new_name) - mapped = True - break - - if mapped: - continue - - if name.endswith(".bias") and name not in params_dict: - continue - if name not in params_dict: - continue - if is_pp_missing_parameter(name, self): - continue - - param = params_dict[name] - weight_loader = getattr(param, "weight_loader", default_weight_loader) - weight_loader(param, loaded_weight) - loaded_params.add(name) - - return loaded_params + loader = AutoWeightsLoader(self) + return loader.load_weights( + _normalized_weights(weights), mapper=self.hf_to_vllm_mapper + ) class SarvamMixtureOfExperts(MixtureOfExperts): @@ -764,9 +686,6 @@ class SarvamMLAForCausalLM(nn.Module, SupportsPP, SupportsLoRA, SarvamMixtureOfE ) return loader.load_weights(weights) - def get_expert_mapping(self) -> list[tuple[str, str, int, str]]: - return self.model.get_expert_mapping() - class SarvamMoEForCausalLM(BailingMoeForCausalLM): """Same as BailingMoeForCausalLM, but normalizes gate expert_bias pre-load.""" diff --git a/vllm/model_executor/models/step3p5.py b/vllm/model_executor/models/step3p5.py index f8bd529e276..07a25d23c8c 100644 --- a/vllm/model_executor/models/step3p5.py +++ b/vllm/model_executor/models/step3p5.py @@ -53,6 +53,7 @@ from .utils import ( PPMissingLayer, WeightsMapper, extract_layer_index, + get_spec_layer_idx_from_weight_name, is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, @@ -978,18 +979,3 @@ class Step3p5ForCausalLM(nn.Module, SupportsPP, MixtureOfExperts): def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: loader = AutoWeightsLoader(self) return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) - - -def get_spec_layer_idx_from_weight_name( - config: ModelConfig, weight_name: str -) -> int | None: - if hasattr(config, "num_nextn_predict_layers") and ( - config.num_nextn_predict_layers > 0 - ): - layer_idx = config.num_hidden_layers - for i in range(config.num_nextn_predict_layers): - if weight_name.startswith( - f"layers.{layer_idx + i}." # Step3p5Model - ) or weight_name.startswith(f"model.layers.{layer_idx + i}."): # Step3p5MTP - return layer_idx + i - return None diff --git a/vllm/model_executor/models/step3p5_mtp.py b/vllm/model_executor/models/step3p5_mtp.py index 092f7a31aa7..b533a9111dc 100644 --- a/vllm/model_executor/models/step3p5_mtp.py +++ b/vllm/model_executor/models/step3p5_mtp.py @@ -18,8 +18,8 @@ from vllm.model_executor.layers.vocab_parallel_embedding import ( from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.sequence import IntermediateTensors -from .step3p5 import Step3p5DecoderLayer, get_spec_layer_idx_from_weight_name -from .utils import maybe_prefix +from .step3p5 import Step3p5DecoderLayer +from .utils import get_spec_layer_idx_from_weight_name, maybe_prefix logger = init_logger(__name__) diff --git a/vllm/model_executor/models/telechat2.py b/vllm/model_executor/models/telechat2.py index 113581d55ff..42fa6d6871a 100644 --- a/vllm/model_executor/models/telechat2.py +++ b/vllm/model_executor/models/telechat2.py @@ -26,19 +26,23 @@ import torch import torch.nn as nn from vllm.config import VllmConfig -from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.model_executor.models.llama import LlamaForCausalLM, LlamaModel from .llama import LlamaDecoderLayer -from .utils import ( - AutoWeightsLoader, - PPMissingLayer, - WeightsMapper, - is_pp_missing_parameter, -) +from .utils import AutoWeightsLoader, PPMissingLayer, WeightsMapper class TeleChat2Model(LlamaModel): + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_stacked={ + ".query": (".qkv_proj", "q"), + ".key": (".qkv_proj", "k"), + ".value": (".qkv_proj", "v"), + ".gate_proj": (".gate_up_proj", 0), + ".up_proj": (".gate_up_proj", 1), + } + ) + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): hf_config = vllm_config.model_config.hf_config @@ -65,62 +69,34 @@ class TeleChat2Model(LlamaModel): layer.mlp.gate_up_proj.bias = None layer.mlp.gate_up_proj.skip_bias_add = True - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - stacked_params_mapping = [ - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] - params_dict = dict(self.named_parameters()) - loaded_params: set[str] = set() + def _split_key_value( + self, weights: Iterable[tuple[str, torch.Tensor]] + ) -> Iterable[tuple[str, torch.Tensor]]: + # TeleChat2 stores k/v as a single per-head-interleaved `key_value` + # tensor. De-interleave it into separate k/v so the qkv_proj mapper + # can stack them. total_num_heads = self.config.n_head head_dim = self.config.hidden_size // total_num_heads for name, loaded_weight in weights: if "self_attn.key_value" in name: - k_weight = [] - v_weight = [] - for i in range(total_num_heads): - start = i * head_dim * 2 - k_weight.append(loaded_weight[start : start + head_dim, :]) - v_weight.append( - loaded_weight[start + head_dim : start + 2 * head_dim :] - ) - k_weight = torch.cat(k_weight, dim=0) - v_weight = torch.cat(v_weight, dim=0) - name = name.replace("key_value", "qkv_proj") - if is_pp_missing_parameter(name, self): - continue - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, k_weight, "k") - weight_loader(param, v_weight, "v") - elif "query" in name: - name = name.replace("query", "qkv_proj") - if is_pp_missing_parameter(name, self): - continue - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, "q") + starts = [i * head_dim * 2 for i in range(total_num_heads)] + k_weight = torch.cat( + [loaded_weight[s : s + head_dim, :] for s in starts], dim=0 + ) + v_weight = torch.cat( + [loaded_weight[s + head_dim : s + 2 * head_dim, :] for s in starts], + dim=0, + ) + yield name.replace("key_value", "key"), k_weight + yield name.replace("key_value", "value"), v_weight else: - for param_name, weight_name, shard_id in stacked_params_mapping: - if weight_name not in name: - continue - name = name.replace(weight_name, param_name) - if is_pp_missing_parameter(name, self): - continue - param = params_dict[name] - weight_loader = param.weight_loader - weight_loader(param, loaded_weight, shard_id) - break - else: - if is_pp_missing_parameter(name, self): - continue - param = params_dict[name] - weight_loader = getattr( - param, "weight_loader", default_weight_loader - ) - weight_loader(param, loaded_weight) - loaded_params.add(name) - return loaded_params + yield name, loaded_weight + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + loader = AutoWeightsLoader(self) + return loader.load_weights( + self._split_key_value(weights), mapper=self.hf_to_vllm_mapper + ) class TeleChat2ForCausalLM(LlamaForCausalLM): diff --git a/vllm/model_executor/models/utils.py b/vllm/model_executor/models/utils.py index 86dc72de76c..a6f39f6ebf8 100644 --- a/vllm/model_executor/models/utils.py +++ b/vllm/model_executor/models/utils.py @@ -35,6 +35,7 @@ if TYPE_CHECKING: from transformers import PretrainedConfig from transformers.conversion_mapping import WeightRenaming + from vllm.config.model import ModelConfig from vllm.model_executor.layers.quantization import QuantizationConfig logger = init_logger(__name__) @@ -424,6 +425,116 @@ class AutoWeightsLoader: return autoloaded_weights +def maybe_fuse_shared_experts( + weights: Iterable[tuple[str, torch.Tensor]], + *, + n_routed_experts: int, + n_shared_experts: int, + ckpt_prefix: str = "mlp.shared_experts", + enabled: bool | None = None, +) -> Iterable[tuple[str, torch.Tensor]]: + """Route AITER fused-shared-expert checkpoint weights into fused slots. + + When AITER fused-shared-experts is active, shared experts are packed into + the routed expert tensor. The checkpoint stores them under `ckpt_prefix` + as a single (possibly widened) tensor; this splits it into + `n_shared_experts` chunks named `mlp.experts.{n_routed_experts + j}` so + the `RoutedExperts` loader treats them as extra experts. Yields the input + unchanged when the fusion is inactive, so callers can wrap unconditionally. + + Args: + weights: Iterable of `(name, tensor)` checkpoint pairs. + n_routed_experts: Number of routed experts; offsets the fused slots. + n_shared_experts: Number of shared experts packed into the tensor. + ckpt_prefix: Checkpoint module name of the shared experts. + enabled: Whether AITER fused-shared-experts is active. Defaults to + `rocm_aiter_ops.is_fusion_moe_shared_experts_enabled()`; pass an + explicit value only when the model gates on something more (e.g. + quant-spec compatibility) and it must match its construction-time + decision. + + Yields: + `(name, tensor)` pairs with shared experts routed to fused slots. + """ + if enabled is None: + from vllm._aiter_ops import rocm_aiter_ops + + enabled = rocm_aiter_ops.is_fusion_moe_shared_experts_enabled() + if not enabled: + yield from weights + return + + # Match on the dotted boundary so e.g. "mlp.shared_expert." does not also + # catch a sibling "mlp.shared_expert_gate". + prefix = f"{ckpt_prefix}." + for name, loaded_weight in weights: + if prefix not in name: + yield name, loaded_weight + continue + # gate/up split on the output dim; down_proj on the input dim. + split_dim = 1 if ("down_proj.weight" in name and loaded_weight.ndim > 1) else 0 + total = loaded_weight.shape[split_dim] + if total % n_shared_experts != 0: + raise ValueError( + f"FSE shared-expert weight {name!r} has size {total} along axis " + f"{split_dim}, not divisible by n_shared_experts={n_shared_experts}." + ) + chunk = total // n_shared_experts + for j in range(n_shared_experts): + sl = slice(j * chunk, (j + 1) * chunk) + if loaded_weight.ndim == 1: + chunk_weight = loaded_weight[sl] + elif split_dim == 0: + chunk_weight = loaded_weight[sl, :] + else: + chunk_weight = loaded_weight[:, sl] + yield ( + name.replace(prefix, f"mlp.experts.{n_routed_experts + j}."), + chunk_weight, + ) + + +def get_spec_layer_idx_from_weight_name( + config: "ModelConfig", weight_name: str +) -> int | None: + """Return the MTP layer index a weight belongs to, or None. + + Args: + config: The model config; must expose `num_hidden_layers` and, for MTP + checkpoints, `num_nextn_predict_layers`. + weight_name: Checkpoint weight name to classify. + + Returns: + The absolute layer index for an MTP-layer weight, else None. + """ + if not (n := getattr(config, "num_nextn_predict_layers", 0)): + return None + base = config.num_hidden_layers + for i in range(n): + if weight_name.startswith((f"model.layers.{base + i}.", f"layers.{base + i}.")): + return base + i + return None + + +def skip_spec_layers( + weights: Iterable[tuple[str, torch.Tensor]], config: "ModelConfig" +) -> Iterable[tuple[str, torch.Tensor]]: + """Drop MTP spec-layer weights (loaded by the MTP head, not the base model). + + Args: + weights: Iterable of `(name, tensor)` checkpoint pairs. + config: The model config, passed to `get_spec_layer_idx_from_weight_name`. + + Yields: + `(name, tensor)` pairs whose weight is not an MTP-layer weight. + """ + return ( + (name, w) + for name, w in weights + if get_spec_layer_idx_from_weight_name(config, name) is None + ) + + def init_vllm_registered_model( vllm_config: VllmConfig, *, From c7ce03bcbd380d0e94490abb111fe48861c16343 Mon Sep 17 00:00:00 2001 From: Michael Goin Date: Sat, 18 Jul 2026 08:59:33 -0400 Subject: [PATCH 21/51] [Bugfix] Bump tml-fa4 for cutlass-dsl 4.6 API compatibility (#48988) Signed-off-by: mgoin --- cmake/external_projects/tml_fa4.cmake | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/cmake/external_projects/tml_fa4.cmake b/cmake/external_projects/tml_fa4.cmake index 59e2e241c92..a965295e2bb 100644 --- a/cmake/external_projects/tml_fa4.cmake +++ b/cmake/external_projects/tml_fa4.cmake @@ -14,7 +14,7 @@ else() FetchContent_Declare( tml_fa4 GIT_REPOSITORY https://github.com/vllm-project/tml-fa4.git - GIT_TAG 13374f0c855acc1add1bf30444bd67aebbc24a8e + GIT_TAG b206834606ed5b5f21f8eed6b0683f528ea9cf7d GIT_PROGRESS TRUE CONFIGURE_COMMAND "" BUILD_COMMAND "") From 29c0ec4d63d1869f54a9fbcdf082f77534d9211a Mon Sep 17 00:00:00 2001 From: "Kevin H. Luu" Date: Sat, 18 Jul 2026 08:43:49 -0700 Subject: [PATCH 22/51] [ci] Move 3 entrypoints tests to h200_35gb queue (#43164) Signed-off-by: Simon Mo Signed-off-by: Simon Mo Co-authored-by: Simon Mo Co-authored-by: OpenAI Codex --- .buildkite/test_areas/entrypoints.yaml | 3 +++ tests/entrypoints/openai/responses/test_parsable_context.py | 2 +- 2 files changed, 4 insertions(+), 1 deletion(-) diff --git a/.buildkite/test_areas/entrypoints.yaml b/.buildkite/test_areas/entrypoints.yaml index 2db4c5ad5a2..499abcc6c82 100644 --- a/.buildkite/test_areas/entrypoints.yaml +++ b/.buildkite/test_areas/entrypoints.yaml @@ -3,6 +3,7 @@ depends_on: - image-build steps: - label: Entrypoints Unit Tests + device: h200_35gb key: entrypoints-unit-tests timeout_in_minutes: 25 working_dir: "/vllm-workspace/tests" @@ -15,6 +16,7 @@ steps: - pytest -v -s entrypoints/weight_transfer - label: Entrypoints Integration (LLM) + device: h200_35gb key: entrypoints-integration-llm timeout_in_minutes: 60 working_dir: "/vllm-workspace/tests" @@ -114,6 +116,7 @@ steps: - image-build-amd - label: Entrypoints Integration (Responses API) + device: h200_35gb key: entrypoints-integration-responses-api timeout_in_minutes: 50 working_dir: "/vllm-workspace/tests" diff --git a/tests/entrypoints/openai/responses/test_parsable_context.py b/tests/entrypoints/openai/responses/test_parsable_context.py index 8ff3a1eeae8..3a17da61276 100644 --- a/tests/entrypoints/openai/responses/test_parsable_context.py +++ b/tests/entrypoints/openai/responses/test_parsable_context.py @@ -68,7 +68,7 @@ async def client(server): async def test_basic(client: OpenAI, model_name: str): response = await client.responses.create( model=model_name, - input="What is 123 * 456?", + input="What is 123 * 456? Answer with only the number.", temperature=0.0, ) assert response is not None From e94243893dd30256f58644ad4ecf779be757dff8 Mon Sep 17 00:00:00 2001 From: frida-andersson Date: Sat, 18 Jul 2026 18:39:37 +0200 Subject: [PATCH 23/51] [ROCm][DSv3.2][Perf] Cap sparse MLA decode KV-splits with a work-per-split heuristic (#46832) Signed-off-by: Frida Andersson --- .../backends/mla/rocm_aiter_mla_sparse.py | 21 +++++++++++++++++++ 1 file changed, 21 insertions(+) diff --git a/vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py b/vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py index bda63e1ae2d..2bbaaad7654 100644 --- a/vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py @@ -369,6 +369,8 @@ class ROCMAiterMLASparseMetadataBuilder( self.num_heads = self.model_config.get_num_attention_heads(parallel_config) self.mla_dims = get_mla_dims(self.model_config) self.topk_tokens = vllm_config.model_config.hf_config.index_topk + # Bounds the KV-split heuristic (see `_sparse_decode_max_split`). + self._num_compute_units = current_platform.num_compute_units() self.max_model_len_tensor = torch.tensor( [self.model_config.max_model_len], device=device, dtype=torch.int32 ) @@ -464,6 +466,21 @@ class ROCMAiterMLASparseMetadataBuilder( self._prev_indices_extent: int = 0 self._prev_metadata_key: tuple | None = None + def _sparse_decode_max_split(self, max_seq_len: int) -> int: + """Cap ``max_split_per_batch`` for the aiter sparse-MLA decode reduce. + + The reduce only covers the selected tokens per row (``<= topk_tokens``), + so aiter's default (``-1`` => split across every CU) over-fragments it. + Mirror ``triton_mla.py``: aim for a minimum work per split, round to a + power of two, and cap by the CU count. Numerics are unchanged. + """ + effective_len = min(max_seq_len, self.topk_tokens) + min_work_per_split = 128 + ideal_splits = triton.next_power_of_2( + max(1, effective_len // min_work_per_split) + ) + return min(ideal_splits, self._num_compute_units) + def build( self, common_prefix_len: int, @@ -547,6 +564,9 @@ class ROCMAiterMLASparseMetadataBuilder( if metadata_key != self._prev_metadata_key: from aiter import get_mla_metadata_v1 + max_split_per_batch = self._sparse_decode_max_split( + int(common_attn_metadata.max_seq_len) + ) get_mla_metadata_v1( qo_indptr, paged_kv_indptr, @@ -565,6 +585,7 @@ class ROCMAiterMLASparseMetadataBuilder( max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=True, + max_split_per_batch=max_split_per_batch, ) # The persistent metadata buffers are read by graph replay. Order # the async metadata write before the graph-captured decode kernel. From a287eb163fb6f8f007a4a78411fb54c8dde64cc7 Mon Sep 17 00:00:00 2001 From: yzong-rh Date: Sat, 18 Jul 2026 13:04:35 -0400 Subject: [PATCH 24/51] [Front-end] [Messages] Populate `num_cache_creation_tokens` (#48535) Signed-off-by: Yifan Zong --- .../engine-core-client/src/protocol/stats.rs | 3 + rust/src/llm/src/request_metrics.rs | 1 + .../test_anthropic_messages_conversion.py | 85 ++++++------------- tests/entrypoints/anthropic/test_messages.py | 52 ++++++++++++ .../chat_completion/test_serving_chat.py | 17 ++-- tests/v1/core/test_async_scheduler.py | 1 + tests/v1/core/test_scheduler.py | 1 + vllm/entrypoints/anthropic/serving.py | 80 +++++++---------- .../openai/chat_completion/serving.py | 12 ++- vllm/entrypoints/openai/engine/protocol.py | 1 + vllm/outputs.py | 7 +- vllm/v1/core/kv_cache_manager.py | 29 +++++++ vllm/v1/core/sched/scheduler.py | 26 ++++-- vllm/v1/engine/output_processor.py | 5 ++ vllm/v1/metrics/stats.py | 8 ++ 15 files changed, 201 insertions(+), 127 deletions(-) diff --git a/rust/src/engine-core-client/src/protocol/stats.rs b/rust/src/engine-core-client/src/protocol/stats.rs index d1f2dd3d353..f02da161ef4 100644 --- a/rust/src/engine-core-client/src/protocol/stats.rs +++ b/rust/src/engine-core-client/src/protocol/stats.rs @@ -103,6 +103,9 @@ pub struct PrefillStats { /// Tokens to be prefilled from external KV transfer. #[serde(default)] pub num_external_cached_tokens: u32, + /// Prompt tokens newly admitted into the local prefix cache. + #[serde(default)] + pub num_cache_creation_tokens: u32, } /// Stats for debugging the metrics calculation. diff --git a/rust/src/llm/src/request_metrics.rs b/rust/src/llm/src/request_metrics.rs index 2cc9c187a87..280fa0f10fb 100644 --- a/rust/src/llm/src/request_metrics.rs +++ b/rust/src/llm/src/request_metrics.rs @@ -375,6 +375,7 @@ mod tests { num_cached_tokens: 4, num_local_cached_tokens: 4, num_external_cached_tokens: 0, + ..Default::default() }), ..Default::default() }, diff --git a/tests/entrypoints/anthropic/test_anthropic_messages_conversion.py b/tests/entrypoints/anthropic/test_anthropic_messages_conversion.py index f89d12553b9..94f5bb6068e 100644 --- a/tests/entrypoints/anthropic/test_anthropic_messages_conversion.py +++ b/tests/entrypoints/anthropic/test_anthropic_messages_conversion.py @@ -23,7 +23,6 @@ from vllm.entrypoints.anthropic.protocol import ( from vllm.entrypoints.anthropic.serving import ( AnthropicServingMessages, _build_anthropic_usage, - _get_cached_tokens, ) from vllm.entrypoints.openai.chat_completion.protocol import ( ChatCompletionResponse, @@ -668,42 +667,6 @@ class TestThinkingBlockConversion: # ====================================================================== -class TestGetCachedTokens: - """Tests for _get_cached_tokens helper.""" - - def test_none_usage(self): - assert _get_cached_tokens(None) is None - - def test_no_prompt_tokens_details(self): - usage = UsageInfo(prompt_tokens=100, completion_tokens=10) - assert _get_cached_tokens(usage) is None - - def test_cached_tokens_present(self): - usage = UsageInfo( - prompt_tokens=100, - completion_tokens=10, - prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=80), - ) - assert _get_cached_tokens(usage) == 80 - - def test_cached_tokens_zero(self): - """Zero cached tokens should return 0, not None.""" - usage = UsageInfo( - prompt_tokens=100, - completion_tokens=10, - prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=0), - ) - assert _get_cached_tokens(usage) == 0 - - def test_cached_tokens_none_in_details(self): - usage = UsageInfo( - prompt_tokens=100, - completion_tokens=10, - prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=None), - ) - assert _get_cached_tokens(usage) is None - - class TestBuildAnthropicUsage: """Tests for _build_anthropic_usage helper. @@ -711,36 +674,32 @@ class TestBuildAnthropicUsage: vLLM's prompt_tokens is the total. """ - def test_no_cache_info(self): - """When cache info is unavailable, return raw prompt_tokens.""" - result = _build_anthropic_usage(100, 10, None) - assert result.input_tokens == 100 - assert result.output_tokens == 10 - assert result.cache_read_input_tokens is None - assert result.cache_creation_input_tokens is None - def test_cache_hit(self): """When cache is hit, input_tokens excludes cached tokens.""" usage = UsageInfo( prompt_tokens=100, completion_tokens=10, - prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=80), + prompt_tokens_details=PromptTokenUsageInfo( + cached_tokens=80, created_cache_tokens=10 + ), ) - result = _build_anthropic_usage(100, 10, usage) - assert result.input_tokens == 20 # 100 - 80 + result = _build_anthropic_usage(usage) + assert result.input_tokens == 10 # 100 - 80 - 10 assert result.output_tokens == 10 assert result.cache_read_input_tokens == 80 - assert result.cache_creation_input_tokens == 0 + assert result.cache_creation_input_tokens == 10 def test_zero_cached_tokens(self): """Zero cached tokens should still set cache_creation to 0.""" usage = UsageInfo( prompt_tokens=100, completion_tokens=10, - prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=0), + prompt_tokens_details=PromptTokenUsageInfo( + cached_tokens=0, created_cache_tokens=0 + ), ) - result = _build_anthropic_usage(100, 10, usage) - assert result.input_tokens == 100 # 100 - 0 + result = _build_anthropic_usage(usage) + assert result.input_tokens == 100 # 100 - 0 - 0 assert result.cache_read_input_tokens == 0 assert result.cache_creation_input_tokens == 0 @@ -749,9 +708,11 @@ class TestBuildAnthropicUsage: usage = UsageInfo( prompt_tokens=100, completion_tokens=10, - prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=100), + prompt_tokens_details=PromptTokenUsageInfo( + cached_tokens=100, created_cache_tokens=0 + ), ) - result = _build_anthropic_usage(100, 10, usage) + result = _build_anthropic_usage(usage) assert result.input_tokens == 0 assert result.cache_read_input_tokens == 100 assert result.cache_creation_input_tokens == 0 @@ -759,7 +720,7 @@ class TestBuildAnthropicUsage: def test_no_prompt_tokens_details(self): """UsageInfo without prompt_tokens_details returns no cache info.""" usage = UsageInfo(prompt_tokens=100, completion_tokens=10) - result = _build_anthropic_usage(100, 10, usage) + result = _build_anthropic_usage(usage) assert result.input_tokens == 100 assert result.cache_read_input_tokens is None assert result.cache_creation_input_tokens is None @@ -1241,7 +1202,9 @@ class TestStreamingCacheUsageSemantics: prompt_tokens=100, completion_tokens=5, total_tokens=105, - prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=80), + prompt_tokens_details=PromptTokenUsageInfo( + cached_tokens=80, created_cache_tokens=10 + ), ), ) yield "data: [DONE]" @@ -1263,9 +1226,9 @@ class TestStreamingCacheUsageSemantics: delta_usage = next( data["usage"] for ev, data in events if ev == "message_delta" ) - assert delta_usage["input_tokens"] == 20 # 100 - 80 + assert delta_usage["input_tokens"] == 10 # 100 - 80 - 10 assert delta_usage["cache_read_input_tokens"] == 80 - assert delta_usage["cache_creation_input_tokens"] == 0 + assert delta_usage["cache_creation_input_tokens"] == 10 @pytest.mark.asyncio async def test_streaming_no_cache_hit(self): @@ -1284,7 +1247,9 @@ class TestStreamingCacheUsageSemantics: prompt_tokens=50, completion_tokens=5, total_tokens=55, - prompt_tokens_details=PromptTokenUsageInfo(cached_tokens=0), + prompt_tokens_details=PromptTokenUsageInfo( + cached_tokens=0, created_cache_tokens=0 + ), ), ) yield "data: [DONE]" @@ -1302,7 +1267,7 @@ class TestStreamingCacheUsageSemantics: assert start_usage["input_tokens"] == 50 assert "cache_read_input_tokens" not in start_usage assert "cache_creation_input_tokens" not in start_usage - assert delta_usage["input_tokens"] == 50 # 50 - 0 + assert delta_usage["input_tokens"] == 50 # 50 - 0 - 0 assert delta_usage["cache_read_input_tokens"] == 0 assert delta_usage["cache_creation_input_tokens"] == 0 diff --git a/tests/entrypoints/anthropic/test_messages.py b/tests/entrypoints/anthropic/test_messages.py index c1f6858e83d..9f5cc2c451d 100644 --- a/tests/entrypoints/anthropic/test_messages.py +++ b/tests/entrypoints/anthropic/test_messages.py @@ -18,6 +18,7 @@ def server(): "--max-model-len", "2048", "--enforce-eager", + "--enable-prompt-tokens-details", "--enable-auto-tool-choice", "--tool-call-parser", "hermes", @@ -191,3 +192,54 @@ async def test_anthropic_structured_output(client: anthropic.AsyncAnthropic): json_obj = json.loads(response.content[0].text) for key in ["name", "email", "plan_interest", "demo_requested"]: assert key in json_obj, f"Missing key in output: {key}" + + +@pytest.mark.asyncio +async def test_anthropic_streaming_cache_usage(client: anthropic.AsyncAnthropic): + async def get_stream_usage(resp): + prompt_tokens = None + usage = None + async for chunk in resp: + if ( + chunk.type == "message_start" + and chunk.message is not None + and chunk.message.usage is not None + ): + prompt_tokens = chunk.message.usage.input_tokens + elif chunk.type == "message_delta" and chunk.usage is not None: + usage = chunk.usage + + assert usage is not None + assert usage.input_tokens >= 0 + assert usage.output_tokens >= 0 + cache_created = usage.cache_creation_input_tokens + cache_read = usage.cache_read_input_tokens + assert cache_read is not None + assert cache_created is not None + assert cache_created >= 0 + assert cache_read >= 0 + assert prompt_tokens == usage.input_tokens + cache_created + cache_read + return usage + + request = dict( + model="claude-3-7-sonnet-latest", + max_tokens=1, + temperature=0.0, + messages=[ + { + "role": "user", + "content": "Cache coverage sentinel. " * 256 + + "Answer with exactly one word: ok.", + } + ], + stream=True, + ) + + cold_usage = await get_stream_usage(await client.messages.create(**request)) + assert cold_usage.cache_read_input_tokens == 0 + assert cold_usage.cache_creation_input_tokens is not None + assert cold_usage.cache_creation_input_tokens > 0 + + warm_usage = await get_stream_usage(await client.messages.create(**request)) + assert warm_usage.cache_read_input_tokens is not None + assert warm_usage.cache_read_input_tokens > 0 diff --git a/tests/entrypoints/openai/chat_completion/test_serving_chat.py b/tests/entrypoints/openai/chat_completion/test_serving_chat.py index 25a9451bc2b..2126acfe027 100644 --- a/tests/entrypoints/openai/chat_completion/test_serving_chat.py +++ b/tests/entrypoints/openai/chat_completion/test_serving_chat.py @@ -831,16 +831,23 @@ def test_mm_prompt_tokens_details(): assert counts == {"image": 600, "video": 1200} # Gated off, or nothing to report -> no details. - assert _make_prompt_tokens_details(False, 5, counts) is None - assert _make_prompt_tokens_details(True, None, None) is None + assert _make_prompt_tokens_details(False, 5, 0, counts) is None + assert _make_prompt_tokens_details(True, None, None, None) is None # Zero cached_tokens is still reported (not None), matching the cached-only # behavior; multimodal counts ride alongside even when cached_tokens is None. - assert _make_prompt_tokens_details(True, 0, None).cached_tokens == 0 - details = _make_prompt_tokens_details(True, None, counts) + details = _make_prompt_tokens_details(True, 0, 0, None) + assert details.cached_tokens == 0 + assert details.created_cache_tokens == 0 + assert details.multimodal_tokens is None + details = _make_prompt_tokens_details(True, None, None, counts) assert details.cached_tokens is None + assert details.created_cache_tokens is None + assert details.multimodal_tokens == {"image": 600, "video": 1200} + details = _make_prompt_tokens_details(True, 3, 0, counts) + assert details.cached_tokens == 3 + assert details.created_cache_tokens == 0 assert details.multimodal_tokens == {"image": 600, "video": 1200} - assert _make_prompt_tokens_details(True, 3, counts).cached_tokens == 3 @pytest.mark.asyncio diff --git a/tests/v1/core/test_async_scheduler.py b/tests/v1/core/test_async_scheduler.py index 3997b85f2d1..cd3efa8ee64 100644 --- a/tests/v1/core/test_async_scheduler.py +++ b/tests/v1/core/test_async_scheduler.py @@ -280,6 +280,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance(): scheduler.waiting = Mock() scheduler.kv_cache_manager = Mock() scheduler.kv_cache_manager.take_events.return_value = None + scheduler.kv_cache_manager.estimate_cached_tokens.return_value = 0 scheduler.kv_event_publisher = Mock() scheduler.finished_req_ids = set() scheduler.finished_req_ids_dict = None diff --git a/tests/v1/core/test_scheduler.py b/tests/v1/core/test_scheduler.py index da1e0c5e76c..408de409280 100644 --- a/tests/v1/core/test_scheduler.py +++ b/tests/v1/core/test_scheduler.py @@ -3036,6 +3036,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance(): scheduler.waiting = Mock() scheduler.kv_cache_manager = Mock() scheduler.kv_cache_manager.take_events.return_value = None + scheduler.kv_cache_manager.estimate_cached_tokens.return_value = 0 scheduler.kv_event_publisher = Mock() scheduler.finished_req_ids = set() scheduler.finished_req_ids_dict = None diff --git a/vllm/entrypoints/anthropic/serving.py b/vllm/entrypoints/anthropic/serving.py index d61a917780d..d516a7530b9 100644 --- a/vllm/entrypoints/anthropic/serving.py +++ b/vllm/entrypoints/anthropic/serving.py @@ -53,47 +53,43 @@ from vllm.renderers.online_renderer import OnlineRenderer logger = logging.getLogger(__name__) -def _get_cached_tokens(usage: UsageInfo | None) -> int | None: - """Extract cached token count from OpenAI UsageInfo.""" - if usage is None or usage.prompt_tokens_details is None: - return None - return usage.prompt_tokens_details.cached_tokens - - def _build_anthropic_usage( - prompt_tokens: int, - completion_tokens: int | None, usage: UsageInfo | None, ) -> AnthropicUsage: - """Build an AnthropicUsage from OpenAI-style token counts. + """Build an AnthropicUsage from UsageInfo. Anthropic defines ``total_input == input_tokens + cache_read + cache_creation``. vLLM's ``prompt_tokens`` is the total, so - ``input_tokens = prompt_tokens - cached_tokens``. + ``input_tokens = prompt_tokens - cache_read - cache_creation``. - OpenAI usage only exposes ``cached_tokens`` (hits); there is no - cache-creation analog, so ``cache_creation_input_tokens`` is ``0`` - when cache info is present. When cache info is absent (e.g. - ``--enable-prompt-tokens-details`` off, or a streaming chunk that - hasn't carried it yet), cache fields are left **unset** so - ``exclude_unset=True`` serialization omits them entirely. + Cache fields are taken from ``UsageInfo.prompt_tokens_details``. + When cache info is absent (e.g. ``--enable-prompt-tokens-details`` + off, or a streaming chunk that hasn't carried it yet), cache fields + are left **unset** so ``exclude_unset=True`` serialization omits them + entirely. ``completion_tokens`` follows ``UsageInfo`` and may be ``None`` on intermediate stream chunks; we coerce to ``0`` for the wire format. """ - output_tokens = completion_tokens or 0 - cached = _get_cached_tokens(usage) - if cached is not None: - return AnthropicUsage( - input_tokens=prompt_tokens - cached, - output_tokens=output_tokens, - cache_read_input_tokens=cached, - cache_creation_input_tokens=0, - ) - return AnthropicUsage( - input_tokens=prompt_tokens, - output_tokens=output_tokens, - ) + kwargs = {} + if usage is None: + kwargs["input_tokens"] = 0 + kwargs["output_tokens"] = 0 + else: + kwargs["output_tokens"] = usage.completion_tokens or 0 + input_tokens = usage.prompt_tokens + + if (details := usage.prompt_tokens_details) is not None: + if (cache_read := details.cached_tokens) is not None: + input_tokens -= cache_read + kwargs["cache_read_input_tokens"] = cache_read + + if (cache_creation := details.created_cache_tokens) is not None: + input_tokens -= cache_creation + kwargs["cache_creation_input_tokens"] = cache_creation + + kwargs["input_tokens"] = max(0, input_tokens) + return AnthropicUsage(**kwargs) def wrap_data_with_event(data: str, event: str): @@ -625,11 +621,7 @@ class AnthropicServingMessages(OpenAIServingChat): id=generator.id, content=[], model=generator.model, - usage=_build_anthropic_usage( - generator.usage.prompt_tokens, - generator.usage.completion_tokens, - generator.usage, - ), + usage=_build_anthropic_usage(generator.usage), kv_transfer_params=generator.kv_transfer_params, ec_transfer_params=generator.ec_transfer_params, ) @@ -816,13 +808,7 @@ class AnthropicServingMessages(OpenAIServingChat): model=origin_chunk.model, stop_reason=None, stop_sequence=None, - usage=_build_anthropic_usage( - origin_chunk.usage.prompt_tokens - if origin_chunk.usage - else 0, - 0, - origin_chunk.usage, - ), + usage=_build_anthropic_usage(origin_chunk.usage), ), ) first_item = False @@ -840,15 +826,7 @@ class AnthropicServingMessages(OpenAIServingChat): chunk = AnthropicStreamEvent( type="message_delta", delta=AnthropicDelta(stop_reason=stop_reason), - usage=_build_anthropic_usage( - origin_chunk.usage.prompt_tokens - if origin_chunk.usage - else 0, - origin_chunk.usage.completion_tokens - if origin_chunk.usage - else 0, - origin_chunk.usage, - ), + usage=_build_anthropic_usage(origin_chunk.usage), ) data = chunk.model_dump_json(exclude_unset=True) yield wrap_data_with_event(data, "message_delta") diff --git a/vllm/entrypoints/openai/chat_completion/serving.py b/vllm/entrypoints/openai/chat_completion/serving.py index 5d821ba5841..4c05db0b8a4 100644 --- a/vllm/entrypoints/openai/chat_completion/serving.py +++ b/vllm/entrypoints/openai/chat_completion/serving.py @@ -90,15 +90,21 @@ def _get_mm_token_counts(engine_input: EngineInput) -> dict[str, int]: def _make_prompt_tokens_details( enable_prompt_tokens_details: bool, num_cached_tokens: int | None, + num_cache_creation_tokens: int | None, mm_token_counts: dict[str, int] | None, ) -> PromptTokenUsageInfo | None: """Build ``prompt_tokens_details`` from cached + multimodal token counts.""" if not enable_prompt_tokens_details: return None - if num_cached_tokens is None and not mm_token_counts: + if ( + num_cached_tokens is None + and num_cache_creation_tokens is None + and not mm_token_counts + ): return None return PromptTokenUsageInfo( cached_tokens=num_cached_tokens, + created_cache_tokens=num_cache_creation_tokens, multimodal_tokens=mm_token_counts or None, ) @@ -427,6 +433,7 @@ class OpenAIServingChat(GenerateBaseServing): finish_reason_sent = [False] * num_choices num_prompt_tokens = 0 num_cached_tokens = None + num_cache_creation_tokens = None tools_streamed = [False] * num_choices if isinstance(request.tool_choice, ChatCompletionNamedToolChoiceParam): @@ -479,6 +486,7 @@ class OpenAIServingChat(GenerateBaseServing): # response (by the try...catch). if first_iteration: num_cached_tokens = res.num_cached_tokens + num_cache_creation_tokens = res.num_cache_creation_tokens # Send first response for each request.n (index) with # the role role = self.get_chat_request_role(request) @@ -756,6 +764,7 @@ class OpenAIServingChat(GenerateBaseServing): final_usage.prompt_tokens_details = _make_prompt_tokens_details( self.enable_prompt_tokens_details, num_cached_tokens, + num_cache_creation_tokens, mm_token_counts, ) @@ -1030,6 +1039,7 @@ class OpenAIServingChat(GenerateBaseServing): usage.prompt_tokens_details = _make_prompt_tokens_details( self.enable_prompt_tokens_details, final_res.num_cached_tokens, + final_res.num_cache_creation_tokens, mm_token_counts, ) diff --git a/vllm/entrypoints/openai/engine/protocol.py b/vllm/entrypoints/openai/engine/protocol.py index 2c32fcf20c6..7d901d19333 100644 --- a/vllm/entrypoints/openai/engine/protocol.py +++ b/vllm/entrypoints/openai/engine/protocol.py @@ -101,6 +101,7 @@ class ModelList(OpenAIBaseModel): class PromptTokenUsageInfo(OpenAIBaseModel): cached_tokens: int | None = None + created_cache_tokens: int | None = None multimodal_tokens: dict[str, int] | None = None """Prompt tokens contributed by each input modality, keyed by modality name (e.g. `image`, `audio`, `video`). A breakdown of the multimodal diff --git a/vllm/outputs.py b/vllm/outputs.py index 5a0f0dec805..feee2a95279 100644 --- a/vllm/outputs.py +++ b/vllm/outputs.py @@ -103,6 +103,8 @@ class RequestOutput: encoder_prompt_token_ids: The token IDs of the encoder prompt. None if decoder-only. num_cached_tokens: The number of tokens with prefix cache hit. + num_cache_creation_tokens: Prompt tokens currently counted as local + prefix-cache writes for this request. kv_transfer_params: The params for remote K/V transfer. ec_transfer_params: The params for remote encoder-cache transfer. """ @@ -120,6 +122,7 @@ class RequestOutput: encoder_prompt: str | None = None, encoder_prompt_token_ids: list[int] | None = None, num_cached_tokens: int | None = None, + num_cache_creation_tokens: int | None = None, *, kv_transfer_params: dict[str, Any] | None = None, ec_transfer_params: dict[str, Any] | None = None, @@ -142,6 +145,7 @@ class RequestOutput: self.encoder_prompt = encoder_prompt self.encoder_prompt_token_ids = encoder_prompt_token_ids self.num_cached_tokens = num_cached_tokens + self.num_cache_creation_tokens = num_cache_creation_tokens self.kv_transfer_params = kv_transfer_params self.ec_transfer_params = ec_transfer_params @@ -188,7 +192,8 @@ class RequestOutput: f"finished={self.finished}, " f"metrics={self.metrics}, " f"lora_request={self.lora_request}, " - f"num_cached_tokens={self.num_cached_tokens})" + f"num_cached_tokens={self.num_cached_tokens}, " + f"num_cache_creation_tokens={self.num_cache_creation_tokens})" ) diff --git a/vllm/v1/core/kv_cache_manager.py b/vllm/v1/core/kv_cache_manager.py index 3d3d8c1573a..568a2774b2d 100644 --- a/vllm/v1/core/kv_cache_manager.py +++ b/vllm/v1/core/kv_cache_manager.py @@ -657,6 +657,35 @@ class KVCacheManager: clipped_block_ids.append(ids[:num_valid_blocks]) return tuple(clipped_block_ids) + def estimate_cached_tokens(self, request: Request) -> int: + """Estimate the number of tokens cached by the request.""" + cached_tokens: int | None = None + for group, blocks in zip( + self.kv_cache_config.kv_cache_groups, + self.get_blocks(request.request_id).blocks, + ): + if isinstance( + group.kv_cache_spec, + (CrossAttentionSpec, EncoderOnlyAttentionSpec), + ): + # Cross-attention and encoder-only groups are not prefix cached. + continue + + group_cached_tokens = 0 + for block in blocks: + group_cached_tokens = max( + group_cached_tokens, + block.block_hash_num_tokens or 0, + ) + + cached_tokens = ( + group_cached_tokens + if cached_tokens is None + else min(cached_tokens, group_cached_tokens) + ) + + return cached_tokens or 0 + def cache_blocks(self, request: Request, num_computed_tokens: int) -> None: """Cache the blocks for the request, if enabled. diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index d6f3c3ad0e3..de076e230e4 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -799,7 +799,10 @@ class Scheduler(SchedulerInterface): continue # Track first scheduled prefill, not post-preemption repeat prefills - if request.prefill_stats is not None: + if ( + request.prefill_stats is not None + and request.num_preemptions <= 0 + ): assert num_computed_tokens <= request.num_prompt_tokens request.prefill_stats.set( num_prompt_tokens=request.num_prompt_tokens, @@ -1708,6 +1711,7 @@ class Scheduler(SchedulerInterface): pooler_output = pooler_outputs[req_index] if pooler_outputs else None kv_transfer_params = None ec_transfer_params = None + prefill_stats = None status_before_stop = request.status num_output_tokens_before = len(request._output_token_ids) @@ -1787,6 +1791,16 @@ class Scheduler(SchedulerInterface): # Normal decode / re-prefill: token(s) at the END. routed_experts = routing_data[end - len(new_token_ids) : end] + should_emit_output = bool( + new_token_ids or pooler_output is not None or stopped + ) + if should_emit_output: + prefill_stats = request.take_prefill_stats() + if prefill_stats is not None: + prefill_stats.finalize( + self.kv_cache_manager.estimate_cached_tokens(request) + ) + finish_reason = None if stopped: # Capture finish_reason BEFORE _handle_stopped_request, which may @@ -1814,13 +1828,7 @@ class Scheduler(SchedulerInterface): # Get prompt logprobs for this request. prompt_logprobs_tensors = prompt_logprobs_dict.get(req_id) - if ( - new_token_ids - or pooler_output is not None - or kv_transfer_params - or ec_transfer_params - or stopped - ): + if should_emit_output: # Add EngineCoreOutput for this Request. outputs[request.client_index].append( EngineCoreOutput( @@ -1832,7 +1840,7 @@ class Scheduler(SchedulerInterface): pooling_output=pooler_output, stop_reason=request.stop_reason, events=request.take_events(), - prefill_stats=request.take_prefill_stats(), + prefill_stats=prefill_stats, kv_transfer_params=kv_transfer_params, ec_transfer_params=ec_transfer_params, trace_headers=request.trace_headers, diff --git a/vllm/v1/engine/output_processor.py b/vllm/v1/engine/output_processor.py index b676c3cd2d3..0d7d5fe18c9 100644 --- a/vllm/v1/engine/output_processor.py +++ b/vllm/v1/engine/output_processor.py @@ -172,6 +172,7 @@ class RequestState: self.is_prefilling = True self.queue = queue self.num_cached_tokens = 0 + self.num_cache_creation_tokens = 0 self.stats = RequestStateStats(arrival_time=arrival_time) if log_stats else None @@ -377,6 +378,7 @@ class RequestState: kv_transfer_params=kv_transfer_params, ec_transfer_params=ec_transfer_params, num_cached_tokens=self.num_cached_tokens, + num_cache_creation_tokens=self.num_cache_creation_tokens, metrics=self.stats, ) @@ -639,6 +641,9 @@ class OutputProcessor: req_state.num_cached_tokens = ( engine_core_output.prefill_stats.num_cached_tokens ) + req_state.num_cache_creation_tokens = ( + engine_core_output.prefill_stats.num_cache_creation_tokens + ) req_state.is_prefilling = False if pooling_output is None: diff --git a/vllm/v1/metrics/stats.py b/vllm/v1/metrics/stats.py index 20bb3e1caa6..3956f7e4413 100644 --- a/vllm/v1/metrics/stats.py +++ b/vllm/v1/metrics/stats.py @@ -265,6 +265,7 @@ class PrefillStats: num_cached_tokens: Tokens to be prefilled without actual compute work. num_local_cached_tokens: Tokens to be prefilled from local prefix cache. num_external_cached_tokens: Tokens to be prefilled from external KV transfer. + num_cache_creation_tokens: Tokens computed and written to the prefix cache. """ num_prompt_tokens: int = 0 @@ -272,6 +273,7 @@ class PrefillStats: num_cached_tokens: int = 0 num_local_cached_tokens: int = 0 num_external_cached_tokens: int = 0 + num_cache_creation_tokens: int = 0 def set( self, @@ -288,6 +290,12 @@ class PrefillStats: self.num_local_cached_tokens = num_local_cached_tokens self.num_external_cached_tokens = num_external_cached_tokens + def finalize(self, num_cached_tokens: int) -> None: + assert num_cached_tokens >= 0 + self.num_cache_creation_tokens = max( + 0, min(num_cached_tokens, self.num_prompt_tokens) - self.num_cached_tokens + ) + @dataclass class PromptTokenStats: From 7c2acd38b72de4a4177f2a4c2b80b216da7d55b0 Mon Sep 17 00:00:00 2001 From: SYLAR <125541396+lishunyang12@users.noreply.github.com> Date: Sun, 19 Jul 2026 01:29:11 +0800 Subject: [PATCH 25/51] [Bugfix] Qwen3-VL/Qwen-Omni: honor max_pixels/min_pixels for video prompts (#49015) --- .../models/qwen2_5_omni_thinker.py | 17 +++++++++++++++++ vllm/model_executor/models/qwen3_vl.py | 13 +++++++++++++ 2 files changed, 30 insertions(+) diff --git a/vllm/model_executor/models/qwen2_5_omni_thinker.py b/vllm/model_executor/models/qwen2_5_omni_thinker.py index 00382cf0a87..a0e9b84f344 100644 --- a/vllm/model_executor/models/qwen2_5_omni_thinker.py +++ b/vllm/model_executor/models/qwen2_5_omni_thinker.py @@ -509,6 +509,23 @@ class Qwen2_5OmniThinkerMultiModalProcessor( **mm_kwargs, ) + merged = self.info.ctx.get_merged_mm_kwargs(mm_kwargs) + if mm_data.get("videos") and ( + merged.keys() & {"size", "min_pixels", "max_pixels"} + ): + mm_kwargs = dict(mm_kwargs) + video_size = dict(self.info.get_hf_processor().video_processor.size) + size_override = merged.get("size") + if size_override is not None: + video_size = video_size | size_override + min_pixels = merged.get("min_pixels") + if min_pixels is not None: + video_size["shortest_edge"] = min_pixels + max_pixels = merged.get("max_pixels") + if max_pixels is not None: + video_size["longest_edge"] = max_pixels + mm_kwargs["size"] = video_size + hf_inputs = super()._call_hf_processor( prompt=prompt, mm_data=mm_data, diff --git a/vllm/model_executor/models/qwen3_vl.py b/vllm/model_executor/models/qwen3_vl.py index 5252afd78e4..f86560e5f4e 100644 --- a/vllm/model_executor/models/qwen3_vl.py +++ b/vllm/model_executor/models/qwen3_vl.py @@ -1272,6 +1272,19 @@ class Qwen3VLMultiModalProcessor(BaseMultiModalProcessor[Qwen3VLProcessingInfo]) # NOTE: a copy of is created to update do_sample_frames, # otherwise mm_hash for the object will be incorrect. video_mm_kwargs = dict(**mm_kwargs) + merged = self.info.ctx.get_merged_mm_kwargs(mm_kwargs) + if merged.keys() & {"size", "min_pixels", "max_pixels"}: + video_size = dict(self.info.get_video_processor().size) + size_override = merged.get("size") + if size_override is not None: + video_size = video_size | size_override + min_pixels = merged.get("min_pixels") + if min_pixels is not None: + video_size["shortest_edge"] = min_pixels + max_pixels = merged.get("max_pixels") + if max_pixels is not None: + video_size["longest_edge"] = max_pixels + video_mm_kwargs["size"] = video_size sampled_fps = video_mm_kwargs.get("fps") if is_list_of(sampled_fps, float): video_mm_kwargs["fps"] = sampled_fps[item_idx] From df362b2d6d091771dbcc364b2fc96d17a78df274 Mon Sep 17 00:00:00 2001 From: Andreas Karatzas Date: Sat, 18 Jul 2026 15:44:05 -0500 Subject: [PATCH 26/51] [ROCm][CI] Ensure sliding window tests release GPU memory (#49055) Signed-off-by: Andreas Karatzas --- .../test_correctness_sliding_window.py | 51 ++++++++++--------- 1 file changed, 28 insertions(+), 23 deletions(-) diff --git a/tests/v1/e2e/general/test_correctness_sliding_window.py b/tests/v1/e2e/general/test_correctness_sliding_window.py index 01d60444170..a8a29203d9e 100644 --- a/tests/v1/e2e/general/test_correctness_sliding_window.py +++ b/tests/v1/e2e/general/test_correctness_sliding_window.py @@ -33,7 +33,7 @@ model_config = { @pytest.mark.parametrize("seed", [1]) @pytest.mark.parametrize("disable_hybrid_kv_cache_manager", [True, False]) def test_sliding_window_retrieval( - model, batch_size, seed, disable_hybrid_kv_cache_manager + model, batch_size, seed, disable_hybrid_kv_cache_manager, vllm_runner ): """ The test does a bunch of assignments "x1 = 10\nx2 = 33\n..." and then @@ -48,34 +48,39 @@ def test_sliding_window_retrieval( test_config = model_config[model] - llm = LLM( - model=model, + with vllm_runner( + model, + max_model_len=None, + enable_chunked_prefill=None, disable_hybrid_kv_cache_manager=disable_hybrid_kv_cache_manager, enforce_eager=enforce_eager, - ) - sampling_params = SamplingParams(temperature=0.0, max_tokens=100) + ) as runner: + llm = runner.get_llm() + sampling_params = SamplingParams(temperature=0.0, max_tokens=100) - prompts, answer, indices = prep_prompts(batch_size, ln_range=test_config.ln_range) + prompts, answer, indices = prep_prompts( + batch_size, ln_range=test_config.ln_range + ) - check_length(prompts, llm, test_config.sliding_window) + check_length(prompts, llm, test_config.sliding_window) - # Fresh generation - responses = llm.generate(prompts, sampling_params) - check_answers( - indices, - answer, - [response.outputs[0].text for response in responses], - accept_rate=1.0, - ) + # Fresh generation + responses = llm.generate(prompts, sampling_params) + check_answers( + indices, + answer, + [response.outputs[0].text for response in responses], + accept_rate=1.0, + ) - # Re-generate with the same prompts to test prefix caching - responses = llm.generate(prompts, sampling_params) - check_answers( - indices, - answer, - [response.outputs[0].text for response in responses], - accept_rate=1.0, - ) + # Re-generate with the same prompts to test prefix caching + responses = llm.generate(prompts, sampling_params) + check_answers( + indices, + answer, + [response.outputs[0].text for response in responses], + accept_rate=1.0, + ) def check_length(prompts: list[str], llm: LLM, sliding_window: int): From 9243e0124e5c5396213258b2fb6d6401ee8965e9 Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Sun, 19 Jul 2026 05:41:04 +0800 Subject: [PATCH 27/51] [Multimodal] Automatically fallback to ViT DP when TP is unavailable (#49046) Signed-off-by: Isotr0py Co-authored-by: Roger Wang --- vllm/model_executor/models/kimi_k25.py | 6 +++--- vllm/model_executor/models/kimi_k25_vit.py | 2 +- vllm/model_executor/models/vision.py | 14 +++++++++++++- 3 files changed, 17 insertions(+), 5 deletions(-) diff --git a/vllm/model_executor/models/kimi_k25.py b/vllm/model_executor/models/kimi_k25.py index e2b74744b41..79b6e272bbb 100644 --- a/vllm/model_executor/models/kimi_k25.py +++ b/vllm/model_executor/models/kimi_k25.py @@ -34,6 +34,7 @@ from vllm.model_executor.models.kimi_k25_vit import ( MoonViT3dPretrainedModel, vision_tower_forward, ) +from vllm.model_executor.models.vision import is_vit_use_data_parallel from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import ( MultiModalFieldConfig, @@ -336,9 +337,8 @@ class KimiK25ForConditionalGeneration( self.config = config quant_config = vllm_config.quant_config - # Check for MoonViT config compatibility - self.use_data_parallel = ( - model_config.multimodal_config.mm_encoder_tp_mode == "data" + self.use_data_parallel = is_vit_use_data_parallel( + config.vision_config.num_attention_heads ) self.hidden_size = config.text_config.hidden_size self.device = current_platform.current_device() diff --git a/vllm/model_executor/models/kimi_k25_vit.py b/vllm/model_executor/models/kimi_k25_vit.py index 29ecb84674a..bcb8dd32de3 100644 --- a/vllm/model_executor/models/kimi_k25_vit.py +++ b/vllm/model_executor/models/kimi_k25_vit.py @@ -357,7 +357,7 @@ class MoonViTEncoderLayer(nn.Module): attn_bias: bool = False, ): super().__init__() - self.use_data_parallel = is_vit_use_data_parallel() + self.use_data_parallel = is_vit_use_data_parallel(num_heads) self.num_heads = num_heads self.hidden_dim = hidden_dim diff --git a/vllm/model_executor/models/vision.py b/vllm/model_executor/models/vision.py index 0582c125c66..18e994c95de 100644 --- a/vllm/model_executor/models/vision.py +++ b/vllm/model_executor/models/vision.py @@ -139,11 +139,23 @@ def get_fp8_padded_hidden_size(num_heads: int, head_dim: int) -> int | None: return num_heads * round_up(head_dim, 16) -def is_vit_use_data_parallel(): +def is_vit_use_data_parallel(num_heads: int | None = None) -> bool: """ Get the tensor parallel type for Vision Transformer. """ mm_cfg = get_multimodal_config() + can_split = ( + num_heads % get_tensor_model_parallel_world_size() == 0 + if num_heads is not None + else None + ) + if num_heads is not None and not can_split: + logger.warning_once( + "The number of vision attention heads is not divisible by " + "the tensor parallel size. Falling back to data parallelism " + "for the vision encoder." + ) + return True return mm_cfg is not None and mm_cfg.mm_encoder_tp_mode == "data" From b6ff8a2f509cc7ac9c58176f5115a836aa1e08bd Mon Sep 17 00:00:00 2001 From: Lucas Wilkinson Date: Sun, 19 Jul 2026 04:53:15 +0200 Subject: [PATCH 28/51] [Core] Add MRV2 virtual-batch PCP for MLA (#46570) Signed-off-by: Lucas Wilkinson Signed-off-by: Matthew Bonanni Co-authored-by: Codex Co-authored-by: OpenAI Codex Co-authored-by: Matthew Bonanni --- .buildkite/test_areas/lm_eval.yaml | 22 + .../configs/GLM-5.2-NVFP4-TP1-PCP4-EP.yaml | 16 + .../configs/GLM-5.2-NVFP4-TP2-PCP2-EP.yaml | 17 + tests/evals/gsm8k/configs/models-pcp.txt | 2 + tests/evals/gsm8k/gsm8k_eval.py | 16 +- tests/evals/gsm8k/test_gsm8k_correctness.py | 1 + .../v1/attention/test_attention_splitting.py | 16 + .../v1/attention/test_sparse_mla_backends.py | 7 + tests/v1/kv_offload/test_factory.py | 8 +- tests/v1/simple_kv_offload/test_scheduler.py | 40 +- vllm/config/parallel.py | 62 +- vllm/config/vllm.py | 5 +- .../device_communicators/all2all.py | 44 +- .../base_device_communicator.py | 11 +- .../kv_connector/v1/offloading/config.py | 9 +- vllm/distributed/parallel_state.py | 25 +- .../layers/attention/mla_attention.py | 72 +- vllm/model_executor/layers/attention/pcp.py | 92 +++ .../layers/attention/sparse_mla_attention.py | 4 +- .../model_executor/layers/fused_moe/config.py | 4 +- .../layers/fused_moe/runner/moe_runner.py | 28 +- .../layers/sparse_attn_indexer.py | 29 +- vllm/v1/attention/backend.py | 16 +- vllm/v1/attention/backends/mla/indexer.py | 23 +- vllm/v1/attention/backends/utils.py | 4 +- vllm/v1/attention/selector.py | 5 +- vllm/v1/core/kv_cache_coordinator.py | 2 - vllm/v1/core/kv_cache_utils.py | 11 +- vllm/v1/core/sched/scheduler.py | 2 +- vllm/v1/core/single_type_kv_cache_manager.py | 12 +- vllm/v1/kv_cache_interface.py | 16 +- vllm/v1/simple_kv_offload/manager.py | 7 +- vllm/v1/worker/block_table.py | 6 +- vllm/v1/worker/cp_utils.py | 25 +- vllm/v1/worker/gpu/attn_utils.py | 7 + vllm/v1/worker/gpu/block_table.py | 22 +- vllm/v1/worker/gpu/model_runner.py | 24 +- vllm/v1/worker/gpu/model_states/default.py | 8 +- vllm/v1/worker/gpu/pcp_manager.py | 680 ++++++++++++++++++ 39 files changed, 1204 insertions(+), 196 deletions(-) create mode 100644 tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP1-PCP4-EP.yaml create mode 100644 tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP2-PCP2-EP.yaml create mode 100644 tests/evals/gsm8k/configs/models-pcp.txt create mode 100644 vllm/model_executor/layers/attention/pcp.py create mode 100644 vllm/v1/worker/gpu/pcp_manager.py diff --git a/.buildkite/test_areas/lm_eval.yaml b/.buildkite/test_areas/lm_eval.yaml index 9c08c96e4c4..95a8bafdd58 100644 --- a/.buildkite/test_areas/lm_eval.yaml +++ b/.buildkite/test_areas/lm_eval.yaml @@ -78,6 +78,28 @@ steps: commands: - pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-small-tp.txt +- label: LM Eval PCP (4xB200) + key: lm-eval-pcp-4xb200 + timeout_in_minutes: 360 + device: b200-k8s + num_devices: 4 + optional: true + source_file_dependencies: + - csrc/ + - tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP2-PCP2-EP.yaml + - tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP1-PCP4-EP.yaml + - tests/evals/gsm8k/configs/models-pcp.txt + - vllm/model_executor/layers/quantization + - vllm/config/parallel.py + - vllm/distributed/parallel_state.py + - vllm/model_executor/layers/attention/mla_attention.py + - vllm/model_executor/layers/attention/pcp.py + - vllm/v1/worker/gpu/model_runner.py + - vllm/v1/worker/gpu/pcp_manager.py + autorun_on_main: true + commands: + - pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-pcp.txt + - label: LM Eval Large Models EP (2xB200) key: lm-eval-large-models-ep-2xb200 timeout_in_minutes: 60 diff --git a/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP1-PCP4-EP.yaml b/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP1-PCP4-EP.yaml new file mode 100644 index 00000000000..13d71af20e6 --- /dev/null +++ b/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP1-PCP4-EP.yaml @@ -0,0 +1,16 @@ +model_name: "nvidia/GLM-5.2-NVFP4" +accuracy_threshold: 0.90 +num_questions: 1319 +num_fewshot: 5 +max_concurrency: 100 +server_args: >- + --enforce-eager + --max-model-len 4096 + --safetensors-load-strategy prefetch + --moe-backend flashinfer_cutlass + --prefill-context-parallel-size 4 + --enable-expert-parallel + --kv-cache-dtype fp8 +env: + VLLM_LOGGING_LEVEL: "DEBUG" + VLLM_USE_V2_MODEL_RUNNER: "1" diff --git a/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP2-PCP2-EP.yaml b/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP2-PCP2-EP.yaml new file mode 100644 index 00000000000..b21d0b12f02 --- /dev/null +++ b/tests/evals/gsm8k/configs/GLM-5.2-NVFP4-TP2-PCP2-EP.yaml @@ -0,0 +1,17 @@ +model_name: "nvidia/GLM-5.2-NVFP4" +accuracy_threshold: 0.90 +num_questions: 1319 +num_fewshot: 5 +max_concurrency: 100 +server_args: >- + --enforce-eager + --max-model-len 4096 + --safetensors-load-strategy prefetch + --moe-backend flashinfer_cutlass + --tensor-parallel-size 2 + --prefill-context-parallel-size 2 + --enable-expert-parallel + --kv-cache-dtype fp8 +env: + VLLM_LOGGING_LEVEL: "DEBUG" + VLLM_USE_V2_MODEL_RUNNER: "1" diff --git a/tests/evals/gsm8k/configs/models-pcp.txt b/tests/evals/gsm8k/configs/models-pcp.txt new file mode 100644 index 00000000000..df645edad57 --- /dev/null +++ b/tests/evals/gsm8k/configs/models-pcp.txt @@ -0,0 +1,2 @@ +GLM-5.2-NVFP4-TP2-PCP2-EP.yaml +GLM-5.2-NVFP4-TP1-PCP4-EP.yaml diff --git a/tests/evals/gsm8k/gsm8k_eval.py b/tests/evals/gsm8k/gsm8k_eval.py index 9c47826850b..a8ee6833ca2 100644 --- a/tests/evals/gsm8k/gsm8k_eval.py +++ b/tests/evals/gsm8k/gsm8k_eval.py @@ -217,6 +217,7 @@ def evaluate_gsm8k( seed: int | None = 42, request_timeout_seconds: float = 600, gen_prefix: str = "", + max_concurrency: int | None = None, ) -> dict[str, float | int]: """ Evaluate GSM8K accuracy using vLLM serve endpoint. @@ -261,7 +262,14 @@ def evaluate_gsm8k( return answer, tokens timeout = aiohttp.ClientTimeout(total=request_timeout_seconds) - async with aiohttp.ClientSession(timeout=timeout) as session: + connector = ( + aiohttp.TCPConnector(limit=max_concurrency) + if max_concurrency is not None + else None + ) + async with aiohttp.ClientSession( + timeout=timeout, connector=connector + ) as session: tasks = [get_answer(session, i) for i in range(num_questions)] await tqdm.gather(*tasks, desc="Evaluating") @@ -343,6 +351,11 @@ def main() -> None: parser.add_argument( "--seed", type=int, default=42, help="Random seed for reproducibility" ) + parser.add_argument( + "--max-concurrency", + type=int, + help="Maximum number of concurrent requests", + ) parser.add_argument("--save-results", type=str, help="Save results to JSON file") args = parser.parse_args() @@ -355,6 +368,7 @@ def main() -> None: port=args.port, temperature=args.temperature, seed=args.seed, + max_concurrency=args.max_concurrency, ) # Print results to terminal diff --git a/tests/evals/gsm8k/test_gsm8k_correctness.py b/tests/evals/gsm8k/test_gsm8k_correctness.py index 4b2b6ad3813..e1ae6ec5427 100644 --- a/tests/evals/gsm8k/test_gsm8k_correctness.py +++ b/tests/evals/gsm8k/test_gsm8k_correctness.py @@ -59,6 +59,7 @@ def run_gsm8k_eval(eval_config: dict, server_url: str) -> dict: seed=eval_config.get("seed", 42), request_timeout_seconds=request_timeout_seconds, gen_prefix=eval_config.get("gen_prefix", ""), + max_concurrency=eval_config.get("max_concurrency"), ) return results diff --git a/tests/v1/attention/test_attention_splitting.py b/tests/v1/attention/test_attention_splitting.py index d6fc59576b2..acbc22a0279 100644 --- a/tests/v1/attention/test_attention_splitting.py +++ b/tests/v1/attention/test_attention_splitting.py @@ -160,6 +160,8 @@ def apply_split_decodes_and_prefills( decode_threshold: int, require_uniform: bool, padded_num_tokens: int | None = None, + is_prefilling: list[bool] | None = None, + treat_short_extends_as_decodes: bool = True, ): """Helper function to apply split_decodes_and_prefills and return the results.""" @@ -173,11 +175,14 @@ def apply_split_decodes_and_prefills( if padded_num_tokens is not None: common_metadata.num_actual_tokens = padded_num_tokens + if is_prefilling is not None: + common_metadata.is_prefilling = torch.tensor(is_prefilling) return split_decodes_and_prefills( common_metadata, decode_threshold=decode_threshold, require_uniform=require_uniform, + treat_short_extends_as_decodes=treat_short_extends_as_decodes, ) @@ -236,6 +241,17 @@ def test_split_decodes_and_prefills_uniform_all_ones(): assert num_prefill_tokens == 0 +def test_split_decodes_and_prefills_uniform_short_extend(): + result = apply_split_decodes_and_prefills( + [1, 1], + decode_threshold=1, + require_uniform=True, + is_prefilling=[False, True], + treat_short_extends_as_decodes=False, + ) + assert result == (1, 1, 1, 1) + + def test_split_decodes_and_prefills_uniform_all_short_decodes(): query_lens = [2, 2, 1, 3, 2, 1, 2] num_decodes, num_prefills, num_decode_tokens, num_prefill_tokens = ( diff --git a/tests/v1/attention/test_sparse_mla_backends.py b/tests/v1/attention/test_sparse_mla_backends.py index 6ec9fc25ed3..3fe5df918ab 100644 --- a/tests/v1/attention/test_sparse_mla_backends.py +++ b/tests/v1/attention/test_sparse_mla_backends.py @@ -1114,6 +1114,13 @@ def test_sparse_backend_prefill_correctness( @pytest.mark.parametrize( "seq_lens,query_lens,workspace_size,max_logits_bytes,expected", [ + ( + torch.tensor([0]), + torch.tensor([0]), + 100, + 1000, + [], + ), # Logits constraint triggers split (M*N exceeds budget) # req0: M=10, N=100 -> 1000 elems (4000 bytes) - fits in 5000 # req1: adding M=10, N=100 -> new_M=20, new_N=200 -> 4000 elems > 1250 diff --git a/tests/v1/kv_offload/test_factory.py b/tests/v1/kv_offload/test_factory.py index 570924cfccc..d08dbee5765 100644 --- a/tests/v1/kv_offload/test_factory.py +++ b/tests/v1/kv_offload/test_factory.py @@ -385,7 +385,7 @@ def test_tiering_spec_aligns_row_size(): assert spec.num_blocks == cpu_bytes_to_use // alignment -def test_offloading_spec_resolves_prefill_context_parallel_block_sizes(): +def test_offloading_spec_kv_sharding_ignores_prefill_context_parallel(): config = _make_layout_vllm_config( cpu_bytes_to_use=65536, extra_config={"block_size": 64}, @@ -394,9 +394,9 @@ def test_offloading_spec_resolves_prefill_context_parallel_block_sizes(): spec = _create_spec(config, _make_kv_cache_config()) - assert spec.tokens_per_block == (32,) - assert spec.tokens_per_hash == 32 - assert spec.blocks_per_chunk == 2 + assert spec.tokens_per_block == (16,) + assert spec.tokens_per_hash == 16 + assert spec.blocks_per_chunk == 4 def test_offloading_config_preserves_data_parallel_index(): diff --git a/tests/v1/simple_kv_offload/test_scheduler.py b/tests/v1/simple_kv_offload/test_scheduler.py index 09586f5e6b4..4f7761cf35b 100644 --- a/tests/v1/simple_kv_offload/test_scheduler.py +++ b/tests/v1/simple_kv_offload/test_scheduler.py @@ -1551,9 +1551,8 @@ def _make_cp_scheduler( num_gpu_blocks: int = 16, lazy: bool = False, ) -> SchedulerFixture: - """Build a SimpleCPUOffloadScheduler with CP-scaled virtual block size.""" - cp_world_size = dcp_world_size * pcp_world_size - virtual_block_size = BLOCK_SIZE * cp_world_size + """Build a SimpleCPUOffloadScheduler with DCP-scaled block size.""" + virtual_block_size = BLOCK_SIZE * dcp_world_size kv_cache_config = _make_kv_cache_config(num_gpu_blocks) vllm_config = _make_cp_vllm_config(dcp_world_size, pcp_world_size) @@ -1645,17 +1644,14 @@ def _allocate_cp_gpu_blocks( ], ) def test_cp_block_size_scaling(dcp_world_size: int, pcp_world_size: int) -> None: - """Verify that the scheduler's block_size and cp_world_size are correctly - scaled when context parallelism is enabled.""" + """Verify block size scaling follows DCP ownership.""" fix = _make_cp_scheduler( dcp_world_size=dcp_world_size, pcp_world_size=pcp_world_size ) sched = fix.scheduler - expected_cp = dcp_world_size * pcp_world_size - assert sched.cp_world_size == expected_cp - assert sched.block_size == BLOCK_SIZE * expected_cp - assert sched.fa_block_size == BLOCK_SIZE * expected_cp + assert sched.cp_world_size == dcp_world_size + assert sched.block_size == BLOCK_SIZE * dcp_world_size # --------------------------------------------------------------------------- @@ -1682,8 +1678,7 @@ def test_cp_eager_store_and_load_roundtrip( lazy=False, ) sched = fix.scheduler - cp = dcp_world_size * pcp_world_size - vbs = BLOCK_SIZE * cp + vbs = BLOCK_SIZE * dcp_world_size num_blocks = 2 req = _make_cp_request(num_blocks, vbs) @@ -1749,18 +1744,16 @@ def test_cp_eager_store_and_load_roundtrip( def test_cp_effective_block_size_store_and_load( dcp_world_size: int, pcp_world_size: int ) -> None: - """Verify ready_blocks_g (store) and n_take_g (load) use the effective - (physical * cp) block size, not the per-rank physical size.""" + """Verify store/load physical block counts scale with DCP, not PCP.""" fix = _make_cp_scheduler( dcp_world_size=dcp_world_size, pcp_world_size=pcp_world_size ) sched = fix.scheduler gpu_pool = fix.gpu_block_pool - cp = dcp_world_size * pcp_world_size - vbs = BLOCK_SIZE * cp + vbs = BLOCK_SIZE * dcp_world_size + expected_blocks = 1 - # Store: allocate 2 blocks, confirm only 1. Without the fix, - # ready_blocks_g = vbs / BLOCK_SIZE = 2, storing both blocks. + # Store one DCP-scaled logical block worth of tokens. req = _make_cp_request(num_blocks=2, virtual_block_size=vbs) gpu_blocks = _allocate_cp_gpu_blocks(gpu_pool, req, 2, vbs) kv = KVCacheBlocks(blocks=(gpu_blocks,)) @@ -1772,12 +1765,11 @@ def test_cp_effective_block_size_store_and_load( new_reqs={req.request_id: kv.get_block_ids()}, ) ) - assert len(m1.store_gpu_blocks) == 1 - assert len(m1.store_cpu_blocks) == 1 + assert len(m1.store_gpu_blocks) == expected_blocks + assert len(m1.store_cpu_blocks) == expected_blocks simulate_store_completion(sched, m1.store_event) - # Load: store 2 blocks from a second request, accept only 1 as external. - # Without the fix, n_take_g = vbs / BLOCK_SIZE = 2, loading both. + # Load one DCP-scaled logical block from a two-block CPU hit. req2 = _make_cp_request(num_blocks=2, virtual_block_size=vbs) kv2 = KVCacheBlocks(blocks=(_allocate_cp_gpu_blocks(gpu_pool, req2, 2, vbs),)) req2.num_computed_tokens = 2 * vbs @@ -1801,7 +1793,7 @@ def test_cp_effective_block_size_store_and_load( hit, _ = sched.get_num_new_matched_tokens(req3, num_computed_tokens=0) assert hit == 2 * vbs - kv3 = KVCacheBlocks(blocks=(gpu_pool.get_new_blocks(2),)) + kv3 = KVCacheBlocks(blocks=(gpu_pool.get_new_blocks(expected_blocks),)) sched.update_state_after_alloc(req3, kv3, num_external_tokens=vbs) m3 = sched.build_connector_meta( make_scheduler_output( @@ -1810,8 +1802,8 @@ def test_cp_effective_block_size_store_and_load( ) ) assert m3.load_event >= 0 - assert len(m3.load_gpu_blocks) == 1 - assert len(m3.load_cpu_blocks) == 1 + assert len(m3.load_gpu_blocks) == expected_blocks + assert len(m3.load_cpu_blocks) == expected_blocks assert m3.load_gpu_blocks == [kv3.get_block_ids()[0][0]] diff --git a/vllm/config/parallel.py b/vllm/config/parallel.py index 7c270b0c0eb..53688c05d92 100644 --- a/vllm/config/parallel.py +++ b/vllm/config/parallel.py @@ -122,10 +122,11 @@ class ParallelConfig: tensor_parallel_size: int = Field(default=1, ge=1) """Number of tensor parallel groups.""" prefill_context_parallel_size: int = Field(default=1, ge=1) - """Number of prefill context parallel groups.""" + """Number of ranks that split prefill sequence computation. PCP expands + the process world size but does not increase the KV-cache shard count.""" data_parallel_size: int = Field(default=1, ge=1) """Number of data parallel groups. MoE layers will be sharded according to - the product of the tensor parallel size and data parallel size.""" + the product of the tensor, prefill-context, and data parallel sizes.""" data_parallel_size_local: int = Field(default=1, ge=0) """Number of local data parallel groups. A value of 0 is a sentinel used by the engine-args layer to signal that data parallelism was specified @@ -337,9 +338,9 @@ class ParallelConfig: connect as clients to exchange self-picked group ports at runtime.""" decode_context_parallel_size: int = Field(default=1, ge=1) - """Number of decode context parallel groups, because the world size does - not change by dcp, it simply reuse the GPUs of TP group, and tp_size - needs to be divisible by dcp_size.""" + """Number of ranks that shard the decode KV cache. DCP does not expand + the process world size. Without PCP, DCP reuses TP ranks. With PCP, DCP + either spans the PCP axis or the full TP x PCP block.""" dcp_kv_cache_interleave_size: int = 1 """ @@ -357,13 +358,11 @@ class ParallelConfig: """ cp_kv_cache_interleave_size: int = 1 - """Interleave size of kv_cache storage while using DCP or PCP. - For `total_cp_rank = pcp_rank * dcp_world_size + dcp_rank`, - and `total_cp_world_size = pcp_world_size * dcp_world_size`. - store interleave_size tokens on total_cp_rank i, - then store next interleave_size tokens on total_cp_rank i+1. + """Interleave size of kv_cache storage while using DCP. + Store interleave_size tokens on dcp_rank i, then store next + interleave_size tokens on dcp_rank i+1. Interleave_size=1: token-level alignment, where token `i` is stored on - total_cp_rank `i % total_cp_world_size`. + dcp_rank `i % dcp_world_size`. Interleave_size=block_size: block-level alignment, where tokens are first populated to the preceding ranks. Tokens are then stored in (rank i+1, block j) only after (rank i, block j) is fully occupied. @@ -480,11 +479,19 @@ class ParallelConfig: ) if not self.enable_expert_parallel: raise ValueError("enable_expert_parallel must be True to use EPLB.") - if self.tensor_parallel_size * self.data_parallel_size <= 1: + # The EP group spans the TP x PCP x DP ranks. EPLB therefore needs + # TP, PCP, or DP > 1. + if ( + self.tensor_parallel_size + * self.prefill_context_parallel_size + * self.data_parallel_size + <= 1 + ): raise ValueError( - "EPLB requires tensor_parallel_size or data_parallel_size " - f"to be greater than 1, but got " - f"TP={self.tensor_parallel_size},DP={self.data_parallel_size}." + "EPLB requires tensor, prefill-context, or data parallelism, " + f"but got TP={self.tensor_parallel_size}, " + f"PCP={self.prefill_context_parallel_size}, " + f"DP={self.data_parallel_size}." ) else: if self.eplb_config.num_redundant_experts != 0: @@ -495,15 +502,21 @@ class ParallelConfig: "num_redundant_experts." ) - # Note(hc): In the current implementation of decode context - # parallel(DCP), tp_size needs to be divisible by dcp_size, - # because the world size does not change by dcp, it simply - # reuses the GPUs of TP group, and split one TP group into - # tp_size//dcp_size DCP groups. - if self.tensor_parallel_size % self.decode_context_parallel_size != 0: + tp = self.tensor_parallel_size + pcp = self.prefill_context_parallel_size + dcp = self.decode_context_parallel_size + if pcp > 1 and self.data_parallel_size > 1: + raise ValueError("PCP does not support data parallelism yet.") + if pcp == 1: + # DCP reuses the TP ranks when PCP is disabled. + if tp % dcp != 0: + raise ValueError(f"tp_size={tp} must be divisible by dcp_size={dcp}.") + elif dcp not in (1, pcp, tp * pcp): raise ValueError( - f"tp_size={self.tensor_parallel_size} must be divisible by" - f"dcp_size={self.decode_context_parallel_size}." + "When PCP is enabled, DCP must be disabled, span the PCP " + "axis, or span the full TP x PCP axis. " + f"Got TP={tp}, PCP={pcp}, DCP={dcp}; valid DCP sizes are " + f"{sorted({1, pcp, tp * pcp})}." ) if self.dcp_comm_backend == "a2a" and self.decode_context_parallel_size <= 1: @@ -515,8 +528,7 @@ class ParallelConfig: @property def world_size_across_dp(self) -> int: - """world_size_across_dp is TPxPPxDP, it is the size of the world - including data parallelism.""" + """Process world size across TP, PCP, PP, and DP.""" return self.world_size * self.data_parallel_size @property diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index cca9adbcf58..4b4a97f41b6 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -2130,9 +2130,10 @@ class VllmConfig: model_config = self.model_config speculative_config = self.speculative_config - if self.parallel_config.prefill_context_parallel_size > 1: + if self.parallel_config.prefill_context_parallel_size > 1 and not ( + model_config is not None and model_config.use_mla + ): unsupported.append("prefill context parallelism") - if self.compilation_config.mode == CompilationMode.STOCK_TORCH_COMPILE: unsupported.append("stock torch.compile") diff --git a/vllm/distributed/device_communicators/all2all.py b/vllm/distributed/device_communicators/all2all.py index 7f540bc4b1f..33841de306e 100644 --- a/vllm/distributed/device_communicators/all2all.py +++ b/vllm/distributed/device_communicators/all2all.py @@ -8,7 +8,7 @@ import torch import torch.distributed as dist import vllm.envs as envs -from vllm.distributed import get_dp_group, get_ep_group +from vllm.distributed import get_dp_group, get_ep_group, get_pcp_group from vllm.distributed.utils import StatelessProcessGroup from vllm.forward_context import get_forward_context from vllm.logger import init_logger @@ -49,6 +49,23 @@ class AgRsAll2AllManager(All2AllManagerBase): def __init__(self, cpu_group, tcp_store_group=None): super().__init__(cpu_group, tcp_store_group) + def _get_comm_group(self, is_sequence_parallel: bool) -> Any: + if is_sequence_parallel: + return get_ep_group() + if self.dp_world_size > 1: + return get_dp_group() + return get_pcp_group() + + def _get_sizes(self, num_local_tokens: int, comm_group: Any) -> list[int]: + if self.dp_world_size == 1: + return [num_local_tokens] * comm_group.world_size + + dp_metadata = get_forward_context().dp_metadata + assert dp_metadata is not None + sizes = dp_metadata.get_chunk_sizes_across_dp_rank() + assert sizes is not None + return sizes + def dispatch_router_logits( self, hidden_states: torch.Tensor, @@ -62,11 +79,8 @@ class AgRsAll2AllManager(All2AllManagerBase): """ Gather hidden_states and router_logits from all dp ranks. """ - dp_metadata = get_forward_context().dp_metadata - assert dp_metadata is not None - sizes = dp_metadata.get_chunk_sizes_across_dp_rank() - assert sizes is not None - dist_group = get_ep_group() if is_sequence_parallel else get_dp_group() + dist_group = self._get_comm_group(is_sequence_parallel) + sizes = self._get_sizes(hidden_states.shape[0], dist_group) assert sizes[dist_group.rank_in_group] == hidden_states.shape[0] tensors_to_gather = [hidden_states, router_logits] @@ -97,11 +111,8 @@ class AgRsAll2AllManager(All2AllManagerBase): """ Gather hidden_states and router_logits from all dp ranks. """ - dp_metadata = get_forward_context().dp_metadata - assert dp_metadata is not None - sizes = dp_metadata.get_chunk_sizes_across_dp_rank() - assert sizes is not None - dist_group = get_ep_group() if is_sequence_parallel else get_dp_group() + dist_group = self._get_comm_group(is_sequence_parallel) + sizes = self._get_sizes(hidden_states.shape[0], dist_group) assert sizes[dist_group.rank_in_group] == hidden_states.shape[0] tensors_to_gather = [hidden_states, topk_weights, topk_ids] @@ -129,12 +140,11 @@ class AgRsAll2AllManager(All2AllManagerBase): """ Reduce-scatter hidden_states across all dp ranks. """ - dp_metadata = get_forward_context().dp_metadata - assert dp_metadata is not None - sizes = dp_metadata.get_chunk_sizes_across_dp_rank() - assert sizes is not None - - dist_group = get_ep_group() if is_sequence_parallel else get_dp_group() + dist_group = self._get_comm_group(is_sequence_parallel) + sizes = self._get_sizes( + hidden_states.shape[0] // dist_group.world_size, + dist_group, + ) hidden_states = dist_group.reduce_scatterv(hidden_states, dim=0, sizes=sizes) return hidden_states diff --git a/vllm/distributed/device_communicators/base_device_communicator.py b/vllm/distributed/device_communicators/base_device_communicator.py index 9a443b7fc16..70f1fb5d62c 100644 --- a/vllm/distributed/device_communicators/base_device_communicator.py +++ b/vllm/distributed/device_communicators/base_device_communicator.py @@ -176,11 +176,16 @@ class DeviceCommunicatorBase: config = get_current_vllm_config_or_none() if config is not None: # initialize the all2all manager for DP or sequence-parallel EP. + parallel_config = config.parallel_config use_ep = ( - config.parallel_config.data_parallel_size > 1 - or config.parallel_config.use_sequence_parallel_moe + parallel_config.data_parallel_size > 1 + or parallel_config.use_sequence_parallel_moe + or ( + parallel_config.enable_expert_parallel + and parallel_config.prefill_context_parallel_size > 1 + ) ) - all2all_backend = config.parallel_config.all2all_backend + all2all_backend = parallel_config.all2all_backend self.is_ep_communicator = unique_name.split(":")[0] == "ep" self.use_all2all = self.is_ep_communicator and use_ep diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/offloading/config.py b/vllm/distributed/kv_transfer/kv_connector/v1/offloading/config.py index 5cf4dffd1af..b807b480f02 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/offloading/config.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/offloading/config.py @@ -36,13 +36,12 @@ def build_offloading_config( engine_id = kv_transfer_config.engine_id parallel_config = vllm_config.parallel_config - context_parallel_factor = ( - parallel_config.decode_context_parallel_size - * parallel_config.prefill_context_parallel_size - ) groups = tuple( OffloadingGroupConfig( - tokens_per_block=(group.kv_cache_spec.block_size * context_parallel_factor), + tokens_per_block=( + group.kv_cache_spec.block_size + * parallel_config.decode_context_parallel_size + ), layer_names=tuple(group.layer_names), ) for group in kv_cache_config.kv_cache_groups diff --git a/vllm/distributed/parallel_state.py b/vllm/distributed/parallel_state.py index e46ca1691c2..f0647323e61 100644 --- a/vllm/distributed/parallel_state.py +++ b/vllm/distributed/parallel_state.py @@ -1784,7 +1784,7 @@ def initialize_model_parallel( get_world_group().device_group ) - # the layout order is: ExternalDP x DP x PP x TP + # the layout order is: ExternalDP x DP x PP x PCP x TP # ExternalDP is the data parallel group that is not part of the model, # every dp rank can generate independently (in verl integration). # DP is the data parallel group that is part of the model, @@ -1821,17 +1821,13 @@ def initialize_model_parallel( # Build the DCP model-parallel groups. global _DCP assert _DCP is None, "decode context model parallel group is already initialized" - # Note(hc): In the current implementation of decode context parallel, - # dcp_size must not exceed tp_size, because the world size does not - # change by DCP, it simply reuses the GPUs of TP group, and split one - # TP group into tp_size//dcp_size DCP groups. - group_ranks = all_ranks.reshape(-1, decode_context_model_parallel_size).unbind(0) + dcp_size = decode_context_model_parallel_size or 1 + dcp_ranks = local_all_ranks if enable_elastic_ep else all_ranks + if dcp_size > 1: + # DCP spans PCP first, then TP for full TP x PCP groups. + dcp_ranks = dcp_ranks.transpose(-1, -2) + group_ranks = dcp_ranks.reshape(-1, dcp_size).unbind(0) group_ranks = [x.tolist() for x in group_ranks] - if enable_elastic_ep: - group_ranks = local_all_ranks.reshape( - -1, decode_context_model_parallel_size - ).unbind(0) - group_ranks = [x.tolist() for x in group_ranks] _DCP = init_model_parallel_group( group_ranks, get_world_group().local_rank, @@ -2005,6 +2001,13 @@ def ensure_model_parallel_initialized( f"{pcp_world_size=} vs. " f"{prefill_context_model_parallel_size=}" ) + dcp_world_size = get_dcp_group().world_size + dcp_model_parallel_size = decode_context_model_parallel_size or 1 + assert dcp_world_size == dcp_model_parallel_size, ( + "decode context parallel group already initialized, but of unexpected size: " + f"{dcp_world_size=} vs. " + f"{dcp_model_parallel_size=}" + ) def prepare_communication_buffer_for_model(model: torch.nn.Module): diff --git a/vllm/model_executor/layers/attention/mla_attention.py b/vllm/model_executor/layers/attention/mla_attention.py index 55ad7aedef2..a450ada185b 100644 --- a/vllm/model_executor/layers/attention/mla_attention.py +++ b/vllm/model_executor/layers/attention/mla_attention.py @@ -211,6 +211,7 @@ from vllm.config import ( from vllm.config.cache import CacheDType from vllm.distributed.parallel_state import ( get_dcp_group, + get_tp_group, is_global_first_rank, ) from vllm.forward_context import ForwardContext, get_forward_context @@ -225,6 +226,10 @@ from vllm.model_executor.layers.attention.attention import ( from vllm.model_executor.layers.attention.kv_transfer_utils import ( maybe_transfer_kv_layer, ) +from vllm.model_executor.layers.attention.pcp import ( + finalize_mla_pcp_decode, + maybe_gather_mla_latent_cache_inputs, +) from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase from vllm.model_executor.layers.linear import ( ColumnParallelLinear, @@ -269,7 +274,7 @@ from vllm.v1.attention.backends.utils import ( get_dcp_local_seq_lens, split_decodes_and_prefills, ) -from vllm.v1.attention.ops.common import cp_lse_ag_out_rs +from vllm.v1.attention.ops.common import cp_lse_ag_out_ar, cp_lse_ag_out_rs from vllm.v1.attention.ops.dcp_alltoall import dcp_a2a_lse_reduce from vllm.v1.attention.ops.merge_attn_states import merge_attn_states from vllm.v1.attention.selector import get_attn_backend @@ -499,6 +504,8 @@ class MLAAttention(nn.Module, AttentionLayerBase): self.use_direct_call = not current_platform.opaque_attention_op() vllm_config = get_current_vllm_config() + parallel_config = vllm_config.parallel_config + self.use_pcp = parallel_config.prefill_context_parallel_size > 1 compilation_config = vllm_config.compilation_config if prefix in compilation_config.static_forward_context: raise ValueError(f"Duplicate layer name: {prefix}") @@ -611,11 +618,23 @@ class MLAAttention(nn.Module, AttentionLayerBase): assert isinstance(slot_mapping, dict), ( f"Expected slot_mapping to be a dict, got {type(slot_mapping)}. " ) + layer_slot_mapping = slot_mapping.get(self.layer_name) + kv_for_cache, kpe_for_cache, layer_slot_mapping = ( + maybe_gather_mla_latent_cache_inputs( + kv_c_normed, + k_pe, + layer_slot_mapping, + attn_metadata.num_decode_tokens + if attn_metadata is not None + else None, + self.use_pcp, + ) + ) self.impl.do_kv_cache_update( # type: ignore[attr-defined] - kv_c_normed, - k_pe, + kv_for_cache, + kpe_for_cache, self_kv_cache, - slot_mapping.get(self.layer_name), + layer_slot_mapping, self.kv_cache_dtype, self._k_scale, ) @@ -710,6 +729,10 @@ class MLAAttention(nn.Module, AttentionLayerBase): fp8_attention = is_quantized_kv_cache(self.kv_cache_dtype) num_actual_toks = attn_metadata.num_actual_tokens + if self.use_pcp and self.impl.dcp_world_size > 1 and quant_key is not None: + raise NotImplementedError( + "MRV2 MLA PCP+DCP does not support fused output quantization yet." + ) # Inputs and outputs may be padded for CUDA graphs output_padded = output @@ -833,11 +856,15 @@ class MLAAttention(nn.Module, AttentionLayerBase): else: mqa_q = (mqa_ql_nope, mqa_q_pe) if self.impl.dcp_world_size > 1: - if isinstance(mqa_q, tuple): - # concatenate mqa_ql_nope and mqa_q_pe -> (B, N, L + P) - mqa_q = torch.cat(mqa_q, dim=-1) - # mqa_q do allgather in head dim. - mqa_q = get_dcp_group().all_gather(mqa_q, dim=1) + if self.use_pcp: + if self.impl.dcp_world_size > self.impl.pcp_world_size: + if isinstance(mqa_q, tuple): + mqa_q = torch.cat(mqa_q, dim=-1) + mqa_q = get_tp_group().all_gather(mqa_q, dim=1) + else: + if isinstance(mqa_q, tuple): + mqa_q = torch.cat(mqa_q, dim=-1) + mqa_q = get_dcp_group().all_gather(mqa_q, dim=1) # call decode attn if not self.impl.is_sparse: @@ -846,6 +873,7 @@ class MLAAttention(nn.Module, AttentionLayerBase): # correct dcp attn_out with lse. if self.impl.dcp_world_size > 1: + assert lse is not None if self.dcp_a2a: attn_out = dcp_a2a_lse_reduce( attn_out, @@ -853,6 +881,13 @@ class MLAAttention(nn.Module, AttentionLayerBase): get_dcp_group(), is_lse_base_on_e=self.impl.lse_base_on_e, ) + elif self.use_pcp: + attn_out = cp_lse_ag_out_ar( + attn_out, + lse, + get_dcp_group(), + is_lse_base_on_e=self.impl.lse_base_on_e, + ) else: attn_out = cp_lse_ag_out_rs( attn_out, @@ -860,6 +895,8 @@ class MLAAttention(nn.Module, AttentionLayerBase): get_dcp_group(), is_lse_base_on_e=self.impl.lse_base_on_e, ) + if self.use_pcp: + attn_out = finalize_mla_pcp_decode(attn_out, self.num_heads) # v_up projection self._v_up_proj(attn_out, out=mqa_output_slice) @@ -903,6 +940,8 @@ class MLAAttention(nn.Module, AttentionLayerBase): raise ValueError(f"Unsupported quant_key: {quant_key}") return quant_output + if self.use_pcp and output_padded.shape[0] > num_actual_toks: + output_padded[num_actual_toks:].zero_() return output_padded def process_weights_after_loading(self, act_dtype: torch.dtype): @@ -1082,8 +1121,17 @@ def unified_mla_kv_cache_update( the data dependency between them to ensure torch.compile preserves ordering. """ layer_name = _resolve_layer_name(layer_name) - _, attn_layer, kv_cache, layer_slot_mapping = get_attention_context(layer_name) + attn_metadata, attn_layer, kv_cache, layer_slot_mapping = get_attention_context( + layer_name + ) if layer_slot_mapping is not None: + kv_c_normed, k_pe, layer_slot_mapping = maybe_gather_mla_latent_cache_inputs( + kv_c_normed, + k_pe, + layer_slot_mapping, + attn_metadata.num_decode_tokens if attn_metadata is not None else None, + attn_layer.use_pcp, + ) attn_layer.impl.do_kv_cache_update( # type: ignore[attr-defined] kv_c_normed, k_pe, @@ -1723,6 +1771,7 @@ class MLACommonMetadataBuilder(AttentionMetadataBuilder[M]): self.compilation_config = vllm_config.compilation_config self.vllm_config = vllm_config self.device = device + self.use_pcp = parallel_config.prefill_context_parallel_size > 1 self.num_heads = self.model_config.get_num_attention_heads(parallel_config) self.mla_dims = get_mla_dims(self.model_config) @@ -1735,11 +1784,9 @@ class MLACommonMetadataBuilder(AttentionMetadataBuilder[M]): try: self.dcp_world_size = get_dcp_group().world_size - self.dcp_rank = get_dcp_group().rank_in_group except AssertionError: # DCP might not be initialized in testing self.dcp_world_size = 1 - self.dcp_rank = 0 self.dcp_local_block_size = parallel_config.cp_kv_cache_interleave_size self.dcp_virtual_block_size = self.dcp_local_block_size * self.dcp_world_size self.cp_kv_cache_interleave_size = parallel_config.cp_kv_cache_interleave_size @@ -1857,6 +1904,7 @@ class MLACommonMetadataBuilder(AttentionMetadataBuilder[M]): common_attn_metadata, decode_threshold=self.reorder_batch_threshold, require_uniform=(self.query_len_support != QueryLenSupport.VARLEN), + treat_short_extends_as_decodes=not self.use_pcp, ) ) diff --git a/vllm/model_executor/layers/attention/pcp.py b/vllm/model_executor/layers/attention/pcp.py new file mode 100644 index 00000000000..75ab1c9e8e1 --- /dev/null +++ b/vllm/model_executor/layers/attention/pcp.py @@ -0,0 +1,92 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import torch + +from vllm.distributed.parallel_state import ( + get_pcp_group, + get_tp_group, +) + + +def _gather_prefill_cache_inputs( + tensors: tuple[torch.Tensor, ...], + slot_mapping: torch.Tensor, + num_decode_tokens: int, +) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]: + """Keep replicated decode writes local and gather partitioned prefills.""" + local_num_tokens = tensors[0].shape[0] + assert all(tensor.shape[0] == local_num_tokens for tensor in tensors) + assert 0 <= num_decode_tokens <= local_num_tokens + + if num_decode_tokens == local_num_tokens: + return tensors, slot_mapping[:num_decode_tokens] + + pcp_group = get_pcp_group() + gathered_prefills = tuple( + pcp_group.all_gather(tensor[num_decode_tokens:].contiguous(), dim=0) + for tensor in tensors + ) + pcp_size = pcp_group.world_size + gathered_slot_mapping = slot_mapping[: pcp_size * local_num_tokens] + if num_decode_tokens == 0: + return gathered_prefills, gathered_slot_mapping + + cache_inputs = tuple( + torch.cat((tensor[:num_decode_tokens], gathered_prefill), dim=0) + for tensor, gathered_prefill in zip(tensors, gathered_prefills) + ) + rank_slot_mappings = gathered_slot_mapping.view(pcp_size, local_num_tokens) + cache_slot_mapping = torch.cat( + ( + rank_slot_mappings[0, :num_decode_tokens], + rank_slot_mappings[:, num_decode_tokens:].flatten(), + ) + ) + return cache_inputs, cache_slot_mapping + + +def maybe_gather_mla_latent_cache_inputs( + kv_c_normed: torch.Tensor, + k_pe: torch.Tensor, + slot_mapping: torch.Tensor | None, + num_decode_tokens: int | None, + use_pcp: bool, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]: + if not use_pcp or num_decode_tokens is None: + return kv_c_normed, k_pe, slot_mapping + assert slot_mapping is not None + num_tokens = kv_c_normed.shape[0] + k_pe_flat = k_pe.reshape(num_tokens, -1) + (cache_kv_c, cache_k_pe_flat), cache_slot_mapping = _gather_prefill_cache_inputs( + (kv_c_normed, k_pe_flat), + slot_mapping, + num_decode_tokens, + ) + cache_k_pe = cache_k_pe_flat.view(-1, *k_pe.shape[1:]) + return cache_kv_c, cache_k_pe, cache_slot_mapping + + +def maybe_gather_indexer_k( + k: torch.Tensor, + slot_mapping: torch.Tensor, + num_decode_tokens: int, + use_pcp: bool, +) -> tuple[torch.Tensor, torch.Tensor]: + if not use_pcp: + return k, slot_mapping + (cache_k,), cache_slot_mapping = _gather_prefill_cache_inputs( + (k,), slot_mapping, num_decode_tokens + ) + return cache_k, cache_slot_mapping + + +def finalize_mla_pcp_decode( + output: torch.Tensor, + num_heads: int, +) -> torch.Tensor: + if output.shape[1] < num_heads: + output = get_pcp_group().all_gather(output, dim=1) + elif output.shape[1] > num_heads: + head_start = get_tp_group().rank_in_group * num_heads + output = output[:, head_start : head_start + num_heads] + return output diff --git a/vllm/model_executor/layers/attention/sparse_mla_attention.py b/vllm/model_executor/layers/attention/sparse_mla_attention.py index 35bbd4d3d4f..19cad7986bf 100644 --- a/vllm/model_executor/layers/attention/sparse_mla_attention.py +++ b/vllm/model_executor/layers/attention/sparse_mla_attention.py @@ -59,13 +59,12 @@ class SparseMLACommonMetadataBuilder(AttentionMetadataBuilder[T]): device=device, ) parallel_config = vllm_config.parallel_config + self.use_pcp = parallel_config.prefill_context_parallel_size > 1 try: self.dcp_world_size = get_dcp_group().world_size - self.dcp_rank = get_dcp_group().rank_in_group except AssertionError: # DCP might not be initialized in testing self.dcp_world_size = 1 - self.dcp_rank = 0 self.cp_kv_cache_interleave_size = parallel_config.cp_kv_cache_interleave_size self.dcp_local_block_size = self.cp_kv_cache_interleave_size self.dcp_virtual_block_size = self.dcp_local_block_size * self.dcp_world_size @@ -174,6 +173,7 @@ class SparseMLACommonMetadataBuilder(AttentionMetadataBuilder[T]): num_decodes, num_prefills, num_decode_tokens, _ = split_decodes_and_prefills( common_attn_metadata, decode_threshold=self.reorder_batch_threshold or 1, + treat_short_extends_as_decodes=not self.use_pcp, require_uniform=self.require_uniform_decodes, ) ( diff --git a/vllm/model_executor/layers/fused_moe/config.py b/vllm/model_executor/layers/fused_moe/config.py index 39b14fc249e..02e6c51116d 100644 --- a/vllm/model_executor/layers/fused_moe/config.py +++ b/vllm/model_executor/layers/fused_moe/config.py @@ -1050,7 +1050,9 @@ class FusedMoEParallelConfig: @property def use_all2all_kernels(self): - return self.use_ep and (self.dp_size > 1 or self.is_sequence_parallel) + return self.use_ep and ( + self.dp_size > 1 or self.pcp_size > 1 or self.is_sequence_parallel + ) @property def use_deepep_ht_kernels(self): diff --git a/vllm/model_executor/layers/fused_moe/runner/moe_runner.py b/vllm/model_executor/layers/fused_moe/runner/moe_runner.py index cf00fa452a9..7942957f327 100644 --- a/vllm/model_executor/layers/fused_moe/runner/moe_runner.py +++ b/vllm/model_executor/layers/fused_moe/runner/moe_runner.py @@ -752,18 +752,12 @@ class MoERunner(MoERunnerInterface): assert len(result) == 2 hidden_states, router_logits = result - # NOTE: Similar with DP, PCP also needs dispatch and combine. For - # simplicity, AgRsAll2All was added separately for PCP here. Maybe - # we should modify All2AllManager abstraction to better support PCP. - if self.moe_config.pcp_size > 1: - hidden_states = get_pcp_group().all_gather( - hidden_states, - dim=0, - ) - router_logits = get_pcp_group().all_gather( - router_logits, - dim=0, - ) + if ( + self.moe_config.pcp_size > 1 + and not self.moe_config.moe_parallel_config.use_all2all_kernels + ): + hidden_states = get_pcp_group().all_gather(hidden_states, dim=0) + router_logits = get_pcp_group().all_gather(router_logits, dim=0) return hidden_states, router_logits @@ -777,11 +771,11 @@ class MoERunner(MoERunnerInterface): hidden_states, self.moe_config.is_sequence_parallel ) - if self.moe_config.pcp_size > 1: - hidden_states = get_pcp_group().reduce_scatter( - hidden_states, - dim=0, - ) + if ( + self.moe_config.pcp_size > 1 + and not self.moe_config.moe_parallel_config.use_all2all_kernels + ): + hidden_states = get_pcp_group().reduce_scatter(hidden_states, dim=0) if self.shared_experts is not None: assert shared_output is not None diff --git a/vllm/model_executor/layers/sparse_attn_indexer.py b/vllm/model_executor/layers/sparse_attn_indexer.py index ceb52e5d329..5b8e2bf008e 100644 --- a/vllm/model_executor/layers/sparse_attn_indexer.py +++ b/vllm/model_executor/layers/sparse_attn_indexer.py @@ -9,10 +9,11 @@ from vllm import _custom_ops as ops from vllm._aiter_ops import rocm_aiter_ops from vllm.compilation.breakable_cudagraph import eager_break_during_capture from vllm.config import get_current_vllm_config -from vllm.distributed import get_dcp_group +from vllm.distributed import get_dcp_group, get_pcp_group from vllm.forward_context import get_forward_context from vllm.logger import init_logger from vllm.model_executor.custom_op import CustomOp +from vllm.model_executor.layers.attention.pcp import maybe_gather_indexer_k from vllm.model_executor.layers.quantization.utils.quant_utils import ( get_fp8_min_max, ) @@ -308,6 +309,7 @@ def sparse_attn_indexer( total_seq_lens: int, topk_indices_buffer: torch.Tensor, skip_k_cache_insert: bool, + use_pcp: bool, use_fp4_cache: bool = False, dcp_rank: int = 0, dcp_world_size: int = 1, @@ -354,6 +356,7 @@ def sparse_attn_indexer( total_seq_lens, topk_indices_buffer, skip_k_cache_insert, + use_pcp, use_fp4_cache, ) attn_metadata_narrowed = attn_metadata[k_cache_prefix] @@ -373,18 +376,28 @@ def sparse_attn_indexer( # During speculative decoding, k may be padded to the CUDA graph batch # size while slot_mapping only covers actual tokens. Truncate k to avoid # out-of-bounds reads in the kernel. + # Keep PCP padding so every rank contributes the same all-gather shape. num_tokens = slot_mapping.shape[0] + if use_pcp: + num_tokens //= get_pcp_group().world_size if k is not None: k = k[:num_tokens] if not skip_k_cache_insert: + assert k is not None + k, slot_mapping_for_cache = maybe_gather_indexer_k( + k, + slot_mapping, + num_decode_tokens, + use_pcp, + ) # scale_fmt can be None, but the function expects str assert scale_fmt is not None assert not use_fp4_cache, "Unfused FP4 Insert is not supported yet" ops.indexer_k_quant_and_cache( k, kv_cache, - slot_mapping, + slot_mapping_for_cache, quant_block_size, scale_fmt, ) @@ -498,7 +511,14 @@ def sparse_attn_indexer( assert decode_metadata is not None kv_cache = kv_cache_as_quant_view(kv_cache, head_dim, use_fp4_cache) decode_lens = decode_metadata.decode_lens - if decode_metadata.requires_padding: + if num_decode_tokens == 0: + padded_q_quant_decode_tokens = q_quant[:1].reshape(1, 1, *q_quant.shape[1:]) + padded_q_scale = ( + q_scale[:1].reshape(1, 1, *q_scale.shape[1:]) + if q_scale is not None + else None + ) + elif decode_metadata.requires_padding: # pad in edge case where we have short chunked prefill length < # decode_threshold since we unstrictly split # prefill and decode by decode_threshold @@ -663,6 +683,7 @@ def sparse_attn_indexer_fake( total_seq_lens: int, topk_indices_buffer: torch.Tensor | None, skip_k_cache_insert: bool, + use_pcp: bool, use_fp4_cache: bool = False, dcp_rank: int = 0, dcp_world_size: int = 1, @@ -725,6 +746,7 @@ class SparseAttnIndexer(CustomOp): self.dcp_world_size = parallel_config.decode_context_parallel_size self.dcp_rank = get_dcp_group().rank_in_group if self.dcp_world_size > 1 else 0 self.cp_kv_cache_interleave_size = parallel_config.cp_kv_cache_interleave_size + self.use_pcp = parallel_config.prefill_context_parallel_size > 1 if current_platform.is_cuda() and not has_deep_gemm(): raise RuntimeError( "Sparse Attention Indexer CUDA op requires DeepGEMM support in " @@ -777,6 +799,7 @@ class SparseAttnIndexer(CustomOp): self.max_total_seq_len, self.topk_indices_buffer, self.skip_k_cache_insert, + self.use_pcp, self.use_fp4_cache, self.dcp_rank, self.dcp_world_size, diff --git a/vllm/v1/attention/backend.py b/vllm/v1/attention/backend.py index a82d49dadb2..6f23cff0b1e 100644 --- a/vllm/v1/attention/backend.py +++ b/vllm/v1/attention/backend.py @@ -281,6 +281,13 @@ class AttentionBackend(ABC): def supports_kv_connector(cls) -> bool: return True + @classmethod + def supports_pcp(cls) -> bool: + try: + return cls.get_impl_cls().supports_pcp + except NotImplementedError: + return False + @classmethod def supports_attn_type(cls, attn_type: str) -> bool: """Check if backend supports a given attention type. @@ -327,6 +334,7 @@ class AttentionBackend(ABC): use_non_causal: bool = False, use_batch_invariant: bool = False, use_kv_connector: bool = False, + use_pcp: bool = False, ) -> list[str]: invalid_reasons = [] if not cls.supports_head_size(head_size): @@ -367,6 +375,8 @@ class AttentionBackend(ABC): invalid_reasons.append("batch invariance not supported") if use_kv_connector and not cls.supports_kv_connector(): invalid_reasons.append("KV connector not supported") + if use_pcp and not cls.supports_pcp(): + invalid_reasons.append("PCP not supported") combination_reason = cls.supports_combination( head_size, dtype, @@ -858,8 +868,8 @@ class AttentionImplBase(ABC, Generic[T]): except AssertionError: self.pcp_world_size = 1 self.pcp_rank = 0 - self.total_cp_world_size = self.pcp_world_size * self.dcp_world_size - self.total_cp_rank = self.pcp_rank * self.dcp_world_size + self.dcp_rank + self.total_cp_world_size = self.dcp_world_size + self.total_cp_rank = self.dcp_rank self.need_to_return_lse_for_decode = ( self.dcp_world_size > 1 and self.can_return_lse_for_decode @@ -987,6 +997,8 @@ class AttentionImpl(AttentionImplBase[T], Generic[T]): class MLAAttentionImpl(AttentionImplBase[T], Generic[T]): """MLA attention implementation with forward_mqa and forward_mha methods.""" + supports_pcp: bool = True + @abstractmethod def __init__( self, diff --git a/vllm/v1/attention/backends/mla/indexer.py b/vllm/v1/attention/backends/mla/indexer.py index 2f578b46775..bf76cbeda85 100644 --- a/vllm/v1/attention/backends/mla/indexer.py +++ b/vllm/v1/attention/backends/mla/indexer.py @@ -6,7 +6,7 @@ import torch import vllm.envs as envs from vllm.config import VllmConfig -from vllm.distributed import get_dcp_group +from vllm.distributed import get_dcp_group, get_pcp_group from vllm.logger import init_logger from vllm.platforms import current_platform from vllm.triton_utils import tl, triton @@ -29,7 +29,7 @@ from vllm.v1.attention.backends.utils import ( split_decodes_and_prefills, ) from vllm.v1.kv_cache_interface import AttentionSpec, MLAAttentionSpec -from vllm.v1.worker.cp_utils import get_total_cp_world_size +from vllm.v1.worker.cp_utils import get_kv_cache_shard_count logger = init_logger(__name__) @@ -109,7 +109,7 @@ def split_indexer_prefill_chunks( end += 1 req_slice = slice(start + request_offset, end + request_offset) - max_q = max(1, max_logits_elems // chunk_n) if chunk_n > 0 else chunk_m + max_q = max(1, max_logits_elems // chunk_n) if chunk_n > 0 else max(1, chunk_m) for q_off in range(0, chunk_m, max_q): sub_m = min(max_q, chunk_m - q_off) chunks.append((req_slice, slice(q_off, q_off + sub_m))) @@ -118,6 +118,10 @@ def split_indexer_prefill_chunks( class DeepseekV32IndexerBackend(AttentionBackend): + @classmethod + def supports_pcp(cls) -> bool: + return True + @staticmethod def get_name() -> str: return "DEEPSEEK_V32_INDEXER" @@ -256,6 +260,8 @@ class DeepseekV32IndexerMetadataBuilder(AttentionMetadataBuilder): parallel_config = self.vllm_config.parallel_config self.dcp_world_size = parallel_config.decode_context_parallel_size self.dcp_rank = get_dcp_group().rank_in_group if self.dcp_world_size > 1 else 0 + self.pcp_world_size = parallel_config.prefill_context_parallel_size + self.use_pcp = self.pcp_world_size > 1 self.cp_kv_cache_interleave_size = parallel_config.cp_kv_cache_interleave_size # The DCP sparse-indexer code is parameterized by interleave size, but # interleave > 1 is not yet validated end-to-end (gsm8k parity fails), @@ -339,7 +345,7 @@ class DeepseekV32IndexerMetadataBuilder(AttentionMetadataBuilder): ) max_num_blocks_per_req = cdiv( self.vllm_config.model_config.max_model_len, - self.kv_cache_spec.block_size * get_total_cp_world_size(), + self.kv_cache_spec.block_size * get_kv_cache_shard_count(), ) self.expanded_block_table_buffer = torch.zeros( ( @@ -570,6 +576,7 @@ class DeepseekV32IndexerMetadataBuilder(AttentionMetadataBuilder): common_attn_metadata, decode_threshold=self.decode_threshold, require_uniform=not self.use_flattening, + treat_short_extends_as_decodes=not self.use_pcp, ) ) @@ -579,6 +586,9 @@ class DeepseekV32IndexerMetadataBuilder(AttentionMetadataBuilder): compressed_slot_mapping = slot_mapping compressed_seq_lens = seq_lens if self.compress_ratio > 1: + padded_num_tokens = num_tokens + if self.pcp_world_size > 1: + padded_num_tokens = slot_mapping.shape[0] // self.pcp_world_size compressed_slot_mapping = get_compressed_slot_mapping( num_tokens, query_start_loc, @@ -588,6 +598,11 @@ class DeepseekV32IndexerMetadataBuilder(AttentionMetadataBuilder): self.compress_ratio, out=self.compressed_slot_mapping_buffer, ) + if self.pcp_world_size > 1: + compressed_slot_mapping = get_pcp_group().all_gather( + self.compressed_slot_mapping_buffer[:padded_num_tokens], + dim=0, + ) compressed_seq_lens = seq_lens // self.compress_ratio prefill_metadata = None diff --git a/vllm/v1/attention/backends/utils.py b/vllm/v1/attention/backends/utils.py index c8c9a7334a2..c6e38c7fbfe 100644 --- a/vllm/v1/attention/backends/utils.py +++ b/vllm/v1/attention/backends/utils.py @@ -612,7 +612,9 @@ def split_decodes_and_prefills( # check if we are in a padded uniform batch; this is used for full-CGs, some # requests may have a query length of 0 but since they are padding its fine # to treat them as decodes (ensures num_decodes matches the captured size) - if torch.all((query_lens == query_lens[0]) | (query_lens == 0)): + if treat_short_extends_as_decodes and torch.all( + (query_lens == query_lens[0]) | (query_lens == 0) + ): return num_reqs, 0, num_tokens, 0 # all decodes is_prefill = query_lens != query_lens[0] else: diff --git a/vllm/v1/attention/selector.py b/vllm/v1/attention/selector.py index 387eac34c9f..734136334aa 100644 --- a/vllm/v1/attention/selector.py +++ b/vllm/v1/attention/selector.py @@ -36,6 +36,7 @@ class AttentionSelectorConfig(NamedTuple): use_non_causal: bool = False use_batch_invariant: bool = False use_kv_connector: bool = False + use_pcp: bool = False def __repr__(self): return ( @@ -52,7 +53,8 @@ class AttentionSelectorConfig(NamedTuple): f"has_sliding_window={self.has_sliding_window}, " f"use_non_causal={self.use_non_causal}, " f"use_batch_invariant={self.use_batch_invariant}, " - f"use_kv_connector={self.use_kv_connector})" + f"use_kv_connector={self.use_kv_connector}, " + f"use_pcp={self.use_pcp})" ) @@ -150,6 +152,7 @@ def get_attn_backend( use_non_causal=vllm_config.attention_config.use_non_causal, use_batch_invariant=envs.VLLM_BATCH_INVARIANT, use_kv_connector=use_kv_connector, + use_pcp=vllm_config.parallel_config.prefill_context_parallel_size > 1, ) # A per-KV-group override (keyed by KVCacheSpecKind) takes precedence over diff --git a/vllm/v1/core/kv_cache_coordinator.py b/vllm/v1/core/kv_cache_coordinator.py index 7d482192d54..df8769c3f3a 100644 --- a/vllm/v1/core/kv_cache_coordinator.py +++ b/vllm/v1/core/kv_cache_coordinator.py @@ -472,8 +472,6 @@ class UnitaryKVCacheCoordinator(KVCacheCoordinator): self.pcp_world_size = pcp_world_size if dcp_world_size > 1: self.block_size *= dcp_world_size - if pcp_world_size > 1: - self.block_size *= pcp_world_size # For models using only Mamba, block_size is set to max_model_len when # prefix caching is disabled, and hash_block_size validation is skipped. assert not enable_caching or (hash_block_size == self.block_size), ( diff --git a/vllm/v1/core/kv_cache_utils.py b/vllm/v1/core/kv_cache_utils.py index dc5ecf3f4d2..d4970d91b2d 100644 --- a/vllm/v1/core/kv_cache_utils.py +++ b/vllm/v1/core/kv_cache_utils.py @@ -631,8 +631,8 @@ def resolve_kv_cache_block_sizes( - ``scheduler_block_size`` is the token-alignment invariant used by the scheduler (e.g. for ``num_computed_tokens`` rounding). Single group: - ``cache_config.block_size * dcp * pcp``. Multiple groups: LCM of every - group's effective block size. Attention groups are scaled by DCP/PCP; + ``cache_config.block_size * dcp``. Multiple groups: LCM of every + group's effective block size. Attention groups are scaled by DCP; Mamba groups keep their full per-rank state and are not scaled. - ``hash_block_size`` is the granularity at which ``Request.block_hashes`` is computed. Single group: equals scheduler block size. Multiple groups: @@ -644,15 +644,14 @@ def resolve_kv_cache_block_sizes( """ cache_config = vllm_config.cache_config dcp = vllm_config.parallel_config.decode_context_parallel_size - pcp = vllm_config.parallel_config.prefill_context_parallel_size groups = kv_cache_config.kv_cache_groups - if len(groups) <= 1: # Single group: block_size * dcp * pcp - bs = cache_config.block_size * dcp * pcp + if len(groups) <= 1: + bs = cache_config.block_size * dcp return bs, bs group_block_sizes = [ - g.kv_cache_spec.block_size * dcp * pcp + g.kv_cache_spec.block_size * dcp if isinstance(g.kv_cache_spec, AttentionSpec) else g.kv_cache_spec.block_size for g in groups diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index de076e230e4..46a999ef7c5 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -271,7 +271,7 @@ class Scheduler(SchedulerInterface): log_stats=self.log_stats, enable_kv_cache_events=self.enable_kv_cache_events, dcp_world_size=self.dcp_world_size, - pcp_world_size=self.pcp_world_size, + pcp_world_size=1, scheduler_block_size=self.block_size, hash_block_size=hash_block_size, metrics_collector=self.kv_metrics_collector, diff --git a/vllm/v1/core/single_type_kv_cache_manager.py b/vllm/v1/core/single_type_kv_cache_manager.py index 24404839094..f8578c68a38 100644 --- a/vllm/v1/core/single_type_kv_cache_manager.py +++ b/vllm/v1/core/single_type_kv_cache_manager.py @@ -75,8 +75,8 @@ class SingleTypeKVCacheManager(ABC): self.block_size = kv_cache_spec.block_size self.dcp_world_size = dcp_world_size self.pcp_world_size = pcp_world_size - if dcp_world_size * pcp_world_size > 1: - self.block_size *= dcp_world_size * pcp_world_size + if dcp_world_size > 1: + self.block_size *= dcp_world_size self.kv_cache_spec = kv_cache_spec self.block_pool = block_pool self.enable_caching = enable_caching @@ -674,10 +674,10 @@ class FullAttentionManager(SingleTypeKVCacheManager): "and chunked local attention groups" ) block_size = kv_cache_spec.block_size - if dcp_world_size * pcp_world_size > 1: - # DCP/PCP shard each block's KV across ranks; hashes must be - # viewed at the sharded (scaled) block size. - block_size *= dcp_world_size * pcp_world_size + if dcp_world_size > 1: + # DCP shards each block's KV across ranks; hashes must be viewed at + # the sharded block size. + block_size *= dcp_world_size block_hashes = resolve_block_hashes( block_hashes, block_pool.hash_block_size, diff --git a/vllm/v1/kv_cache_interface.py b/vllm/v1/kv_cache_interface.py index d7f195e35b8..7a44d4e1d25 100644 --- a/vllm/v1/kv_cache_interface.py +++ b/vllm/v1/kv_cache_interface.py @@ -218,14 +218,9 @@ class AttentionSpec(KVCacheSpec): ) def max_num_blocks_per_req(self, vllm_config: VllmConfig, max_len: int) -> int: - # Attention KV is token-interleaved across DCP/PCP ranks, so each rank - # only stores max_len // (dcp * pcp) tokens per request. parallel_config = vllm_config.parallel_config - total_cp_size = ( - parallel_config.decode_context_parallel_size - * parallel_config.prefill_context_parallel_size - ) - return cdiv(max_len, self.block_size * total_cp_size) + kv_shard_count = parallel_config.decode_context_parallel_size + return cdiv(max_len, self.block_size * kv_shard_count) @dataclass(frozen=True, kw_only=True) @@ -263,11 +258,8 @@ class FullAttentionSpec(AttentionSpec): def max_memory_usage_bytes(self, vllm_config: VllmConfig) -> int: max_model_len = vllm_config.model_config.max_model_len dcp_world_size = vllm_config.parallel_config.decode_context_parallel_size - pcp_world_size = vllm_config.parallel_config.prefill_context_parallel_size - # Note(hc): each dcp rank only need save - # (max_model_len//dcp_world_size) tokens locally. - if dcp_world_size * pcp_world_size > 1: - max_model_len = cdiv(max_model_len, dcp_world_size * pcp_world_size) + if dcp_world_size > 1: + max_model_len = cdiv(max_model_len, dcp_world_size) return cdiv(max_model_len, self.block_size) * self.page_size_bytes @classmethod diff --git a/vllm/v1/simple_kv_offload/manager.py b/vllm/v1/simple_kv_offload/manager.py index 1d2b8a6b7e9..515a2f19b0e 100644 --- a/vllm/v1/simple_kv_offload/manager.py +++ b/vllm/v1/simple_kv_offload/manager.py @@ -83,8 +83,7 @@ class SimpleCPUOffloadScheduler: and vllm_config.kv_events_config.enable_kv_cache_events ) dcp_world_size = vllm_config.parallel_config.decode_context_parallel_size - pcp_world_size = vllm_config.parallel_config.prefill_context_parallel_size - self.cp_world_size = dcp_world_size * pcp_world_size + self.cp_world_size = dcp_world_size self.block_size = scheduler_block_size self.hash_block_size = hash_block_size assert self.block_size % self.hash_block_size == 0 @@ -127,7 +126,7 @@ class SimpleCPUOffloadScheduler: enable_caching=True, enable_kv_cache_events=self.enable_kv_cache_events, dcp_world_size=dcp_world_size, - pcp_world_size=pcp_world_size, + pcp_world_size=1, scheduler_block_size=self.block_size, hash_block_size=self.hash_block_size, ) @@ -354,7 +353,7 @@ class SimpleCPUOffloadScheduler: ) assert num_external_tokens % g_block_size == 0, ( f"num_external_tokens={num_external_tokens} not aligned to " - f"group {g} effective block_size={g_block_size}" + f"group {g} block_size={g_block_size}" ) n_take_g = num_external_tokens // g_block_size cpu_hit_blocks.append(cpu_hit_blocks_full[g][:n_take_g]) diff --git a/vllm/v1/worker/block_table.py b/vllm/v1/worker/block_table.py index d40887879fc..332eda4cbdf 100644 --- a/vllm/v1/worker/block_table.py +++ b/vllm/v1/worker/block_table.py @@ -163,8 +163,6 @@ class BlockTable: return assert self.slot_mapping_mode == SlotMappingMode.TOKEN_TO_KV_SLOT - total_cp_world_size = self.pcp_world_size * self.dcp_world_size - total_cp_rank = self.pcp_rank * self.dcp_world_size + self.dcp_rank _compute_slot_mapping_kernel[(num_reqs + 1,)]( num_tokens, self.max_num_batched_tokens, @@ -176,8 +174,8 @@ class BlockTable: self.slot_mapping.gpu, KV_CACHE_BLOCK_SIZE=self.kv_cache_block_size, BLOCKS_PER_KV_BLOCK=self.blocks_per_kv_block, - TOTAL_CP_WORLD_SIZE=total_cp_world_size, - TOTAL_CP_RANK=total_cp_rank, + TOTAL_CP_WORLD_SIZE=self.dcp_world_size, + TOTAL_CP_RANK=self.dcp_rank, CP_KV_CACHE_INTERLEAVE_SIZE=self.cp_kv_cache_interleave_size, PAD_ID=PAD_SLOT_ID, BLOCK_SIZE=1024, diff --git a/vllm/v1/worker/cp_utils.py b/vllm/v1/worker/cp_utils.py index 11edd86a5db..92d8383c1f1 100644 --- a/vllm/v1/worker/cp_utils.py +++ b/vllm/v1/worker/cp_utils.py @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Any, cast import torch from vllm.config import VllmConfig, get_layers_from_vllm_config -from vllm.distributed import get_dcp_group, get_pcp_group +from vllm.distributed import get_dcp_group from vllm.logger import init_logger from vllm.v1.attention.backend import CommonAttentionMetadata from vllm.v1.attention.backends.utils import split_decodes_prefills_and_extends @@ -27,6 +27,13 @@ def check_attention_cp_compatibility(vllm_config: VllmConfig) -> None: layer_type = cast(type[Any], AttentionLayerBase) layers = get_layers_from_vllm_config(vllm_config, layer_type) for layer in layers.values(): + get_attn_backend = getattr(layer, "get_attn_backend", None) + if pcp_size > 1 and get_attn_backend is not None: + backend = get_attn_backend() + assert backend.supports_pcp(), ( + "PCP requires attention backend support, " + f"but {backend.get_name()} does not support PCP." + ) layer_impl = getattr(layer, "impl", None) if layer_impl is None: continue @@ -44,26 +51,14 @@ def check_attention_cp_compatibility(vllm_config: VllmConfig) -> None: "--attention-backend or disable DCP." ) - if pcp_size > 1: - assert layer_impl.supports_pcp, ( - "PCP requires attention impls' support, " - f"but the impl {layer_impl.__class__.__name__} " - "does not support PCP." - ) - -def get_total_cp_world_size(): - try: - pcp_world_size = get_pcp_group().world_size - except AssertionError: - # PCP might not be initialized in testing - pcp_world_size = 1 +def get_kv_cache_shard_count() -> int: try: dcp_world_size = get_dcp_group().world_size except AssertionError: # DCP might not be initialized in testing dcp_world_size = 1 - return dcp_world_size * pcp_world_size + return dcp_world_size def get_dcp_dummy_context_len( diff --git a/vllm/v1/worker/gpu/attn_utils.py b/vllm/v1/worker/gpu/attn_utils.py index 5c07860b3ba..a97a4e39597 100644 --- a/vllm/v1/worker/gpu/attn_utils.py +++ b/vllm/v1/worker/gpu/attn_utils.py @@ -578,6 +578,7 @@ def build_attn_metadata( seq_lens_cpu_upper_bound: torch.Tensor | None = None, dcp_local_seq_lens: torch.Tensor | None = None, positions: torch.Tensor | None = None, + is_prefilling: torch.Tensor | None = None, mm_req_doc_ranges: dict[int, list[tuple[int, int]]] | None = None, model_specific_attn_metadata: ModelSpecificAttnMetadata | None = None, for_cudagraph_capture: bool = False, @@ -605,6 +606,11 @@ def build_attn_metadata( if model_specific_attn_metadata is not None else {} ) + # Model-specific metadata (e.g. Mamba hybrid) may supply its own + # padding-aware is_prefilling, which takes precedence over the default. + group_is_prefilling = common_attn_metadata_extra_kwargs.pop( + "is_prefilling", is_prefilling + ) common_attn_metadata = CommonAttentionMetadata( query_start_loc=query_start_loc_gpu, query_start_loc_cpu=query_start_loc_cpu, @@ -619,6 +625,7 @@ def build_attn_metadata( causal=group_causal, dcp_local_seq_lens=dcp_local_seq_lens, positions=positions, + is_prefilling=group_is_prefilling, mm_req_doc_ranges=mm_req_doc_ranges, rswa_prefix_lens=rswa_prefix_lens, **common_attn_metadata_extra_kwargs, diff --git a/vllm/v1/worker/gpu/block_table.py b/vllm/v1/worker/gpu/block_table.py index 8d41ba5a36a..22c4afc11bc 100644 --- a/vllm/v1/worker/gpu/block_table.py +++ b/vllm/v1/worker/gpu/block_table.py @@ -135,20 +135,28 @@ class BlockTables: self, idx_mapping: torch.Tensor, num_reqs_padded: int, + out: tuple[torch.Tensor, ...] | None = None, + out_ptrs: torch.Tensor | None = None, ) -> tuple[torch.Tensor, ...]: + if out is None: + out = tuple(self.input_block_tables) + out_ptrs = self.input_block_table_ptrs + else: + assert out_ptrs is not None + assert len(out) == self.num_kv_cache_groups num_reqs = idx_mapping.shape[0] # Launch kernel with num_reqs_padded to fuse zeroing of padded rows. _gather_block_tables_kernel[(self.num_kv_cache_groups, num_reqs_padded)]( idx_mapping, self.block_table_ptrs, - self.input_block_table_ptrs, + out_ptrs, self.block_table_strides, self.num_blocks.gpu, self.num_blocks.gpu.stride(0), num_reqs, BLOCK_SIZE=1024, # type: ignore ) - return tuple(bt[:num_reqs_padded] for bt in self.input_block_tables) + return tuple(bt[:num_reqs_padded] for bt in out) def get_dummy_block_tables(self, num_reqs: int) -> tuple[torch.Tensor, ...]: # NOTE(woosuk): The output may be used for CUDA graph capture. @@ -163,26 +171,28 @@ class BlockTables: query_start_loc: torch.Tensor, positions: torch.Tensor, num_tokens_padded: int, + out: torch.Tensor | None = None, ) -> torch.Tensor: num_reqs = idx_mapping.shape[0] num_groups = self.num_kv_cache_groups + slot_mappings = self.slot_mappings if out is None else out _compute_slot_mappings_kernel[(num_groups, num_reqs + 1)]( - self.max_num_batched_tokens, + slot_mappings.shape[1], idx_mapping, query_start_loc, positions, self.block_table_ptrs, self.block_table_strides, self.block_sizes_tensor, - self.slot_mappings, - self.slot_mappings.stride(0), + slot_mappings, + slot_mappings.stride(0), self.cp_rank, CP_SIZE=self.cp_size, CP_INTERLEAVE=self.cp_interleave, PAD_ID=PAD_SLOT_ID, TRITON_BLOCK_SIZE=1024, # type: ignore ) - return self.slot_mappings[:, :num_tokens_padded] + return slot_mappings[:, :num_tokens_padded] def get_dummy_slot_mappings(self, num_tokens: int) -> torch.Tensor: # Fill the entire slot_mappings tensor, not just the first `num_tokens` entries. diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 518d12a6b28..293b06ff97b 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -52,6 +52,7 @@ from vllm.v1.core.sched.output import GrammarOutput, SchedulerOutput from vllm.v1.kv_cache_interface import KVCacheConfig, MambaSpec from vllm.v1.outputs import DraftTokenIds, ModelRunnerOutput from vllm.v1.worker.cp_utils import check_attention_cp_compatibility +from vllm.v1.worker.gpu import pcp_manager as pcp from vllm.v1.worker.gpu.async_utils import AsyncOutput, AsyncPoolingOutput from vllm.v1.worker.gpu.attn_utils import ( build_slot_mappings_by_layer, @@ -205,6 +206,8 @@ class GPUModelRunner(LoRAModelRunnerMixin): # Draft tokens propagation - for spec-dec + struct outputs. self.draft_tokens_handler = DraftTokensHandler(self.device) + self.pcp_manager: pcp.PCPManager | None = None + # Pooling models. self.is_pooling_model = self.model_config.runner_type == "pooling" self.pooling_runner: PoolingRunner | None = None @@ -456,6 +459,13 @@ class GPUModelRunner(LoRAModelRunnerMixin): cp_rank=self.dcp_rank, cp_interleave=self.cp_interleave, ) + self.pcp_manager = pcp.maybe_build_pcp_manager( + self.vllm_config, + self.device, + self.supports_mm_inputs, + self.req_states, + self.block_tables, + ) initialize_mamba_ssu_backend( self.vllm_config.mamba_config, self.kv_cache_config ) @@ -1008,7 +1018,7 @@ class GPUModelRunner(LoRAModelRunnerMixin): # prompt_lens is only used in R-SWA case. prompt_lens = self.req_states.prompt_len.gpu[idx_mapping] - return InputBatch( + input_batch = InputBatch( req_ids=req_ids, num_reqs=num_reqs, num_reqs_after_padding=num_reqs_padded, @@ -1040,10 +1050,14 @@ class GPUModelRunner(LoRAModelRunnerMixin): has_structured_output_reqs=scheduler_output.has_structured_output_requests, prompt_lens=prompt_lens, ) + return pcp.maybe_partition_pcp_batch(self.pcp_manager, input_batch) def prepare_attn( self, input_batch: InputBatch ) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]: + if self.pcp_manager is not None: + return self.pcp_manager.prepare_attn(input_batch) + # Block tables: num_kv_cache_groups x [num_reqs_padded, max_num_blocks]. block_tables = self.block_tables.gather_block_tables( input_batch.idx_mapping, @@ -1063,8 +1077,8 @@ class GPUModelRunner(LoRAModelRunnerMixin): self, input_batch: InputBatch ) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]: block_tables = self.block_tables.get_dummy_block_tables(input_batch.num_reqs) - slot_mappings = self.block_tables.get_dummy_slot_mappings( - input_batch.num_tokens + slot_mappings = pcp.maybe_get_pcp_dummy_slot_mappings( + self.pcp_manager, self.block_tables, input_batch.num_tokens ) return block_tables, slot_mappings @@ -1412,6 +1426,10 @@ class GPUModelRunner(LoRAModelRunnerMixin): return ModelRunnerOutput.with_kv_conn_output_only(kv_connector_output) # Last rank: sample tokens + hidden_states, input_batch = pcp.maybe_restore_pcp_for_sampling( + self.pcp_manager, hidden_states, input_batch + ) + sampler_output, num_sampled, num_rejected = self.sample( hidden_states, input_batch, grammar_output ) diff --git a/vllm/v1/worker/gpu/model_states/default.py b/vllm/v1/worker/gpu/model_states/default.py index 18eb40640ad..053f37aaf10 100644 --- a/vllm/v1/worker/gpu/model_states/default.py +++ b/vllm/v1/worker/gpu/model_states/default.py @@ -151,7 +151,10 @@ class DefaultModelState(ModelState): # For piecewise cudagraphs and eager, use unpadded sizes. num_reqs = input_batch.num_reqs num_tokens = input_batch.num_tokens - query_start_loc_cpu = torch.from_numpy(input_batch.query_start_loc_np) + query_start_loc_cpu = torch.from_numpy( + input_batch.query_start_loc_np[: num_reqs + 1] + ) + query_start_loc_gpu = input_batch.query_start_loc[: num_reqs + 1] max_query_len = input_batch.num_scheduled_tokens.max().item() seq_lens_cpu_upper_bound = input_batch.seq_lens_cpu_upper_bound if for_capture: @@ -174,7 +177,7 @@ class DefaultModelState(ModelState): attn_groups=attn_groups, num_reqs=num_reqs, num_tokens=num_tokens, - query_start_loc_gpu=input_batch.query_start_loc, + query_start_loc_gpu=query_start_loc_gpu, query_start_loc_cpu=query_start_loc_cpu, max_query_len=max_query_len, seq_lens=input_batch.seq_lens, @@ -185,6 +188,7 @@ class DefaultModelState(ModelState): seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, dcp_local_seq_lens=input_batch.dcp_local_seq_lens, positions=input_batch.positions, + is_prefilling=torch.from_numpy(input_batch.is_prefilling_np), mm_req_doc_ranges=req_doc_ranges, for_cudagraph_capture=for_capture, rswa_prefix_lens=input_batch.prompt_lens, diff --git a/vllm/v1/worker/gpu/pcp_manager.py b/vllm/v1/worker/gpu/pcp_manager.py new file mode 100644 index 00000000000..f50e4874942 --- /dev/null +++ b/vllm/v1/worker/gpu/pcp_manager.py @@ -0,0 +1,680 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from dataclasses import dataclass, replace + +import numpy as np +import torch + +from vllm.config import CUDAGraphMode, VllmConfig +from vllm.distributed.parallel_state import get_dcp_group, get_pcp_group +from vllm.logger import init_logger +from vllm.v1.attention.backends.utils import PAD_SLOT_ID +from vllm.v1.worker.gpu.block_table import BlockTables +from vllm.v1.worker.gpu.buffer_utils import async_copy_to_gpu +from vllm.v1.worker.gpu.cp_utils import prepare_dcp_local_seq_lens +from vllm.v1.worker.gpu.input_batch import ( + InputBatch, + InputBuffers, + combine_sampled_and_draft_tokens, + prepare_pos_seq_lens, +) +from vllm.v1.worker.gpu.states import RequestState + +logger = init_logger(__name__) + + +@dataclass(frozen=True) +class RankSegment: + global_batch_req_idx: int + global_batch_slice: slice + rank_local_batch_slice: slice + + @property + def num_tokens(self) -> int: + return self.global_batch_slice.stop - self.global_batch_slice.start + + +class PCPManager: + """MRV2 PC batch manager. + + The model runner keeps the global scheduled batch. This manager rewrites only + the per-step InputBatch into rank-local DualChunkSwap rows and keeps the + global-batch view private to restore to the global batch shape before + sampling/postprocess. + """ + + def __init__( + self, + pcp_world_size: int, + pcp_rank: int, + device: torch.device, + req_states: RequestState | None = None, + max_num_reqs: int | None = None, + max_num_tokens: int | None = None, + block_tables: BlockTables | None = None, + dcp_world_size: int = 1, + dcp_rank: int = 0, + cp_interleave: int = 1, + ) -> None: + self.pcp_world_size = pcp_world_size + self.pcp_rank = pcp_rank + self.device = device + self.dcp_world_size = dcp_world_size + self.dcp_rank = dcp_rank + self.cp_interleave = cp_interleave + + self._global_batch: InputBatch | None = None + self._req_states = req_states + self._block_tables = block_tables + self._hidden_restore_idx: torch.Tensor | None = None + self._padded_gather_idx: torch.Tensor | None = None + self._gathered_kv_write_mask: torch.Tensor | None = None + self._pad_slot_id = torch.tensor(PAD_SLOT_ID, dtype=torch.int64, device=device) + + max_num_local_reqs = 2 * max_num_reqs if max_num_reqs is not None else None + self._input_buffers = ( + InputBuffers(max_num_local_reqs, max_num_tokens, device) + if max_num_local_reqs is not None and max_num_tokens is not None + else None + ) + self._local_req_idx = ( + torch.arange(max_num_local_reqs, dtype=torch.int32, device=device) + if max_num_local_reqs is not None + else None + ) + self._local_block_tables: tuple[torch.Tensor, ...] | None + self._local_block_table_ptrs: torch.Tensor | None + if block_tables is not None and max_num_local_reqs is not None: + self._local_block_tables = tuple( + table.new_zeros((max_num_local_reqs, table.shape[1])) + for table in block_tables.input_block_tables + ) + self._local_block_table_ptrs = torch.tensor( + [table.data_ptr() for table in self._local_block_tables], + dtype=torch.uint64, + device=device, + ) + else: + self._local_block_tables = None + self._local_block_table_ptrs = None + num_kv_cache_groups = ( + block_tables.num_kv_cache_groups if block_tables is not None else 0 + ) + self._global_batch_slot_mappings = ( + torch.empty( + num_kv_cache_groups, + max_num_tokens, + dtype=torch.int64, + device=device, + ) + if max_num_tokens is not None and num_kv_cache_groups > 0 + else None + ) + self._gathered_kv_slot_mappings = ( + torch.empty( + num_kv_cache_groups, + max_num_tokens * pcp_world_size, + dtype=torch.int64, + device=device, + ) + if max_num_tokens is not None and num_kv_cache_groups > 0 + else None + ) + + @staticmethod + def validate_config( + vllm_config: VllmConfig, + supports_mm_inputs: bool, + ) -> None: + parallel_config = vllm_config.parallel_config + model_config = vllm_config.model_config + pcp_size = parallel_config.prefill_context_parallel_size + if pcp_size <= 1: + return + + if not model_config.use_mla: + raise NotImplementedError("MRV2 PCP currently supports MLA models only.") + if parallel_config.pipeline_parallel_size > 1: + raise NotImplementedError("MRV2 PCP does not support PP yet.") + if model_config.is_encoder_decoder: + raise NotImplementedError( + "MRV2 PCP does not support encoder-decoder models yet." + ) + if supports_mm_inputs: + raise NotImplementedError("MRV2 PCP does not support MM inputs yet.") + if vllm_config.lora_config is not None: + raise NotImplementedError("MRV2 PCP does not support LoRA yet.") + if vllm_config.speculative_config is not None: + raise NotImplementedError( + "MRV2 PCP does not support speculative decoding yet." + ) + is_sparse_mla = hasattr(model_config.hf_text_config, "index_topk") + if ( + is_sparse_mla + and vllm_config.compilation_config.cudagraph_mode != CUDAGraphMode.NONE + ): + raise NotImplementedError( + "MRV2 sparse MLA PCP does not support CUDA graphs yet. " + "Set -cc.cudagraph_mode=NONE." + ) + if vllm_config.compilation_config.cudagraph_mode.has_full_cudagraphs(): + raise NotImplementedError("MRV2 PCP supports PIECEWISE CUDA graphs only.") + + @staticmethod + def _reorder_segments( + segments: list[RankSegment], + num_computed_tokens: np.ndarray, + is_prefilling: np.ndarray, + query_start_loc_np: np.ndarray, + ) -> list[RankSegment]: + """Move pure prefills last to match the batch ordering expected by + attention backends like MLA and sparse MLA. + """ + + def is_pure_prefill(segment: RankSegment) -> bool: + req_idx = segment.global_batch_req_idx + start_pos = ( + num_computed_tokens[req_idx] + + segment.global_batch_slice.start + - query_start_loc_np[req_idx] + ) + return is_prefilling[req_idx] and start_pos == 0 + + segments.sort(key=is_pure_prefill) + rank_offset = 0 + for index, segment in enumerate(segments): + segments[index] = replace( + segment, + rank_local_batch_slice=slice( + rank_offset, rank_offset + segment.num_tokens + ), + ) + rank_offset += segment.num_tokens + return segments + + def _get_rank_segments( + self, + rank: int, + num_scheduled_tokens: np.ndarray, + num_computed_tokens: np.ndarray, + is_prefilling: np.ndarray, + query_start_loc_np: np.ndarray, + ) -> list[RankSegment]: + """Build one rank's attention-compatible DualChunkSwap rows. + + PCP=4 partitions each prefill into eight chunks: + + full: | 0 | 1 | 2 | 3 | 4 | 5 | 6 | 7 | + rank 0: 0 7 + rank 1: 1 6 + rank 2: 2 5 + rank 3: 3 4 + """ + rank_segments = [] + rank_offset = 0 + num_chunks = 2 * self.pcp_world_size + for global_batch_req_idx, num_tokens in enumerate(num_scheduled_tokens): + query_len = int(num_tokens) + if query_len == 0: + continue + global_batch_start = int(query_start_loc_np[global_batch_req_idx]) + chunk_indices: tuple[int, ...] + if bool(is_prefilling[global_batch_req_idx]): + chunk_size = (query_len + num_chunks - 1) // num_chunks + chunk_indices = (rank, num_chunks - 1 - rank) + else: # decodes are replicated + chunk_size = query_len + chunk_indices = (0,) + + for chunk_idx in chunk_indices: + chunk_offset = chunk_idx * chunk_size + chunk_len = min(chunk_size, query_len - chunk_offset) + if chunk_len <= 0: + continue + chunk_start = global_batch_start + chunk_offset + rank_segments.append( + RankSegment( + global_batch_req_idx=global_batch_req_idx, + global_batch_slice=slice(chunk_start, chunk_start + chunk_len), + rank_local_batch_slice=slice( + rank_offset, rank_offset + chunk_len + ), + ) + ) + rank_offset += chunk_len + return self._reorder_segments( + rank_segments, + num_computed_tokens, + is_prefilling, + query_start_loc_np, + ) + + def _build_batch_layout( + self, + num_scheduled_tokens: np.ndarray, + num_computed_tokens: np.ndarray, + is_prefilling: np.ndarray, + query_start_loc_np: np.ndarray, + ) -> tuple[list[list[RankSegment]], list[int]]: + segments_by_rank = [] + per_rank_num_tokens = [] + for rank in range(self.pcp_world_size): + segments = self._get_rank_segments( + rank, + num_scheduled_tokens, + num_computed_tokens, + is_prefilling, + query_start_loc_np, + ) + num_rank_tokens = sum(segment.num_tokens for segment in segments) + segments_by_rank.append(segments) + per_rank_num_tokens.append(num_rank_tokens) + + # PCP=2 example: + # global batch: [A B C D E F G] + # rank 0 / rank 1: [A B G] / [C D E F] + # padded gathered: [A B G _ | C D E F] + # hidden_restore_idx: [0, 1, 4, 5, 6, 7, 2] + # padded_gather_idx: [0, 1, 6, 0, 2, 3, 4, 5] + # Therefore global = gathered[hidden_restore_idx] and + # padded_gathered = global[padded_gather_idx]. + hidden_restore_idx = np.empty(int(query_start_loc_np[-1]), dtype=np.int64) + padded_num_tokens = max(per_rank_num_tokens) + num_expanded_tokens = padded_num_tokens * self.pcp_world_size + padded_gather_idx = np.zeros(num_expanded_tokens, dtype=np.int64) + gathered_kv_write_mask = np.zeros(num_expanded_tokens, dtype=np.bool_) + for rank, segments in enumerate(segments_by_rank): + expanded_rank_offset = rank * padded_num_tokens + for segment in segments: + padded_gathered_slice = slice( + expanded_rank_offset + segment.rank_local_batch_slice.start, + expanded_rank_offset + segment.rank_local_batch_slice.stop, + ) + padded_gather_idx[padded_gathered_slice] = np.arange( + segment.global_batch_slice.start, + segment.global_batch_slice.stop, + dtype=np.int64, + ) + # Cache insertion pairs one slot entry with each rank's local decode. + if not bool(is_prefilling[segment.global_batch_req_idx]) and rank != 0: + continue + gathered_kv_write_mask[padded_gathered_slice] = True + hidden_restore_idx[segment.global_batch_slice] = np.arange( + padded_gathered_slice.start, + padded_gathered_slice.stop, + dtype=np.int64, + ) + + self._hidden_restore_idx = async_copy_to_gpu( + hidden_restore_idx, device=self.device + ) + self._padded_gather_idx = async_copy_to_gpu( + padded_gather_idx, device=self.device + ) + self._gathered_kv_write_mask = async_copy_to_gpu( + gathered_kv_write_mask, device=self.device + ) + return segments_by_rank, per_rank_num_tokens + + def partition_batch(self, input_batch: InputBatch) -> InputBatch: + assert self._req_states is not None + assert self._input_buffers is not None + req_states = self._req_states + input_buffers = self._input_buffers + if input_batch.num_draft_tokens > 0: + raise NotImplementedError("MRV2 PCP does not support spec decode yet.") + + global_batch = input_batch + self._global_batch = global_batch + + num_scheduled_tokens = global_batch.num_scheduled_tokens + num_computed_tokens = global_batch.num_computed_tokens_np + is_prefilling = global_batch.is_prefilling_np + + segments_by_rank, per_rank_num_tokens = self._build_batch_layout( + num_scheduled_tokens, + num_computed_tokens, + is_prefilling, + global_batch.query_start_loc_np, + ) + + local_segments = segments_by_rank[self.pcp_rank] + if not local_segments: + local_segments = [ + RankSegment( + global_batch_req_idx=0, + global_batch_slice=slice(0, 0), + rank_local_batch_slice=slice(0, 0), + ) + ] + + num_local_reqs = len(local_segments) + if num_local_reqs > input_buffers.max_num_reqs: + raise RuntimeError( + "PCP local request count exceeds the MRV2 input buffer size: " + f"{num_local_reqs} > {input_buffers.max_num_reqs}." + ) + + local_to_global_batch_req_idx_np = np.fromiter( + (segment.global_batch_req_idx for segment in local_segments), + dtype=np.int32, + count=num_local_reqs, + ) + local_start_pos_np = np.fromiter( + ( + num_computed_tokens[segment.global_batch_req_idx] + + segment.global_batch_slice.start + - global_batch.query_start_loc_np[segment.global_batch_req_idx] + for segment in local_segments + ), + dtype=np.int32, + count=num_local_reqs, + ) + local_num_scheduled_tokens = np.fromiter( + (segment.num_tokens for segment in local_segments), + dtype=np.int32, + count=num_local_reqs, + ) + local_to_global_req_idx_np = global_batch.idx_mapping_np[ + local_to_global_batch_req_idx_np + ] + local_req_ids = [ + global_batch.req_ids[global_batch_req_idx] + for global_batch_req_idx in local_to_global_batch_req_idx_np + ] + + num_local_tokens = int(local_num_scheduled_tokens.sum()) + num_local_tokens_padded = max(per_rank_num_tokens) + fresh_prefills = int( + np.count_nonzero(is_prefilling & (num_computed_tokens == 0)) + ) + continued_prefills = int( + np.count_nonzero(is_prefilling & (num_computed_tokens > 0)) + ) + logger.debug( + "PCP batch: rank=%d global_batch_reqs=%d fresh_prefills=%d " + "continued_prefills=%d decodes=%d local_reqs=%d " + "local_tokens=%d per_rank_tokens=%s", + self.pcp_rank, + global_batch.num_reqs, + fresh_prefills, + continued_prefills, + global_batch.num_reqs - fresh_prefills - continued_prefills, + num_local_reqs, + num_local_tokens, + per_rank_num_tokens, + ) + if num_local_tokens_padded > input_buffers.max_num_tokens: + raise RuntimeError( + "PCP local token count exceeds the MRV2 input buffer size: " + f"{num_local_tokens_padded} > {input_buffers.max_num_tokens}." + ) + rank_token_start = self.pcp_rank * num_local_tokens_padded + assert self._padded_gather_idx is not None + local_gather_idx = self._padded_gather_idx[ + rank_token_start : rank_token_start + num_local_tokens_padded + ] + torch.index_select( + global_batch.input_ids, + 0, + local_gather_idx, + out=input_buffers.input_ids[:num_local_tokens_padded], + ) + + local_query_start_loc_np = np.empty( + input_buffers.max_num_reqs + 1, dtype=np.int32 + ) + local_query_start_loc_np[0] = 0 + local_query_start_loc_out = local_query_start_loc_np[1 : num_local_reqs + 1] + np.cumsum(local_num_scheduled_tokens, out=local_query_start_loc_out) + local_query_start_loc_np[num_local_reqs + 1 :] = num_local_tokens + async_copy_to_gpu(local_query_start_loc_np, out=input_buffers.query_start_loc) + local_query_start_loc = input_buffers.query_start_loc[: num_local_reqs + 1] + + local_to_global_req_idx = async_copy_to_gpu( + local_to_global_req_idx_np, device=self.device + ) + local_start_pos = async_copy_to_gpu(local_start_pos_np, device=self.device) + + assert self._local_req_idx is not None + prepare_pos_seq_lens( + self._local_req_idx[:num_local_reqs], + local_query_start_loc, + local_start_pos, + input_buffers.positions, + input_buffers.seq_lens[:num_local_reqs], + ) + seq_lens = input_buffers.seq_lens[:num_local_reqs] + is_padding = input_buffers.is_padding[:num_local_tokens_padded] + is_padding[:num_local_tokens].fill_(False) + is_padding[num_local_tokens:].fill_(True) + if num_local_tokens_padded > num_local_tokens: + input_buffers.input_ids[:num_local_tokens_padded].masked_fill_( + is_padding, 0 + ) + input_buffers.positions[:num_local_tokens_padded].masked_fill_( + is_padding, 0 + ) + + total_num_logits = num_local_reqs if num_local_tokens > 0 else 0 + if total_num_logits > 0: + cu_num_logits_np = np.arange(num_local_reqs + 1, dtype=np.int32) + cu_num_logits = torch.arange( + num_local_reqs + 1, device=self.device, dtype=torch.int32 + ) + else: + cu_num_logits_np = np.zeros(num_local_reqs + 1, dtype=np.int32) + cu_num_logits = torch.zeros( + num_local_reqs + 1, device=self.device, dtype=torch.int32 + ) + logits_indices = combine_sampled_and_draft_tokens( + input_buffers.input_ids, + local_to_global_req_idx, + req_states.last_sampled_tokens, + local_query_start_loc, + seq_lens, + req_states.prefill_len.gpu, + req_states.draft_tokens, + cu_num_logits, + total_num_logits, + 1, + ) + + local_prefill_len_np = global_batch.prefill_len_np[ + local_to_global_batch_req_idx_np + ] + local_num_computed_prefill_tokens_np = np.minimum( + local_start_pos_np, local_prefill_len_np + ) + local_is_prefilling_np = ( + local_num_computed_prefill_tokens_np < local_prefill_len_np + ) + seq_lens_cpu_upper_bound_np = np.zeros(num_local_reqs, dtype=np.int32) + seq_lens_cpu_upper_bound_np[:] = local_start_pos_np + local_num_scheduled_tokens + + dcp_local_seq_lens = None + if self.dcp_world_size > 1: + prepare_dcp_local_seq_lens( + input_buffers.dcp_local_seq_lens, + seq_lens, + num_local_reqs, + self.dcp_world_size, + self.dcp_rank, + self.cp_interleave, + ) + dcp_local_seq_lens = input_buffers.dcp_local_seq_lens[:num_local_reqs] + + return replace( + input_batch, + req_ids=local_req_ids, + num_reqs=num_local_reqs, + num_reqs_after_padding=num_local_reqs, + idx_mapping=local_to_global_req_idx, + idx_mapping_np=local_to_global_req_idx_np, + expanded_idx_mapping=local_to_global_req_idx, + expanded_local_pos=torch.zeros( + num_local_reqs, dtype=torch.int32, device=self.device + ), + num_scheduled_tokens=local_num_scheduled_tokens, + num_tokens=num_local_tokens, + num_tokens_after_padding=num_local_tokens_padded, + num_draft_tokens=0, + num_draft_tokens_per_req=None, + query_start_loc=local_query_start_loc, + query_start_loc_np=local_query_start_loc_np[: num_local_reqs + 1], + seq_lens=seq_lens, + seq_lens_cpu_upper_bound=torch.from_numpy(seq_lens_cpu_upper_bound_np), + dcp_local_seq_lens=dcp_local_seq_lens, + num_computed_tokens_np=local_start_pos_np, + prefill_len_np=local_prefill_len_np, + num_computed_prefill_tokens_np=local_num_computed_prefill_tokens_np, + is_prefilling_np=local_is_prefilling_np, + max_seq_len_np=global_batch.max_seq_len_np[local_to_global_batch_req_idx_np] + if global_batch.max_seq_len_np is not None + else None, + input_ids=input_buffers.input_ids[:num_local_tokens_padded], + positions=input_buffers.positions[:num_local_tokens_padded], + is_padding=is_padding, + logits_indices=logits_indices, + cu_num_logits=cu_num_logits, + cu_num_logits_np=cu_num_logits_np, + prompt_lens=None, + ) + + def prepare_attn( + self, input_batch: InputBatch + ) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]: + assert self._block_tables is not None + assert self._local_block_tables is not None + assert self._local_block_table_ptrs is not None + block_tables = self._block_tables.gather_block_tables( + input_batch.idx_mapping, + input_batch.num_reqs_after_padding, + out=self._local_block_tables, + out_ptrs=self._local_block_table_ptrs, + ) + slot_mappings = self.prepare_slot_mappings() + return block_tables, slot_mappings + + def prepare_slot_mappings(self) -> torch.Tensor: + assert self._block_tables is not None + assert self._global_batch_slot_mappings is not None + assert self._global_batch is not None + global_batch = self._global_batch + global_batch_slot_mappings = self._block_tables.compute_slot_mappings( + global_batch.idx_mapping, + global_batch.query_start_loc, + global_batch.positions, + global_batch.num_tokens, + out=self._global_batch_slot_mappings, + ) + return self._convert_to_gathered_slot_mappings(global_batch_slot_mappings) + + def get_dummy_slot_mappings(self, num_tokens: int) -> torch.Tensor: + assert self._gathered_kv_slot_mappings is not None + self._gathered_kv_slot_mappings.fill_(PAD_SLOT_ID) + return self._gathered_kv_slot_mappings[:, : num_tokens * self.pcp_world_size] + + def _convert_to_gathered_slot_mappings( + self, + global_batch_slot_mappings: torch.Tensor, + ) -> torch.Tensor: + assert self._padded_gather_idx is not None + assert self._gathered_kv_write_mask is not None + padded_gather_idx = self._padded_gather_idx + num_expanded_tokens = padded_gather_idx.shape[0] + if self._gathered_kv_slot_mappings is None: + self._gathered_kv_slot_mappings = global_batch_slot_mappings.new_empty( + global_batch_slot_mappings.shape[0], num_expanded_tokens + ) + gathered_kv_slot_mappings = self._gathered_kv_slot_mappings[ + :, :num_expanded_tokens + ] + torch.index_select( + global_batch_slot_mappings, + 1, + padded_gather_idx, + out=gathered_kv_slot_mappings, + ) + torch.where( + self._gathered_kv_write_mask.unsqueeze(0), + gathered_kv_slot_mappings, + self._pad_slot_id, + out=gathered_kv_slot_mappings, + ) + return gathered_kv_slot_mappings + + def restore_hidden_states(self, hidden_states: torch.Tensor) -> torch.Tensor: + if self._hidden_restore_idx is None: + return hidden_states + gathered = get_pcp_group().all_gather(hidden_states, dim=0) + return gathered[self._hidden_restore_idx] + + def restore_for_sampling( + self, + hidden_states: torch.Tensor, + ) -> tuple[torch.Tensor, InputBatch]: + assert self._global_batch is not None + return self.restore_hidden_states(hidden_states), self._global_batch + + +def maybe_partition_pcp_batch( + manager: PCPManager | None, + input_batch: InputBatch, +) -> InputBatch: + if manager is None: + return input_batch + return manager.partition_batch(input_batch) + + +def maybe_get_pcp_dummy_slot_mappings( + manager: PCPManager | None, + block_tables: BlockTables, + num_tokens: int, +) -> torch.Tensor: + if manager is None: + return block_tables.get_dummy_slot_mappings(num_tokens) + return manager.get_dummy_slot_mappings(num_tokens) + + +def maybe_restore_pcp_for_sampling( + manager: PCPManager | None, + hidden_states: torch.Tensor | None, + input_batch: InputBatch, +) -> tuple[torch.Tensor, InputBatch]: + assert hidden_states is not None + if manager is None: + return hidden_states, input_batch + return manager.restore_for_sampling(hidden_states) + + +def maybe_build_pcp_manager( + vllm_config: VllmConfig, + device: torch.device, + supports_mm_inputs: bool, + req_states: RequestState, + block_tables: BlockTables, +) -> PCPManager | None: + parallel_config = vllm_config.parallel_config + pcp_size = parallel_config.prefill_context_parallel_size + if pcp_size <= 1: + return None + + PCPManager.validate_config(vllm_config, supports_mm_inputs) + + pcp_rank = get_pcp_group().rank_in_group + dcp_size = parallel_config.decode_context_parallel_size + dcp_rank = get_dcp_group().rank_in_group if dcp_size > 1 else 0 + + return PCPManager( + pcp_world_size=pcp_size, + pcp_rank=pcp_rank, + device=device, + req_states=req_states, + max_num_reqs=vllm_config.scheduler_config.max_num_seqs, + max_num_tokens=vllm_config.scheduler_config.max_num_batched_tokens, + block_tables=block_tables, + dcp_world_size=dcp_size, + dcp_rank=dcp_rank, + cp_interleave=parallel_config.cp_kv_cache_interleave_size, + ) From ac5f38a0f7af29cb5bc15a1e73623922d9832500 Mon Sep 17 00:00:00 2001 From: yzong-rh Date: Sun, 19 Jul 2026 08:18:00 -0400 Subject: [PATCH 29/51] [Refactor] Extract StructuredOutputsParams creation logic from Request.to_sampling_params (#49003) Signed-off-by: Yifan Zong --- .../openai/responses/test_sampling_params.py | 15 ++++++ .../openai/chat_completion/protocol.py | 46 ++++-------------- .../entrypoints/openai/completion/protocol.py | 46 ++++-------------- vllm/entrypoints/openai/engine/protocol.py | 38 ++++++++++++++- vllm/entrypoints/openai/responses/protocol.py | 48 ++++++++++--------- 5 files changed, 95 insertions(+), 98 deletions(-) diff --git a/tests/entrypoints/openai/responses/test_sampling_params.py b/tests/entrypoints/openai/responses/test_sampling_params.py index 87910271dd7..5a68e3a9c0d 100644 --- a/tests/entrypoints/openai/responses/test_sampling_params.py +++ b/tests/entrypoints/openai/responses/test_sampling_params.py @@ -132,6 +132,21 @@ class TestResponsesRequestSamplingParams: assert sampling_params.structured_outputs is not None assert sampling_params.structured_outputs.grammar == "root ::= 'hello'" + def test_text_format_json_object_enables_structured_outputs(self): + """text.format json_object enables structured outputs for sampling.""" + request = ResponsesRequest( + model="test-model", + input="test input", + text=ResponseTextConfig.model_validate({"format": {"type": "json_object"}}), + ) + + sampling_params = request.to_sampling_params(default_max_tokens=1000) + + assert sampling_params.structured_outputs is not None + assert sampling_params.structured_outputs.json_object is True + assert sampling_params.structured_outputs.json is None + assert request.structured_outputs is None + def test_structured_outputs_and_json_schema_conflict(self): """Test that specifying both structured_outputs and json_schema raises.""" structured_outputs = StructuredOutputsParams(grammar="root ::= 'hello'") diff --git a/vllm/entrypoints/openai/chat_completion/protocol.py b/vllm/entrypoints/openai/chat_completion/protocol.py index 76a03dd7202..8c1694bbf73 100644 --- a/vllm/entrypoints/openai/chat_completion/protocol.py +++ b/vllm/entrypoints/openai/chat_completion/protocol.py @@ -3,7 +3,6 @@ # Adapted from # https://github.com/lm-sys/FastChat/blob/168ccc29d3f7edc50823016105c024fe2282732a/fastchat/protocol/openai_api_protocol.py -import json import time from typing import Annotated, Any, ClassVar, Literal @@ -14,7 +13,6 @@ from openai.types.chat.chat_completion_message import Annotation as OpenAIAnnota from pydantic import Field, PrivateAttr, model_serializer, model_validator from vllm.config import ModelConfig -from vllm.config.utils import replace from vllm.entrypoints.chat_utils import ( ChatCompletionMessageParam, ChatTemplateContentFormatOption, @@ -24,13 +22,12 @@ from vllm.entrypoints.openai.engine.protocol import ( DeltaMessage, FunctionCall, FunctionDefinition, - LegacyStructuralTagResponseFormat, OpenAIBaseModel, PerRequestTimingMetrics, StreamOptions, - StructuralTagResponseFormat, ToolCall, UsageInfo, + structured_outputs_from_response_format, validate_structural_tag_response_format, validate_structured_outputs_structural_tag, ) @@ -607,6 +604,13 @@ class ChatCompletionRequest(OpenAIBaseModel): include_stop_str_in_output=self.include_stop_str_in_output, ) + def extract_structured_outputs(self) -> StructuredOutputsParams | None: + """Normalize request constraints into ``StructuredOutputsParams``.""" + return structured_outputs_from_response_format( + self.structured_outputs, + self.response_format, + ) + def to_sampling_params( self, max_tokens: int, @@ -651,38 +655,6 @@ class ChatCompletionRequest(OpenAIBaseModel): if prompt_logprobs is None and self.echo: prompt_logprobs = self.top_logprobs - response_format = self.response_format - if response_format is not None: - structured_outputs_kwargs = dict[str, Any]() - - # Set structured output params for response format - if response_format.type == "json_object": - structured_outputs_kwargs["json_object"] = True - elif response_format.type == "json_schema": - json_schema = response_format.json_schema - assert json_schema is not None - structured_outputs_kwargs["json"] = json_schema.json_schema - elif response_format.type == "structural_tag": - structural_tag = response_format - assert structural_tag is not None and isinstance( - structural_tag, - ( - LegacyStructuralTagResponseFormat, - StructuralTagResponseFormat, - ), - ) - s_tag_obj = structural_tag.model_dump(by_alias=True) - structured_outputs_kwargs["structural_tag"] = json.dumps(s_tag_obj) - - # If structured outputs wasn't already enabled, - # we must enable it for these features to work - if len(structured_outputs_kwargs) > 0: - self.structured_outputs = ( - StructuredOutputsParams(**structured_outputs_kwargs) - if self.structured_outputs is None - else replace(self.structured_outputs, **structured_outputs_kwargs) - ) - extra_args: dict[str, Any] = self.vllm_xargs if self.vllm_xargs else {} if self.kv_transfer_params: # Pass in kv_transfer_params via extra_args @@ -718,7 +690,7 @@ class ChatCompletionRequest(OpenAIBaseModel): output_kind=RequestOutputKind.DELTA if self.stream else RequestOutputKind.FINAL_ONLY, - structured_outputs=self.structured_outputs, + structured_outputs=self.extract_structured_outputs(), logit_bias=self.logit_bias, bad_words=self.bad_words, thinking_token_budget=self.thinking_token_budget, diff --git a/vllm/entrypoints/openai/completion/protocol.py b/vllm/entrypoints/openai/completion/protocol.py index 1784d0f5364..73677b16af9 100644 --- a/vllm/entrypoints/openai/completion/protocol.py +++ b/vllm/entrypoints/openai/completion/protocol.py @@ -3,7 +3,6 @@ # Adapted from # https://github.com/lm-sys/FastChat/blob/168ccc29d3f7edc50823016105c024fe2282732a/fastchat/protocol/openai_api_protocol.py -import json import time from typing import Annotated, Any, Literal @@ -11,15 +10,13 @@ from pydantic import Field, model_validator import vllm.envs as envs from vllm.config import ModelConfig -from vllm.config.utils import replace from vllm.entrypoints.openai.engine.protocol import ( AnyResponseFormat, - LegacyStructuralTagResponseFormat, OpenAIBaseModel, PerRequestTimingMetrics, StreamOptions, - StructuralTagResponseFormat, UsageInfo, + structured_outputs_from_response_format, validate_structural_tag_response_format, validate_structured_outputs_structural_tag, ) @@ -281,6 +278,13 @@ class CompletionRequest(OpenAIBaseModel): include_stop_str_in_output=self.include_stop_str_in_output, ) + def extract_structured_outputs(self) -> StructuredOutputsParams | None: + """Normalize request constraints into ``StructuredOutputsParams``.""" + return structured_outputs_from_response_format( + self.structured_outputs, + self.response_format, + ) + def to_sampling_params( self, max_tokens: int, @@ -330,38 +334,6 @@ class CompletionRequest(OpenAIBaseModel): echo_without_generation = self.echo and self.max_tokens == 0 - response_format = self.response_format - if response_format is not None: - structured_outputs_kwargs = dict[str, Any]() - - # Set structured output params for response format - if response_format.type == "json_object": - structured_outputs_kwargs["json_object"] = True - elif response_format.type == "json_schema": - json_schema = response_format.json_schema - assert json_schema is not None - structured_outputs_kwargs["json"] = json_schema.json_schema - elif response_format.type == "structural_tag": - structural_tag = response_format - assert isinstance( - structural_tag, - ( - LegacyStructuralTagResponseFormat, - StructuralTagResponseFormat, - ), - ) - s_tag_obj = structural_tag.model_dump(by_alias=True) - structured_outputs_kwargs["structural_tag"] = json.dumps(s_tag_obj) - - # If structured outputs wasn't already enabled, - # we must enable it for these features to work - if len(structured_outputs_kwargs) > 0: - self.structured_outputs = ( - StructuredOutputsParams(**structured_outputs_kwargs) - if self.structured_outputs is None - else replace(self.structured_outputs, **structured_outputs_kwargs) - ) - extra_args: dict[str, Any] = self.vllm_xargs if self.vllm_xargs else {} if self.kv_transfer_params: # Pass in kv_transfer_params via extra_args @@ -393,7 +365,7 @@ class CompletionRequest(OpenAIBaseModel): output_kind=RequestOutputKind.DELTA if self.stream else RequestOutputKind.FINAL_ONLY, - structured_outputs=self.structured_outputs, + structured_outputs=self.extract_structured_outputs(), logit_bias=self.logit_bias, allowed_token_ids=self.allowed_token_ids, bad_words=self.bad_words, diff --git a/vllm/entrypoints/openai/engine/protocol.py b/vllm/entrypoints/openai/engine/protocol.py index 7d901d19333..05536f14217 100644 --- a/vllm/entrypoints/openai/engine/protocol.py +++ b/vllm/entrypoints/openai/engine/protocol.py @@ -3,6 +3,7 @@ # Adapted from # https://github.com/lm-sys/FastChat/blob/168ccc29d3f7edc50823016105c024fe2282732a/fastchat/protocol/openai_api_protocol.py +import json import time from http import HTTPStatus from typing import Any, ClassVar, Literal, TypeAlias @@ -16,9 +17,11 @@ from pydantic import ( model_validator, ) +from vllm.config.utils import replace from vllm.entrypoints.chat_utils import make_tool_call_id from vllm.exceptions import VLLMValidationError from vllm.logger import init_logger +from vllm.sampling_params import StructuredOutputsParams from vllm.utils import random_uuid from vllm.utils.import_utils import resolve_obj_by_qualname @@ -173,6 +176,39 @@ AnyResponseFormat: TypeAlias = ( ) +def structured_outputs_from_response_format( + structured_outputs: StructuredOutputsParams | None, + response_format: AnyResponseFormat | None, +) -> StructuredOutputsParams | None: + """Apply ``response_format`` overrides to ``structured_outputs``.""" + if response_format is None or response_format.type == "text": + return structured_outputs + + overrides: dict[str, Any] + if response_format.type == "json_object": + overrides = {"json_object": True} + elif response_format.type == "json_schema": + json_schema = response_format.json_schema + assert json_schema is not None + overrides = {"json": json_schema.json_schema} + else: + assert isinstance( + response_format, + ( + LegacyStructuralTagResponseFormat, + StructuralTagResponseFormat, + ), + ) + overrides = { + "structural_tag": json.dumps(response_format.model_dump(by_alias=True)) + } + + if structured_outputs is None: + return StructuredOutputsParams(**overrides) + + return replace(structured_outputs, **overrides) + + def validate_structural_tag_response_format( response_format: AnyStructuralTagResponseFormat | dict[str, Any], ) -> None: @@ -181,8 +217,6 @@ def validate_structural_tag_response_format( Engine-side validation reports malformed structural tags as generation failures. OpenAI request parsing should classify them as bad requests. """ - import json - from pydantic import TypeAdapter, ValidationError if isinstance(response_format, dict): diff --git a/vllm/entrypoints/openai/responses/protocol.py b/vllm/entrypoints/openai/responses/protocol.py index d4708a5fb3e..3f6857dcc32 100644 --- a/vllm/entrypoints/openai/responses/protocol.py +++ b/vllm/entrypoints/openai/responses/protocol.py @@ -354,6 +354,31 @@ class ResponsesRequest(OpenAIBaseModel): "top_k": 0, } + def extract_structured_outputs(self) -> StructuredOutputsParams | None: + """Normalize request constraints into ``StructuredOutputsParams``.""" + if self.text is None or self.text.format is None: + return self.structured_outputs + + if self.structured_outputs is not None: + raise VLLMValidationError( + "Cannot specify both structured_outputs and text.format", + parameter="structured_outputs", + ) + + response_format = self.text.format + if response_format.type == "json_object": + return StructuredOutputsParams(json_object=True) + if ( + response_format.type == "json_schema" + and response_format.schema_ is not None + ): + return StructuredOutputsParams( + json=response_format.schema_ # type: ignore[call-arg] + # --follow-imports skip hides the class definition but also hides + # multiple third party conflicts, so best of both evils + ) + return None + def to_sampling_params( self, default_max_tokens: int, @@ -387,27 +412,6 @@ class ResponsesRequest(OpenAIBaseModel): if (frequency_penalty := self.frequency_penalty) is None: frequency_penalty = default_sampling_params.get("frequency_penalty", 0.0) - # Structured output - structured_outputs = self.structured_outputs - - # Also check text.format for OpenAI-style json_schema - if self.text is not None and self.text.format is not None: - if structured_outputs is not None: - raise VLLMValidationError( - "Cannot specify both structured_outputs and text.format", - parameter="structured_outputs", - ) - response_format = self.text.format - if ( - response_format.type == "json_schema" - and response_format.schema_ is not None - ): - structured_outputs = StructuredOutputsParams( - json=response_format.schema_ # type: ignore[call-arg] - # --follow-imports skip hides the class definition but also hides - # multiple third party conflicts, so best of both evils - ) - stop = self.stop if self.stop else [] if isinstance(stop, str): stop = [stop] @@ -433,7 +437,7 @@ class ResponsesRequest(OpenAIBaseModel): output_kind=( RequestOutputKind.DELTA if self.stream else RequestOutputKind.FINAL_ONLY ), - structured_outputs=structured_outputs, + structured_outputs=self.extract_structured_outputs(), logit_bias=self.logit_bias, extra_args=extra_args, skip_clone=True, # Created fresh per request, safe to skip clone From e6d1310b2ac52e56e266450693af9de12a871a51 Mon Sep 17 00:00:00 2001 From: Taneem Ibrahim Date: Sun, 19 Jul 2026 07:18:03 -0500 Subject: [PATCH 30/51] [Bugfix] Reject removed pooling parameters (#48984) Signed-off-by: Taneem Ibrahim --- tests/test_pooling_params.py | 38 +++++++++++++++++++ vllm/config/pooler.py | 18 ++++++++- vllm/entrypoints/pooling/base/protocol.py | 13 +++++++ vllm/entrypoints/pooling/classify/protocol.py | 12 +++--- vllm/entrypoints/pooling/embed/protocol.py | 10 +++-- vllm/entrypoints/pooling/pooling/protocol.py | 12 +++--- vllm/entrypoints/pooling/pooling/serving.py | 4 +- vllm/pooling_params.py | 3 +- vllm/tasks.py | 15 ++++++++ 9 files changed, 108 insertions(+), 17 deletions(-) diff --git a/tests/test_pooling_params.py b/tests/test_pooling_params.py index 6bd97db03dc..17d04078b4e 100644 --- a/tests/test_pooling_params.py +++ b/tests/test_pooling_params.py @@ -1,12 +1,18 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from dataclasses import dataclass +from typing import Any import pytest +from pydantic import TypeAdapter, ValidationError from tests.models.utils import EmbedModelInfo from vllm import PoolingParams from vllm.config import ModelConfig, PoolerConfig +from vllm.entrypoints.pooling.classify.protocol import ClassificationRequest +from vllm.entrypoints.pooling.embed.protocol import EmbeddingRequest +from vllm.entrypoints.pooling.pooling.protocol import PoolingRequest +from vllm.exceptions import VLLMValidationError EMBEDDING_MODELS = [ EmbedModelInfo("intfloat/multilingual-e5-small", is_matryoshka=False), @@ -27,6 +33,38 @@ class MockModelConfig: pooler_config: PoolerConfig +@pytest.mark.parametrize( + ("parameter", "value", "message"), + [ + ( + "normalize", + False, + "Parameter `normalize` was removed; use `use_activation` instead.", + ), + ("task", "score", "`score` task was removed; use `classify` instead."), + ( + "task", + "encode", + "`encode` task was removed; use `token_embed` or `token_classify` instead.", + ), + ], +) +def test_removed_pooling_parameters(parameter: str, value: Any, message: str): + data = {"input": "hello", parameter: value} + for request_type in (EmbeddingRequest, ClassificationRequest, PoolingRequest): + with pytest.raises(ValidationError, match=message) as exc_info: + TypeAdapter(request_type).validate_python(data) + assert len(exc_info.value.errors()) == 1 + + with pytest.raises(ValidationError, match=message) as exc_info: + TypeAdapter(PoolerConfig).validate_python({parameter: value}) + assert len(exc_info.value.errors()) == 1 + + if parameter == "task": + with pytest.raises(VLLMValidationError, match=message): + PoolingParams(task=value) + + def test_embed(): task = "embed" model_config = MockModelConfig(pooler_config=PoolerConfig(seq_pooling_type="CLS")) diff --git a/vllm/config/pooler.py b/vllm/config/pooler.py index f7985a52e4a..9ca0ba457eb 100644 --- a/vllm/config/pooler.py +++ b/vllm/config/pooler.py @@ -3,9 +3,12 @@ from typing import Any, Literal, get_args +from pydantic import model_validator +from pydantic_core import ArgsKwargs + from vllm.config.utils import config from vllm.logger import init_logger -from vllm.tasks import PoolingTask +from vllm.tasks import PoolingTask, check_removed_pooling_task from vllm.utils.hashing import safe_hash logger = init_logger(__name__) @@ -112,6 +115,19 @@ class PoolerConfig: `math-shepherd-mistral-7b-prm` model. """ + @model_validator(mode="before") + @classmethod + def reject_removed_parameters(cls, data): + values = data.kwargs if isinstance(data, ArgsKwargs) else data + if not isinstance(values, dict): + return data + if "normalize" in values: + raise ValueError( + "Parameter `normalize` was removed; use `use_activation` instead." + ) + check_removed_pooling_task(values.get("task")) + return data + def __post_init__(self) -> None: if self.logit_sigma is not None and self.logit_sigma == 0: raise ValueError("logit_sigma cannot be 0 (division by zero)") diff --git a/vllm/entrypoints/pooling/base/protocol.py b/vllm/entrypoints/pooling/base/protocol.py index 04ebc18817b..c0fff022d26 100644 --- a/vllm/entrypoints/pooling/base/protocol.py +++ b/vllm/entrypoints/pooling/base/protocol.py @@ -14,10 +14,23 @@ from vllm.entrypoints.chat_utils import ( from vllm.entrypoints.openai.engine.protocol import OpenAIBaseModel from vllm.exceptions import VLLMValidationError from vllm.renderers import ChatParams, TokenizeParams, merge_kwargs +from vllm.tasks import check_removed_pooling_task from vllm.utils import random_uuid from vllm.utils.serial_utils import EmbedDType, EncodingFormat, Endianness +def reject_removed_pooling_parameters(data): + if not isinstance(data, dict): + return data + if "normalize" in data: + raise VLLMValidationError( + "Parameter `normalize` was removed; use `use_activation` instead.", + parameter="normalize", + ) + check_removed_pooling_task(data.get("task")) + return data + + class PoolingBasicRequestMixin(OpenAIBaseModel): # --8<-- [start:pooling-common-params] model: str | None = None diff --git a/vllm/entrypoints/pooling/classify/protocol.py b/vllm/entrypoints/pooling/classify/protocol.py index 39cacdd835e..1b09319c29d 100644 --- a/vllm/entrypoints/pooling/classify/protocol.py +++ b/vllm/entrypoints/pooling/classify/protocol.py @@ -2,9 +2,9 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import time -from typing import TypeAlias +from typing import Annotated, TypeAlias -from pydantic import Field +from pydantic import BeforeValidator, Field from vllm import PoolingParams from vllm.entrypoints.openai.engine.protocol import OpenAIBaseModel, UsageInfo @@ -17,6 +17,7 @@ from ..base.protocol import ( CompletionRequestMixin, FixedMaxLenTokenizeParamsMixin, PoolingBasicRequestMixin, + reject_removed_pooling_parameters, ) logger = init_logger(__name__) @@ -48,9 +49,10 @@ class ClassificationChatRequest( ) -ClassificationRequest: TypeAlias = ( - ClassificationCompletionRequest | ClassificationChatRequest -) +ClassificationRequest: TypeAlias = Annotated[ + ClassificationCompletionRequest | ClassificationChatRequest, + BeforeValidator(reject_removed_pooling_parameters), +] class ClassificationData(OpenAIBaseModel): diff --git a/vllm/entrypoints/pooling/embed/protocol.py b/vllm/entrypoints/pooling/embed/protocol.py index 2dcc848c8c7..b2912e544f2 100644 --- a/vllm/entrypoints/pooling/embed/protocol.py +++ b/vllm/entrypoints/pooling/embed/protocol.py @@ -13,7 +13,7 @@ from collections.abc import Sequence from typing import Annotated, Any, Literal, TypeAlias import pybase64 as base64 -from pydantic import BaseModel, Field, model_validator +from pydantic import BaseModel, BeforeValidator, Field, model_validator from vllm import PoolingParams from vllm.entrypoints.chat_utils import ChatCompletionMessageParam @@ -27,6 +27,7 @@ from ..base.protocol import ( EmbeddingTokenizeParamsMixin, EmbedRequestMixin, PoolingBasicRequestMixin, + reject_removed_pooling_parameters, ) @@ -154,13 +155,14 @@ class EmbeddingBatchChatInputRequest(EmbeddingBatchChatRequest): return normalized -EmbeddingRequest: TypeAlias = ( +EmbeddingRequest: TypeAlias = Annotated[ EmbeddingCompletionRequest | EmbeddingChatRequest | EmbeddingBatchChatRequest | EmbeddingChatInputRequest - | EmbeddingBatchChatInputRequest -) + | EmbeddingBatchChatInputRequest, + BeforeValidator(reject_removed_pooling_parameters), +] # --------------------------------------------------------------------------- diff --git a/vllm/entrypoints/pooling/pooling/protocol.py b/vllm/entrypoints/pooling/pooling/protocol.py index b2a43b1935e..b7edda18c74 100644 --- a/vllm/entrypoints/pooling/pooling/protocol.py +++ b/vllm/entrypoints/pooling/pooling/protocol.py @@ -1,9 +1,9 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import time -from typing import Generic, TypeAlias, TypeVar +from typing import Annotated, Generic, TypeAlias, TypeVar -from pydantic import Field +from pydantic import BeforeValidator, Field from vllm import PoolingParams from vllm.config import ModelConfig @@ -20,6 +20,7 @@ from ..base.protocol import ( EncodingRequestMixin, FixedMaxLenTokenizeParamsMixin, PoolingBasicRequestMixin, + reject_removed_pooling_parameters, ) @@ -92,9 +93,10 @@ class IOProcessorResponse(OpenAIBaseModel, Generic[T]): """ -PoolingRequest: TypeAlias = ( - PoolingCompletionRequest | PoolingChatRequest | IOProcessorRequest -) +PoolingRequest: TypeAlias = Annotated[ + PoolingCompletionRequest | PoolingChatRequest | IOProcessorRequest, + BeforeValidator(reject_removed_pooling_parameters), +] class PoolingResponseData(OpenAIBaseModel): diff --git a/vllm/entrypoints/pooling/pooling/serving.py b/vllm/entrypoints/pooling/pooling/serving.py index a73ebf77233..e049d1ba18f 100644 --- a/vllm/entrypoints/pooling/pooling/serving.py +++ b/vllm/entrypoints/pooling/pooling/serving.py @@ -1,5 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import cast + from fastapi.responses import JSONResponse, Response, StreamingResponse from typing_extensions import assert_never @@ -52,7 +54,7 @@ class ServingPooling(PoolingBaseServing): self.json_response_cls = get_json_response_cls() def get_io_processor(self, request: AnyPoolingRequest) -> PoolingIOProcessor: - assert isinstance(request, PoolingRequest) + request = cast(PoolingRequest, request) pooling_task = self._verify_pooling_task(request) return self.io_processors[pooling_task] diff --git a/vllm/pooling_params.py b/vllm/pooling_params.py index 0da436a871d..6cb130fdbbb 100644 --- a/vllm/pooling_params.py +++ b/vllm/pooling_params.py @@ -9,7 +9,7 @@ import msgspec from vllm.config import ModelConfig, PoolerConfig from vllm.logger import init_logger from vllm.sampling_params import RequestOutputKind -from vllm.tasks import PoolingTask +from vllm.tasks import PoolingTask, check_removed_pooling_task logger = init_logger(__name__) @@ -229,6 +229,7 @@ class PoolingParams( ) def __post_init__(self) -> None: + check_removed_pooling_task(self.task) if self.output_kind != RequestOutputKind.FINAL_ONLY: raise ValueError( "For pooling output_kind has to be FINAL_ONLY, " diff --git a/vllm/tasks.py b/vllm/tasks.py index 017bc31197b..366ef0c672f 100644 --- a/vllm/tasks.py +++ b/vllm/tasks.py @@ -2,6 +2,8 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from typing import Literal, get_args +from vllm.exceptions import VLLMValidationError + GenerationTask = Literal["generate", "transcription", "realtime"] GENERATION_TASKS: tuple[GenerationTask, ...] = get_args(GenerationTask) @@ -15,6 +17,19 @@ PoolingTask = Literal[ ] POOLING_TASKS: tuple[PoolingTask, ...] = get_args(PoolingTask) +_REMOVED_POOLING_TASK_MESSAGES = { + "score": "`score` task was removed; use `classify` instead.", + "encode": ( + "`encode` task was removed; use `token_embed` or `token_classify` instead." + ), +} + + +def check_removed_pooling_task(task: object) -> None: + if isinstance(task, str) and (message := _REMOVED_POOLING_TASK_MESSAGES.get(task)): + raise VLLMValidationError(message, parameter="task") + + ScoreType = Literal["bi-encoder", "cross-encoder", "late-interaction"] SCORE_TYPE_MAP: dict[PoolingTask, ScoreType] = { "embed": "bi-encoder", From ef0aa7ca2feb75051b30ea3cef4e9950252b1441 Mon Sep 17 00:00:00 2001 From: TJian Date: Sun, 19 Jul 2026 11:38:04 -0700 Subject: [PATCH 31/51] [ROCm] [Release] [Per-commit] Reenable per commit rocm wheel (#49044) Signed-off-by: tjtanaa --- .buildkite/release-pipeline.yaml | 683 +++++++++++++++---------------- 1 file changed, 340 insertions(+), 343 deletions(-) diff --git a/.buildkite/release-pipeline.yaml b/.buildkite/release-pipeline.yaml index 1be6f60fe25..84bf492fccd 100644 --- a/.buildkite/release-pipeline.yaml +++ b/.buildkite/release-pipeline.yaml @@ -590,373 +590,370 @@ steps: # # ============================================================================= - - block: "Unblock ROCm wheel/image prerequisites" + - group: "Build ROCm Wheel / Image " + key: "build-rocm-wheel-image" depends_on: ~ - key: block-build-rocm - if: build.env("NIGHTLY") != "1" + steps: + # ROCm Job 1: Build ROCm Base Wheels (with S3 caching) + - label: ":rocm: Build ROCm Base Image & Wheels" + id: build-rocm-base-wheels + depends_on: ~ + agents: + queue: cpu_queue_release + commands: + - | + set -euo pipefail - # ROCm Job 1: Build ROCm Base Wheels (with S3 caching) - - label: ":rocm: Build ROCm Base Image & Wheels" - id: build-rocm-base-wheels - depends_on: - - step: block-build-rocm - allow_failure: true - agents: - queue: cpu_queue_release - commands: - - | - set -euo pipefail + # Generate cache key + CACHE_KEY=$$(.buildkite/scripts/cache-rocm-base-wheels.sh key) + ECR_CACHE_TAG="public.ecr.aws/q9t5s3a7/vllm-release-repo:$${CACHE_KEY}-rocm-base" - # Generate cache key - CACHE_KEY=$$(.buildkite/scripts/cache-rocm-base-wheels.sh key) - ECR_CACHE_TAG="public.ecr.aws/q9t5s3a7/vllm-release-repo:$${CACHE_KEY}-rocm-base" + echo "========================================" + echo "ROCm Base Build Configuration" + echo "========================================" + echo " CACHE_KEY: $${CACHE_KEY}" + echo " ECR_CACHE_TAG: $${ECR_CACHE_TAG}" + echo "========================================" + + # Login to ECR + aws ecr-public get-login-password --region us-east-1 | \ + docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7 + + IMAGE_EXISTS=false + WHEELS_EXIST=false + + # Check ECR for Docker image - echo "========================================" - echo "ROCm Base Build Configuration" - echo "========================================" - echo " CACHE_KEY: $${CACHE_KEY}" - echo " ECR_CACHE_TAG: $${ECR_CACHE_TAG}" - echo "========================================" - - # Login to ECR - aws ecr-public get-login-password --region us-east-1 | \ - docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7 - - IMAGE_EXISTS=false - WHEELS_EXIST=false - - # Check ECR for Docker image + if docker manifest inspect "$${ECR_CACHE_TAG}" > /dev/null 2>&1; then + IMAGE_EXISTS=true + echo "ECR image cache HIT" + fi + + # Check S3 for wheels + WHEEL_CACHE_STATUS=$(.buildkite/scripts/cache-rocm-base-wheels.sh check) + if [ "$${WHEEL_CACHE_STATUS}" = "hit" ]; then + WHEELS_EXIST=true + echo "S3 wheels cache HIT" + fi - if docker manifest inspect "$${ECR_CACHE_TAG}" > /dev/null 2>&1; then - IMAGE_EXISTS=true - echo "ECR image cache HIT" - fi - - # Check S3 for wheels - WHEEL_CACHE_STATUS=$(.buildkite/scripts/cache-rocm-base-wheels.sh check) - if [ "$${WHEEL_CACHE_STATUS}" = "hit" ]; then - WHEELS_EXIST=true - echo "S3 wheels cache HIT" - fi + + # Scenario 1: Both cached (best case) + if [ "$${IMAGE_EXISTS}" = "true" ] && [ "$${WHEELS_EXIST}" = "true" ]; then + echo "" + echo "FULL CACHE HIT - Reusing both image and wheels" + echo "" - - # Scenario 1: Both cached (best case) - if [ "$${IMAGE_EXISTS}" = "true" ] && [ "$${WHEELS_EXIST}" = "true" ]; then - echo "" - echo "FULL CACHE HIT - Reusing both image and wheels" - echo "" + # Download wheels + .buildkite/scripts/cache-rocm-base-wheels.sh download + + # Save ECR tag for downstream jobs + buildkite-agent meta-data set "rocm-base-image-tag" "$${ECR_CACHE_TAG}" + + # Scenario 2: Full rebuild needed + else + echo "" + echo " CACHE MISS - Building from scratch..." + echo "" + + # Build full base image and push to ECR + DOCKER_BUILDKIT=1 docker buildx build \ + --file docker/Dockerfile.rocm_base \ + --tag "$${ECR_CACHE_TAG}" \ + --build-arg USE_SCCACHE=1 \ + --build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \ + --build-arg SCCACHE_REGION_NAME=us-west-2 \ + --build-arg SCCACHE_S3_NO_CREDENTIALS=0 \ + --push \ + . + + # Build wheel extraction stage + DOCKER_BUILDKIT=1 docker buildx build \ + --file docker/Dockerfile.rocm_base \ + --tag rocm-base-debs:$${BUILDKITE_BUILD_NUMBER} \ + --target debs_wheel_release \ + --build-arg USE_SCCACHE=1 \ + --build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \ + --build-arg SCCACHE_REGION_NAME=us-west-2 \ + --build-arg SCCACHE_S3_NO_CREDENTIALS=0 \ + --load \ + . + + # Extract and upload wheels + mkdir -p artifacts/rocm-base-wheels + cid=$(docker create rocm-base-debs:$${BUILDKITE_BUILD_NUMBER}) + docker cp $${cid}:/app/debs/. artifacts/rocm-base-wheels/ + docker rm $${cid} + + .buildkite/scripts/cache-rocm-base-wheels.sh upload - # Download wheels - .buildkite/scripts/cache-rocm-base-wheels.sh download - - # Save ECR tag for downstream jobs - buildkite-agent meta-data set "rocm-base-image-tag" "$${ECR_CACHE_TAG}" - - # Scenario 2: Full rebuild needed - else - echo "" - echo " CACHE MISS - Building from scratch..." - echo "" - - # Build full base image and push to ECR - DOCKER_BUILDKIT=1 docker buildx build \ - --file docker/Dockerfile.rocm_base \ - --tag "$${ECR_CACHE_TAG}" \ - --build-arg USE_SCCACHE=1 \ - --build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \ - --build-arg SCCACHE_REGION_NAME=us-west-2 \ - --build-arg SCCACHE_S3_NO_CREDENTIALS=0 \ - --push \ - . - - # Build wheel extraction stage - DOCKER_BUILDKIT=1 docker buildx build \ - --file docker/Dockerfile.rocm_base \ - --tag rocm-base-debs:$${BUILDKITE_BUILD_NUMBER} \ - --target debs_wheel_release \ - --build-arg USE_SCCACHE=1 \ - --build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \ - --build-arg SCCACHE_REGION_NAME=us-west-2 \ - --build-arg SCCACHE_S3_NO_CREDENTIALS=0 \ - --load \ - . - - # Extract and upload wheels - mkdir -p artifacts/rocm-base-wheels - cid=$(docker create rocm-base-debs:$${BUILDKITE_BUILD_NUMBER}) - docker cp $${cid}:/app/debs/. artifacts/rocm-base-wheels/ - docker rm $${cid} - - .buildkite/scripts/cache-rocm-base-wheels.sh upload + # Cache base docker image to ECR + docker push "$${ECR_CACHE_TAG}" + + buildkite-agent meta-data set "rocm-base-image-tag" "$${ECR_CACHE_TAG}" + + echo "" + echo " Build complete - Image and wheels cached" + fi - # Cache base docker image to ECR - docker push "$${ECR_CACHE_TAG}" - - buildkite-agent meta-data set "rocm-base-image-tag" "$${ECR_CACHE_TAG}" - - echo "" - echo " Build complete - Image and wheels cached" - fi + artifact_paths: + - "artifacts/rocm-base-wheels/*.whl" + env: + DOCKER_BUILDKIT: "1" + S3_BUCKET: "vllm-wheels" - artifact_paths: - - "artifacts/rocm-base-wheels/*.whl" - env: - DOCKER_BUILDKIT: "1" - S3_BUCKET: "vllm-wheels" + # ROCm Job 2: Build vLLM ROCm Wheel + - label: ":python: Build vLLM ROCm Wheel - x86_64" + id: build-rocm-vllm-wheel + depends_on: + - step: build-rocm-base-wheels + allow_failure: false + agents: + queue: cpu_queue_release + timeout_in_minutes: 180 + commands: + # Download artifacts and prepare Docker image + - | + set -euo pipefail - # ROCm Job 2: Build vLLM ROCm Wheel - - label: ":python: Build vLLM ROCm Wheel - x86_64" - id: build-rocm-vllm-wheel - depends_on: - - step: build-rocm-base-wheels - allow_failure: false - agents: - queue: cpu_queue_release - timeout_in_minutes: 180 - commands: - # Download artifacts and prepare Docker image - - | - set -euo pipefail + # Ensure git tags are up-to-date (Buildkite's default fetch doesn't update tags) + # This fixes version detection when tags are moved/force-pushed + echo "Fetching latest tags from origin..." + git fetch --tags --force origin + + # Log tag information for debugging version detection + echo "========================================" + echo "Git Tag Verification" + echo "========================================" + echo "Current HEAD: $(git rev-parse HEAD)" + echo "git describe --tags: $(git describe --tags 2>/dev/null || echo 'No tags found')" + echo "" + echo "Recent tags (pointing to commits near HEAD):" + git tag -l --sort=-creatordate | head -5 + echo "setuptools_scm version detection:" + pip install -q setuptools_scm 2>/dev/null || true + python3 -c "import setuptools_scm; print(' Detected version:', setuptools_scm.get_version())" 2>/dev/null || echo " (setuptools_scm not available in this environment)" + echo "========================================" - # Ensure git tags are up-to-date (Buildkite's default fetch doesn't update tags) - # This fixes version detection when tags are moved/force-pushed - echo "Fetching latest tags from origin..." - git fetch --tags --force origin - - # Log tag information for debugging version detection - echo "========================================" - echo "Git Tag Verification" - echo "========================================" - echo "Current HEAD: $(git rev-parse HEAD)" - echo "git describe --tags: $(git describe --tags 2>/dev/null || echo 'No tags found')" - echo "" - echo "Recent tags (pointing to commits near HEAD):" - git tag -l --sort=-creatordate | head -5 - echo "setuptools_scm version detection:" - pip install -q setuptools_scm 2>/dev/null || true - python3 -c "import setuptools_scm; print(' Detected version:', setuptools_scm.get_version())" 2>/dev/null || echo " (setuptools_scm not available in this environment)" - echo "========================================" + # Download wheel artifacts from current build + echo "Downloading wheel artifacts from current build" + buildkite-agent artifact download "artifacts/rocm-base-wheels/*.whl" . - # Download wheel artifacts from current build - echo "Downloading wheel artifacts from current build" - buildkite-agent artifact download "artifacts/rocm-base-wheels/*.whl" . + # Get ECR image tag from metadata (set by build-rocm-base-wheels) + ECR_IMAGE_TAG="$$(buildkite-agent meta-data get rocm-base-image-tag 2>/dev/null || echo '')" + if [ -z "$${ECR_IMAGE_TAG}" ]; then + echo "ERROR: rocm-base-image-tag metadata not found" + echo "This should have been set by the build-rocm-base-wheels job" + exit 1 + fi + + echo "Pulling base Docker image from ECR: $${ECR_IMAGE_TAG}" + + # Login to ECR + aws ecr-public get-login-password --region us-east-1 | \ + docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7 + + # Pull base Docker image from ECR + docker pull "$${ECR_IMAGE_TAG}" + + echo "Loaded base image: $${ECR_IMAGE_TAG}" + + # Prepare base wheels for Docker build context + mkdir -p docker/context/base-wheels + touch docker/context/base-wheels/.keep + cp artifacts/rocm-base-wheels/*.whl docker/context/base-wheels/ + echo "Base wheels for vLLM build:" + ls -lh docker/context/base-wheels/ - # Get ECR image tag from metadata (set by build-rocm-base-wheels) - ECR_IMAGE_TAG="$$(buildkite-agent meta-data get rocm-base-image-tag 2>/dev/null || echo '')" - if [ -z "$${ECR_IMAGE_TAG}" ]; then - echo "ERROR: rocm-base-image-tag metadata not found" - echo "This should have been set by the build-rocm-base-wheels job" - exit 1 - fi - - echo "Pulling base Docker image from ECR: $${ECR_IMAGE_TAG}" - - # Login to ECR - aws ecr-public get-login-password --region us-east-1 | \ - docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7 - - # Pull base Docker image from ECR - docker pull "$${ECR_IMAGE_TAG}" - - echo "Loaded base image: $${ECR_IMAGE_TAG}" - - # Prepare base wheels for Docker build context - mkdir -p docker/context/base-wheels - touch docker/context/base-wheels/.keep - cp artifacts/rocm-base-wheels/*.whl docker/context/base-wheels/ - echo "Base wheels for vLLM build:" - ls -lh docker/context/base-wheels/ + echo "========================================" + echo "Building vLLM wheel with:" + echo " BUILDKITE_COMMIT: $${BUILDKITE_COMMIT}" + echo " BUILDKITE_BRANCH: $${BUILDKITE_BRANCH}" + echo " BASE_IMAGE: $${ECR_IMAGE_TAG}" + echo "========================================" - echo "========================================" - echo "Building vLLM wheel with:" - echo " BUILDKITE_COMMIT: $${BUILDKITE_COMMIT}" - echo " BUILDKITE_BRANCH: $${BUILDKITE_BRANCH}" - echo " BASE_IMAGE: $${ECR_IMAGE_TAG}" - echo "========================================" + # Build vLLM wheel using local checkout (REMOTE_VLLM=0) + DOCKER_BUILDKIT=1 docker build \ + --file docker/Dockerfile.rocm \ + --target export_vllm_wheel_release \ + --output type=local,dest=rocm-dist \ + --build-arg BASE_IMAGE="$${ECR_IMAGE_TAG}" \ + --build-arg REMOTE_VLLM=0 \ + --build-arg GIT_REPO_CHECK=1 \ + --build-arg USE_SCCACHE=1 \ + --build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \ + --build-arg SCCACHE_REGION_NAME=us-west-2 \ + --build-arg SCCACHE_S3_NO_CREDENTIALS=0 \ + . + echo "Built vLLM wheel:" + ls -lh rocm-dist/*.whl + # Copy wheel to artifacts directory + mkdir -p artifacts/rocm-vllm-wheel + cp rocm-dist/*.whl artifacts/rocm-vllm-wheel/ + echo "Final vLLM wheel:" + ls -lh artifacts/rocm-vllm-wheel/ + artifact_paths: + - "artifacts/rocm-vllm-wheel/*.whl" + env: + DOCKER_BUILDKIT: "1" + S3_BUCKET: "vllm-wheels" - # Build vLLM wheel using local checkout (REMOTE_VLLM=0) - DOCKER_BUILDKIT=1 docker build \ - --file docker/Dockerfile.rocm \ - --target export_vllm_wheel_release \ - --output type=local,dest=rocm-dist \ - --build-arg BASE_IMAGE="$${ECR_IMAGE_TAG}" \ - --build-arg REMOTE_VLLM=0 \ - --build-arg GIT_REPO_CHECK=1 \ - --build-arg USE_SCCACHE=1 \ - --build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \ - --build-arg SCCACHE_REGION_NAME=us-west-2 \ - --build-arg SCCACHE_S3_NO_CREDENTIALS=0 \ - . - echo "Built vLLM wheel:" - ls -lh rocm-dist/*.whl - # Copy wheel to artifacts directory - mkdir -p artifacts/rocm-vllm-wheel - cp rocm-dist/*.whl artifacts/rocm-vllm-wheel/ - echo "Final vLLM wheel:" - ls -lh artifacts/rocm-vllm-wheel/ - artifact_paths: - - "artifacts/rocm-vllm-wheel/*.whl" - env: - DOCKER_BUILDKIT: "1" - S3_BUCKET: "vllm-wheels" + # ROCm Job 3: Upload Wheels to S3 + - label: ":s3: Upload ROCm Wheels to S3" + id: upload-rocm-wheels + depends_on: + - step: build-rocm-vllm-wheel + allow_failure: false + agents: + queue: cpu_queue_release + timeout_in_minutes: 60 + commands: + # Download all wheel artifacts and run upload + - | + set -euo pipefail - # ROCm Job 3: Upload Wheels to S3 - - label: ":s3: Upload ROCm Wheels to S3" - id: upload-rocm-wheels - depends_on: - - step: build-rocm-vllm-wheel - allow_failure: false - agents: - queue: cpu_queue_release - timeout_in_minutes: 60 - commands: - # Download all wheel artifacts and run upload - - | - set -euo pipefail + # Download artifacts from current build + echo "Downloading artifacts from current build" + # buildkite-agent artifact download "artifacts/rocm-base-wheels/*.whl" . + # buildkite-agent artifact download "artifacts/rocm-vllm-wheel/*.whl" . - # Download artifacts from current build - echo "Downloading artifacts from current build" - buildkite-agent artifact download "artifacts/rocm-base-wheels/*.whl" . - buildkite-agent artifact download "artifacts/rocm-vllm-wheel/*.whl" . + # # Run upload script + bash .buildkite/scripts/upload-rocm-wheels.sh + env: + DOCKER_BUILDKIT: "1" + S3_BUCKET: "vllm-wheels" - # Run upload script - bash .buildkite/scripts/upload-rocm-wheels.sh - env: - DOCKER_BUILDKIT: "1" - S3_BUCKET: "vllm-wheels" + # ROCm Job 4: Annotate ROCm Wheel Release + - label: ":memo: Annotate ROCm wheel release" + id: annotate-rocm-release + depends_on: + - upload-rocm-wheels + agents: + queue: cpu_queue_release + commands: + - "bash .buildkite/scripts/annotate-rocm-release.sh" + env: + S3_BUCKET: "vllm-wheels" - # ROCm Job 4: Annotate ROCm Wheel Release - - label: ":memo: Annotate ROCm wheel release" - id: annotate-rocm-release - depends_on: - - upload-rocm-wheels - agents: - queue: cpu_queue_release - commands: - - "bash .buildkite/scripts/annotate-rocm-release.sh" - env: - S3_BUCKET: "vllm-wheels" + # ROCm Job 5: Generate Root Index for ROCm Wheels (for release only) + # This is the job to create https://wheels.vllm.ai/rocm/ index allowing + # users to install with `uv pip install vllm --extra-index-url https://wheels.vllm.ai/rocm/` + - block: "Generate Root Index for ROCm Wheels for Release" + key: block-generate-root-index-rocm-wheels + depends_on: upload-rocm-wheels - # ROCm Job 5: Generate Root Index for ROCm Wheels (for release only) - # This is the job to create https://wheels.vllm.ai/rocm/ index allowing - # users to install with `uv pip install vllm --extra-index-url https://wheels.vllm.ai/rocm/` - - block: "Generate Root Index for ROCm Wheels for Release" - key: block-generate-root-index-rocm-wheels - depends_on: upload-rocm-wheels + - label: ":package: Generate Root Index for ROCm Wheels for Release" + depends_on: block-generate-root-index-rocm-wheels + id: generate-root-index-rocm-wheels + agents: + queue: cpu_queue_release + commands: + - "bash tools/vllm-rocm/generate-rocm-wheels-root-index.sh" + env: + S3_BUCKET: "vllm-wheels" + VARIANT: "rocm723" - - label: ":package: Generate Root Index for ROCm Wheels for Release" - depends_on: block-generate-root-index-rocm-wheels - id: generate-root-index-rocm-wheels - agents: - queue: cpu_queue_release - commands: - - "bash tools/vllm-rocm/generate-rocm-wheels-root-index.sh" - env: - S3_BUCKET: "vllm-wheels" - VARIANT: "rocm723" + # ROCm Job 6: Build ROCm Release Docker Image + - label: ":docker: Build release image - x86_64 - ROCm" + id: build-rocm-release-image + depends_on: + - step: block-build-release-images + allow_failure: true + - step: build-rocm-base-wheels + allow_failure: false + agents: + queue: cpu_queue_release + timeout_in_minutes: 60 + commands: + - | + set -euo pipefail + + # Login to ECR + aws ecr-public get-login-password --region us-east-1 | \ + docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7 + + # Get ECR image tag from metadata (set by build-rocm-base-wheels) + ECR_IMAGE_TAG="$$(buildkite-agent meta-data get rocm-base-image-tag 2>/dev/null || echo '')" + if [ -z "$${ECR_IMAGE_TAG}" ]; then + echo "ERROR: rocm-base-image-tag metadata not found" + echo "This should have been set by the build-rocm-base-wheels job" + exit 1 + fi + + echo "Pulling base Docker image from ECR: $${ECR_IMAGE_TAG}" + + # Pull base Docker image from ECR + docker pull "$${ECR_IMAGE_TAG}" + + echo "Loaded base image: $${ECR_IMAGE_TAG}" + + # Pass the base image ECR tag to downstream steps (nightly publish) + buildkite-agent meta-data set "rocm-base-ecr-tag" "$${ECR_IMAGE_TAG}" + + echo "========================================" + echo "Building vLLM ROCm release image with:" + echo " BASE_IMAGE: $${ECR_IMAGE_TAG}" + echo " BUILDKITE_COMMIT: $${BUILDKITE_COMMIT}" + echo "========================================" + + # Build vLLM ROCm release image using cached base + DOCKER_BUILDKIT=1 docker build \ + --build-arg max_jobs=16 \ + --build-arg BASE_IMAGE="$${ECR_IMAGE_TAG}" \ + --build-arg USE_SCCACHE=1 \ + --build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \ + --build-arg SCCACHE_REGION_NAME=us-west-2 \ + --build-arg SCCACHE_S3_NO_CREDENTIALS=0 \ + --tag public.ecr.aws/q9t5s3a7/vllm-release-repo:$${BUILDKITE_COMMIT}-rocm \ + --target vllm-openai \ + --progress plain \ + -f docker/Dockerfile.rocm . + + # Push to ECR + docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$${BUILDKITE_COMMIT}-rocm - # ROCm Job 6: Build ROCm Release Docker Image - - label: ":docker: Build release image - x86_64 - ROCm" - id: build-rocm-release-image - depends_on: - - step: block-build-release-images - allow_failure: true - - step: build-rocm-base-wheels - allow_failure: false - agents: - queue: cpu_queue_release - timeout_in_minutes: 60 - commands: - - | - set -euo pipefail - - # Login to ECR - aws ecr-public get-login-password --region us-east-1 | \ - docker login --username AWS --password-stdin public.ecr.aws/q9t5s3a7 - - # Get ECR image tag from metadata (set by build-rocm-base-wheels) - ECR_IMAGE_TAG="$$(buildkite-agent meta-data get rocm-base-image-tag 2>/dev/null || echo '')" - if [ -z "$${ECR_IMAGE_TAG}" ]; then - echo "ERROR: rocm-base-image-tag metadata not found" - echo "This should have been set by the build-rocm-base-wheels job" - exit 1 - fi - - echo "Pulling base Docker image from ECR: $${ECR_IMAGE_TAG}" - - # Pull base Docker image from ECR - docker pull "$${ECR_IMAGE_TAG}" - - echo "Loaded base image: $${ECR_IMAGE_TAG}" - - # Pass the base image ECR tag to downstream steps (nightly publish) - buildkite-agent meta-data set "rocm-base-ecr-tag" "$${ECR_IMAGE_TAG}" - - echo "========================================" - echo "Building vLLM ROCm release image with:" - echo " BASE_IMAGE: $${ECR_IMAGE_TAG}" - echo " BUILDKITE_COMMIT: $${BUILDKITE_COMMIT}" - echo "========================================" - - # Build vLLM ROCm release image using cached base - DOCKER_BUILDKIT=1 docker build \ - --build-arg max_jobs=16 \ - --build-arg BASE_IMAGE="$${ECR_IMAGE_TAG}" \ - --build-arg USE_SCCACHE=1 \ - --build-arg SCCACHE_BUCKET_NAME=vllm-build-sccache \ - --build-arg SCCACHE_REGION_NAME=us-west-2 \ - --build-arg SCCACHE_S3_NO_CREDENTIALS=0 \ - --tag public.ecr.aws/q9t5s3a7/vllm-release-repo:$${BUILDKITE_COMMIT}-rocm \ - --target vllm-openai \ - --progress plain \ - -f docker/Dockerfile.rocm . - - # Push to ECR - docker push public.ecr.aws/q9t5s3a7/vllm-release-repo:$${BUILDKITE_COMMIT}-rocm + echo "" + echo " Successfully built and pushed ROCm release image" + echo " Image: public.ecr.aws/q9t5s3a7/vllm-release-repo:$${BUILDKITE_COMMIT}-rocm" + echo "" + env: + DOCKER_BUILDKIT: "1" + S3_BUCKET: "vllm-wheels" - echo "" - echo " Successfully built and pushed ROCm release image" - echo " Image: public.ecr.aws/q9t5s3a7/vllm-release-repo:$${BUILDKITE_COMMIT}-rocm" - echo "" - env: - DOCKER_BUILDKIT: "1" - S3_BUCKET: "vllm-wheels" + - label: "Publish nightly XPU image to DockerHub" + depends_on: + - create-manifest-xpu + if: build.env("NIGHTLY") == "1" + agents: + queue: small_cpu_queue_release + commands: + - "bash .buildkite/scripts/xpu/push-nightly-builds-xpu.sh" + - "bash .buildkite/scripts/cleanup-nightly-builds.sh nightly- vllm/vllm-openai-xpu" + plugins: + - docker-login#v3.0.0: + username: vllmbot + password-env: DOCKERHUB_TOKEN + env: + DOCKER_BUILDKIT: "1" + DOCKERHUB_USERNAME: "vllmbot" - - label: "Publish nightly XPU image to DockerHub" - depends_on: - - create-manifest-xpu - if: build.env("NIGHTLY") == "1" - agents: - queue: small_cpu_queue_release - commands: - - "bash .buildkite/scripts/xpu/push-nightly-builds-xpu.sh" - - "bash .buildkite/scripts/cleanup-nightly-builds.sh nightly- vllm/vllm-openai-xpu" - plugins: - - docker-login#v3.0.0: - username: vllmbot - password-env: DOCKERHUB_TOKEN - env: - DOCKER_BUILDKIT: "1" - DOCKERHUB_USERNAME: "vllmbot" - - - label: "Publish nightly ROCm image to DockerHub" - depends_on: - - build-rocm-release-image - if: build.env("NIGHTLY") == "1" - agents: - queue: small_cpu_queue_release - commands: - - "bash .buildkite/scripts/push-nightly-builds-rocm.sh" - # Clean up old nightly builds (keep only last 14) - - "bash .buildkite/scripts/cleanup-nightly-builds.sh nightly- vllm/vllm-openai-rocm" - - "bash .buildkite/scripts/cleanup-nightly-builds.sh base-nightly- vllm/vllm-openai-rocm" - plugins: - - docker-login#v3.0.0: - username: vllmbot - password-env: DOCKERHUB_TOKEN - env: - DOCKER_BUILDKIT: "1" - DOCKERHUB_USERNAME: "vllmbot" + - label: "Publish nightly ROCm image to DockerHub" + depends_on: + - build-rocm-release-image + if: build.env("NIGHTLY") == "1" + agents: + queue: small_cpu_queue_release + commands: + - "bash .buildkite/scripts/push-nightly-builds-rocm.sh" + # Clean up old nightly builds (keep only last 14) + - "bash .buildkite/scripts/cleanup-nightly-builds.sh nightly- vllm/vllm-openai-rocm" + - "bash .buildkite/scripts/cleanup-nightly-builds.sh base-nightly- vllm/vllm-openai-rocm" + plugins: + - docker-login#v3.0.0: + username: vllmbot + password-env: DOCKERHUB_TOKEN + env: + DOCKER_BUILDKIT: "1" + DOCKERHUB_USERNAME: "vllmbot" # ============================================================================= # Publish to DockerHub and PyPI (at the end so all builds complete first) From ace9fda495bc9a132584e054cc909d504164141f Mon Sep 17 00:00:00 2001 From: rasmith Date: Sun, 19 Jul 2026 13:41:52 -0500 Subject: [PATCH 32/51] [CI/Build][BugFix][The Rock][AMD] Add spawn method in vision examples to avoid reinitialization (#47932) Signed-off-by: Randall Smith Co-authored-by: Andreas Karatzas --- .../multimodal/vision_language_multi_image_offline.py | 5 +++++ examples/generate/multimodal/vision_language_offline.py | 3 +++ 2 files changed, 8 insertions(+) diff --git a/examples/generate/multimodal/vision_language_multi_image_offline.py b/examples/generate/multimodal/vision_language_multi_image_offline.py index c3541427742..1b9fddc99ed 100644 --- a/examples/generate/multimodal/vision_language_multi_image_offline.py +++ b/examples/generate/multimodal/vision_language_multi_image_offline.py @@ -17,6 +17,7 @@ from transformers import AutoProcessor, AutoTokenizer from vllm import LLM, EngineArgs, SamplingParams from vllm.lora.request import LoRARequest from vllm.multimodal.utils import fetch_image +from vllm.platforms import current_platform from vllm.utils.argparse_utils import FlexibleArgumentParser QUESTION = "What is the content of each image?" @@ -1443,6 +1444,8 @@ def run_generate( engine_args.seed = seed if tensor_parallel_size is not None: engine_args.tensor_parallel_size = tensor_parallel_size + if current_platform.is_rocm(): + os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "spawn" llm = LLM.from_engine_args(engine_args) sampling_params = SamplingParams( @@ -1484,6 +1487,8 @@ def run_chat( engine_args.seed = seed if tensor_parallel_size is not None: engine_args.tensor_parallel_size = tensor_parallel_size + if current_platform.is_rocm(): + os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "spawn" llm = LLM.from_engine_args(engine_args) sampling_params = ( diff --git a/examples/generate/multimodal/vision_language_offline.py b/examples/generate/multimodal/vision_language_offline.py index bddf6388ae6..661401046bd 100644 --- a/examples/generate/multimodal/vision_language_offline.py +++ b/examples/generate/multimodal/vision_language_offline.py @@ -21,6 +21,7 @@ from vllm.assets.image import ImageAsset from vllm.assets.video import VideoAsset from vllm.lora.request import LoRARequest from vllm.multimodal.image import convert_image_mode +from vllm.platforms import current_platform from vllm.utils.argparse_utils import FlexibleArgumentParser @@ -2646,6 +2647,8 @@ def main(args): if args.tensor_parallel_size is not None: engine_args.tensor_parallel_size = args.tensor_parallel_size engine_args = maybe_add_vit_cuda_graph_compilation_config(args, engine_args) + if current_platform.is_rocm(): + os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "spawn" llm = LLM.from_engine_args(engine_args) # Don't want to check the flag multiple times, so just hijack `prompts`. From 1dcbbd9caca3ddabd99a22747e1a889304cbeab5 Mon Sep 17 00:00:00 2001 From: "Kevin H. Luu" Date: Sun, 19 Jul 2026 19:21:25 -0700 Subject: [PATCH 33/51] [CI] Move compatible 1xL4 jobs to H200 35GB MIG (#43024) Signed-off-by: Simon Mo Co-authored-by: Simon Mo Co-authored-by: OpenAI Codex Co-authored-by: OpenAI Codex --- .buildkite/test_areas/cuda.yaml | 3 +- .buildkite/test_areas/entrypoints.yaml | 4 ++ .buildkite/test_areas/kernels.yaml | 2 + .buildkite/test_areas/misc.yaml | 3 +- .buildkite/test_areas/model_executor.yaml | 1 + .buildkite/test_areas/models_language.yaml | 21 ++++++++-- .buildkite/test_areas/models_multimodal.yaml | 1 + .buildkite/test_areas/pytorch.yaml | 40 ++++++++++++++++++- .buildkite/test_areas/quantization.yaml | 19 +++++++-- .buildkite/test_areas/rust_frontend.yaml | 1 + tests/kernels/mamba/test_mamba_ssm.py | 6 ++- .../models/language/generation/test_common.py | 11 +++++ .../models/language/generation/test_hybrid.py | 7 +++- tests/models/quantization/test_fp8.py | 2 + tests/quantization/test_fp8.py | 3 ++ tests/quantization/test_per_token_kv_cache.py | 5 ++- tests/utils.py | 15 ++++++- 17 files changed, 129 insertions(+), 15 deletions(-) diff --git a/.buildkite/test_areas/cuda.yaml b/.buildkite/test_areas/cuda.yaml index 99e1949fef7..2076eb27fc8 100644 --- a/.buildkite/test_areas/cuda.yaml +++ b/.buildkite/test_areas/cuda.yaml @@ -18,6 +18,7 @@ steps: - pytest -v -s cuda/test_platform_no_cuda_init.py - label: Cudagraph + device: h200_35gb key: cudagraph timeout_in_minutes: 30 source_file_dependencies: @@ -28,4 +29,4 @@ steps: commands: - pytest -v -s v1/cudagraph/test_cudagraph_dispatch.py - pytest -v -s v1/cudagraph/test_cudagraph_mode.py - - pytest -v -s v1/cudagraph/test_breakable_cudagraph.py \ No newline at end of file + - pytest -v -s v1/cudagraph/test_breakable_cudagraph.py diff --git a/.buildkite/test_areas/entrypoints.yaml b/.buildkite/test_areas/entrypoints.yaml index 499abcc6c82..503180b9130 100644 --- a/.buildkite/test_areas/entrypoints.yaml +++ b/.buildkite/test_areas/entrypoints.yaml @@ -57,6 +57,7 @@ steps: - image-build-amd - label: Entrypoints Integration (API Server OpenAI - Part 1) + device: h200_35gb key: entrypoints-integration-api-server-openai-part-1 timeout_in_minutes: 45 working_dir: "/vllm-workspace/tests" @@ -75,6 +76,7 @@ steps: - image-build-amd - label: Entrypoints Integration (API Server OpenAI - Part 2) + device: h200_35gb key: entrypoints-integration-api-server-openai-part-2 timeout_in_minutes: 45 working_dir: "/vllm-workspace/tests" @@ -94,6 +96,7 @@ steps: - image-build-amd - label: Entrypoints Integration (API Server Generate) + device: h200_35gb key: entrypoints-integration-api-server-generate timeout_in_minutes: 50 working_dir: "/vllm-workspace/tests" @@ -151,6 +154,7 @@ steps: - pytest -v -s entrypoints/multimodal - label: Entrypoints Integration (Pooling) + device: h200_35gb key: entrypoints-integration-pooling timeout_in_minutes: 50 working_dir: "/vllm-workspace/tests" diff --git a/.buildkite/test_areas/kernels.yaml b/.buildkite/test_areas/kernels.yaml index 150c3da57a3..1749b71abe9 100644 --- a/.buildkite/test_areas/kernels.yaml +++ b/.buildkite/test_areas/kernels.yaml @@ -15,6 +15,7 @@ steps: - pytest -v -s tests/kernels/ir - label: Kernels Core Operation Test + device: h200_35gb key: kernels-core-operation-test timeout_in_minutes: 120 source_file_dependencies: @@ -163,6 +164,7 @@ steps: - image-build-amd - label: Kernels Mamba Test + device: h200_35gb key: kernels-mamba-test timeout_in_minutes: 40 source_file_dependencies: diff --git a/.buildkite/test_areas/misc.yaml b/.buildkite/test_areas/misc.yaml index fb08373c9a6..521040c4ef2 100644 --- a/.buildkite/test_areas/misc.yaml +++ b/.buildkite/test_areas/misc.yaml @@ -64,8 +64,9 @@ steps: - image-build-amd - label: V1 Core + KV + Metrics + device: h200_35gb key: v1-core-kv-metrics - timeout_in_minutes: 60 + timeout_in_minutes: 80 source_file_dependencies: - vllm/config/ - vllm/distributed/ diff --git a/.buildkite/test_areas/model_executor.yaml b/.buildkite/test_areas/model_executor.yaml index 1f511590ffb..1e54538cc37 100644 --- a/.buildkite/test_areas/model_executor.yaml +++ b/.buildkite/test_areas/model_executor.yaml @@ -3,6 +3,7 @@ depends_on: - image-build steps: - label: Model Executor + device: h200_35gb key: model-executor timeout_in_minutes: 45 source_file_dependencies: diff --git a/.buildkite/test_areas/models_language.yaml b/.buildkite/test_areas/models_language.yaml index 6066c9bc708..b37f2dbd20f 100644 --- a/.buildkite/test_areas/models_language.yaml +++ b/.buildkite/test_areas/models_language.yaml @@ -21,6 +21,7 @@ steps: - image-build-amd - label: Language Models Tests (Extra Standard) %N + device: h200_35gb key: language-models-tests-extra-standard timeout_in_minutes: 40 source_file_dependencies: @@ -51,8 +52,8 @@ steps: - tests/models/language/pooling/test_classification.py - vllm/_aiter_ops.py - vllm/platforms/rocm.py - - label: Language Models Tests (Hybrid) %N + device: h200_35gb key: language-models-tests-hybrid timeout_in_minutes: 65 source_file_dependencies: @@ -63,8 +64,8 @@ steps: # Note: also needed to run plamo2 model in vLLM - uv pip install --system --no-build-isolation 'git+https://github.com/state-spaces/mamba@v2.3.0' - uv pip install --system --no-build-isolation 'git+https://github.com/Dao-AILab/causal-conv1d@v1.6.0' - # Shard hybrid language model tests - - pytest -v -s models/language/generation -m hybrid_model --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --shard-id=$$BUILDKITE_PARALLEL_JOB + # Shard the hybrid language model tests that are numerically stable on Hopper. + - pytest -v -s models/language/generation -m hybrid_model -k 'not granite-4.0-tiny-preview' --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --shard-id=$$BUILDKITE_PARALLEL_JOB parallelism: 2 mirror: amd: @@ -77,6 +78,20 @@ steps: - uv pip install --system --no-build-isolation 'git+https://github.com/Dao-AILab/causal-conv1d@v1.6.0' - pytest -v -s models/language/generation -m hybrid_model --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --shard-id=$$BUILDKITE_PARALLEL_JOB +# Granite 4 hybrid generation is sensitive to hardware-specific Triton SSD +# autotuning (https://github.com/vllm-project/vllm/issues/25194). Keep this one +# correctness test on L4 until its H200 output matches the Transformers reference. +- label: Language Models Tests (Granite L4 Compatibility) + key: language-models-tests-granite-l4-compatibility + timeout_in_minutes: 65 + source_file_dependencies: + - vllm/ + - tests/models/language/generation + commands: + - uv pip install --system --no-build-isolation 'git+https://github.com/state-spaces/mamba@v2.3.0' + - uv pip install --system --no-build-isolation 'git+https://github.com/Dao-AILab/causal-conv1d@v1.6.0' + - pytest -v -s models/language/generation -m hybrid_model -k 'granite-4.0-tiny-preview' + - label: Language Models Test (Extended Generation) # 80min device: h200_35gb key: language-models-test-extended-generation diff --git a/.buildkite/test_areas/models_multimodal.yaml b/.buildkite/test_areas/models_multimodal.yaml index 32a27a56fde..2a73eb4a47e 100644 --- a/.buildkite/test_areas/models_multimodal.yaml +++ b/.buildkite/test_areas/models_multimodal.yaml @@ -119,6 +119,7 @@ steps: - vllm/model_executor/model_loader/ - label: Multi-Modal Models (Extended Generation 1) + device: h200_35gb key: multi-modal-models-extended-generation-1 optional: true source_file_dependencies: diff --git a/.buildkite/test_areas/pytorch.yaml b/.buildkite/test_areas/pytorch.yaml index cdc5a1fd044..ab707c37ba4 100644 --- a/.buildkite/test_areas/pytorch.yaml +++ b/.buildkite/test_areas/pytorch.yaml @@ -116,8 +116,9 @@ steps: - image-build-amd - label: PyTorch Fullgraph Smoke Test + device: h200_35gb key: pytorch-fullgraph-smoke-test - timeout_in_minutes: 60 + timeout_in_minutes: 90 source_file_dependencies: - vllm/__init__.py - vllm/_aiter_ops.py @@ -149,7 +150,42 @@ steps: # as it is a heavy test that is covered in other steps. # Use `find` to launch multiple instances of pytest so that # they do not suffer from https://github.com/vllm-project/vllm/issues/28965 - - "find compile/fullgraph/ -name 'test_*.py' -not -name 'test_full_graph.py' -print0 | xargs -0 -n1 -I{} pytest -s -v '{}'" + - "find compile/fullgraph/ -name 'test_*.py' -not -name 'test_full_cudagraph.py' -not -name 'test_full_graph.py' -print0 | xargs -0 -n1 -I{} pytest -s -v '{}'" + +# Hopper-only DeepSeek-V2-Lite cases in this file require two 29.3-GiB model +# instances and cannot fit a 35GB MIG slice. L4 retains the original coverage: +# those SM90 cases skip while the architecture-compatible cases still run. +- label: PyTorch Fullgraph CUDAGraph (L4 Compatibility) + key: pytorch-fullgraph-cudagraph-l4-compatibility + timeout_in_minutes: 60 + source_file_dependencies: + - vllm/__init__.py + - vllm/_aiter_ops.py + - vllm/_custom_ops.py + - vllm/compilation/ + - vllm/config/ + - vllm/distributed/ + - vllm/engine/ + - vllm/env_override.py + - vllm/envs.py + - vllm/forward_context.py + - vllm/inputs/ + - vllm/ir/ + - vllm/kernels/ + - vllm/logger.py + - vllm/model_executor/ + - vllm/multimodal/ + - vllm/platforms/ + - vllm/plugins/ + - vllm/sampling_params.py + - vllm/sequence.py + - vllm/transformers_utils/ + - vllm/triton_utils/ + - vllm/utils/ + - vllm/v1/ + - tests/compile + commands: + - pytest -s -v compile/fullgraph/test_full_cudagraph.py - label: PyTorch Fullgraph key: pytorch-fullgraph diff --git a/.buildkite/test_areas/quantization.yaml b/.buildkite/test_areas/quantization.yaml index ce3e58e501b..16782f727f0 100644 --- a/.buildkite/test_areas/quantization.yaml +++ b/.buildkite/test_areas/quantization.yaml @@ -3,8 +3,11 @@ depends_on: - image-build steps: - label: Quantization + device: h200_35gb key: quantization - timeout_in_minutes: 60 + timeout_in_minutes: 75 + env: + VLLM_USE_V2_MODEL_RUNNER: "0" source_file_dependencies: - csrc/ - vllm/model_executor/layers/quantization @@ -19,9 +22,13 @@ steps: # TODO(jerryzh168): resolve the above comment - uv pip install --system torchao==0.17.0 --index-url https://download.pytorch.org/whl/cu130 - uv pip install --system conch-triton-kernels - - VLLM_TEST_FORCE_LOAD_FORMAT=auto pytest -v -s quantization/ --ignore quantization/test_blackwell_moe.py + # The SM90-only checkpoint currently contains a removed weight_chan_scale + # parameter. It was not exercised by the previous L4 job. + - VLLM_TEST_FORCE_LOAD_FORMAT=auto pytest -v -s quantization/ --ignore quantization/test_blackwell_moe.py -k 'not test_compressed_tensors_w4a8_fp8' --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT + parallelism: 8 - label: Quantized Fusions + device: h200_35gb key: quantized-fusions timeout_in_minutes: 20 source_file_dependencies: @@ -52,10 +59,14 @@ steps: - pytest -s -v tests/quantization/test_blackwell_moe.py - label: Quantized Models Test + device: h200_35gb key: quantized-models-test - timeout_in_minutes: 50 + timeout_in_minutes: 65 + env: + VLLM_USE_V2_MODEL_RUNNER: "0" source_file_dependencies: - vllm/model_executor/layers/quantization - tests/models/quantization commands: - - pytest -v -s models/quantization + - pytest -v -s models/quantization --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT + parallelism: 3 diff --git a/.buildkite/test_areas/rust_frontend.yaml b/.buildkite/test_areas/rust_frontend.yaml index 9e5e09c3ec3..c1599c1c6ea 100644 --- a/.buildkite/test_areas/rust_frontend.yaml +++ b/.buildkite/test_areas/rust_frontend.yaml @@ -81,6 +81,7 @@ steps: - pytest -s entrypoints/openai/correctness/test_lmeval.py::test_lm_eval_accuracy_v1_engine - label: Rust Frontend Tool Use + device: h200_35gb timeout_in_minutes: 25 working_dir: "/vllm-workspace/tests" source_file_dependencies: diff --git a/tests/kernels/mamba/test_mamba_ssm.py b/tests/kernels/mamba/test_mamba_ssm.py index 7350b646523..81b57d4b16b 100644 --- a/tests/kernels/mamba/test_mamba_ssm.py +++ b/tests/kernels/mamba/test_mamba_ssm.py @@ -347,7 +347,11 @@ def test_selective_state_update(dim, dstate, has_z, itype): rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2) if itype == torch.bfloat16: rtol, atol = 1e-2, 5e-2 - if current_platform.is_rocm() or current_platform.is_xpu(): + if ( + current_platform.is_rocm() + or current_platform.is_xpu() + or current_platform.is_device_capability_family(90) + ): atol *= 2 # set seed set_random_seed(0) diff --git a/tests/models/language/generation/test_common.py b/tests/models/language/generation/test_common.py index 50c87d7729e..18d06fe0fd3 100644 --- a/tests/models/language/generation/test_common.py +++ b/tests/models/language/generation/test_common.py @@ -188,6 +188,16 @@ def test_models( prompt_embeds.append(embed.squeeze(0)) + vllm_kwargs = {} + if ( + model == "bigscience/bloom-560m" + and current_platform.is_device_capability_family(90) + ): + # On SM90, the metadata builder otherwise selects FA3 AOT scheduling + # before Bloom's ALiBi layers fall back to FA2. Pinning FA2 keeps the + # builder and layer consistent and preserves the L4 test path. + vllm_kwargs["attention_config"] = {"flash_attn_version": 2} + with vllm_runner( model, tokenizer_name=model_info.tokenizer or model, @@ -200,6 +210,7 @@ def test_models( max_num_seqs=1 if current_platform.is_rocm() else 2, enable_prompt_embeds=use_prompt_embeds, compilation_config={"cudagraph_capture_sizes": [1, 2]}, + **vllm_kwargs, ) as vllm_model: vllm_outputs = vllm_model.generate_greedy_logprobs( example_prompts, max_tokens, num_logprobs diff --git a/tests/models/language/generation/test_hybrid.py b/tests/models/language/generation/test_hybrid.py index f06998e07f6..3fee0662bc6 100644 --- a/tests/models/language/generation/test_hybrid.py +++ b/tests/models/language/generation/test_hybrid.py @@ -384,8 +384,13 @@ def test_fp32_cache_state( example_prompts, max_tokens, num_logprobs ) + # Leave enough headroom for repeated engine initialization on a + # 32.5 GiB MIG. with vllm_runner( - model, max_num_seqs=MAX_NUM_SEQS, **{cache_dtype_param: "float32"} + model, + max_num_seqs=MAX_NUM_SEQS, + gpu_memory_utilization=0.9, + **{cache_dtype_param: "float32"}, ) as vllm_model: vllm_outputs = vllm_model.generate_greedy_logprobs( example_prompts, max_tokens, num_logprobs diff --git a/tests/models/quantization/test_fp8.py b/tests/models/quantization/test_fp8.py index 5f3c6547612..6a13794427a 100644 --- a/tests/models/quantization/test_fp8.py +++ b/tests/models/quantization/test_fp8.py @@ -67,6 +67,8 @@ def test_models( if kv_cache_dtype == "fp8_e5m2" and current_platform.is_rocm(): pytest.skip(f"{kv_cache_dtype} is currently not supported on ROCm/HIP.") + if kv_cache_dtype == "fp8_e5m2" and current_platform.is_cuda(): + pytest.skip(f"{kv_cache_dtype} is not supported by FLASH_ATTN on CUDA.") if not ( current_platform.is_xpu() diff --git a/tests/quantization/test_fp8.py b/tests/quantization/test_fp8.py index abb995257e3..0ae51652c62 100644 --- a/tests/quantization/test_fp8.py +++ b/tests/quantization/test_fp8.py @@ -93,6 +93,9 @@ def test_online_quantization( use_rocm_aiter: bool, monkeypatch, ) -> None: + if kv_cache_dtype == "fp8" and current_platform.is_device_capability_family(90): + pytest.skip("FA3 currently rejects FP8 KV cache output dtype on SM90") + if use_rocm_aiter: monkeypatch.setenv("VLLM_ROCM_USE_AITER", "1") diff --git a/tests/quantization/test_per_token_kv_cache.py b/tests/quantization/test_per_token_kv_cache.py index 3715aaf3a52..9fa4d97add4 100644 --- a/tests/quantization/test_per_token_kv_cache.py +++ b/tests/quantization/test_per_token_kv_cache.py @@ -717,7 +717,10 @@ def test_triton_unified_attention_per_token_head_scale( # Coarser quantization → wider tolerance. if is_int4: - atol, rtol = 0.5, 0.5 + # Hopper's attention reduction order can move a few BF16 elements by + # just over 1.0 after INT4 quantization. + atol = 1.1 if current_platform.is_device_capability_family(90) else 0.5 + rtol = 0.5 else: atol, rtol = 5e-2, 5e-2 torch.testing.assert_close(output_q, output_ref, atol=atol, rtol=rtol) diff --git a/tests/utils.py b/tests/utils.py index 2a3bdb91fe0..645e988dbb8 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -108,6 +108,7 @@ if current_platform.is_rocm(): elif current_platform.is_cuda(): from vllm.third_party.pynvml import ( nvmlDeviceGetHandleByIndex, + nvmlDeviceGetHandleByUUID, nvmlDeviceGetMemoryInfo, nvmlInit, nvmlShutdown, @@ -1521,6 +1522,18 @@ def get_physical_device_indices(devices: list[int]): return [index_mapping[i] for i in devices if i in index_mapping] +def get_nvml_device_handle(device: int): + visible_devices = os.environ.get("NVIDIA_VISIBLE_DEVICES") + if visible_devices is not None: + identifiers = visible_devices.split(",") + if device < len(identifiers): + identifier = identifiers[device] + if identifier.startswith(("GPU-", "MIG-")): + return nvmlDeviceGetHandleByUUID(identifier) + + return nvmlDeviceGetHandleByIndex(device) + + @_nvml() def record_gpu_memory_usage_stats( *, @@ -1534,7 +1547,7 @@ def record_gpu_memory_usage_stats( gb_used = mem_info["vram_used"] / 2**10 gb_total = mem_info["vram_total"] / 2**10 else: - dev_handle = nvmlDeviceGetHandleByIndex(device) + dev_handle = get_nvml_device_handle(device) mem_info = nvmlDeviceGetMemoryInfo(dev_handle) gb_used = mem_info.used / 2**30 gb_total = mem_info.total / 2**30 From 2730b657c4aa92a1a342a243514bbbd68804866f Mon Sep 17 00:00:00 2001 From: Thien Tran Date: Mon, 20 Jul 2026 10:45:56 +0800 Subject: [PATCH 34/51] [Bugfix] Fix broken NVVM caused by CuteDSL 4.6.0 (#49108) Signed-off-by: Thien Tran --- .buildkite/test_areas/kernels.yaml | 7 +++++++ vllm/cute_utils/_tcgen05.py | 8 ++++---- 2 files changed, 11 insertions(+), 4 deletions(-) diff --git a/.buildkite/test_areas/kernels.yaml b/.buildkite/test_areas/kernels.yaml index 1749b71abe9..3618cb6e969 100644 --- a/.buildkite/test_areas/kernels.yaml +++ b/.buildkite/test_areas/kernels.yaml @@ -237,6 +237,11 @@ steps: - vllm/model_executor/kernels/linear/cute_dsl/ll_bf16.py - vllm/model_executor/kernels/linear/cute_dsl/_ll_bf16_dotprod.py - vllm/model_executor/kernels/linear/cute_dsl/_ll_bf16_splitk.py + - vllm/cute_utils/ + - vllm/model_executor/layers/mamba/ops/gdn_chunk_cutedsl/ + - vllm/model_executor/layers/fused_moe/router/bf16x3_router_gemm_cutedsl.py + - tests/kernels/mamba/test_gdn_prefill_cutedsl.py + - tests/kernels/test_bf16x3_router_gemm_cutedsl.py - tests/kernels/test_ll_bf16_gemm.py - tests/kernels/test_top_k_per_row.py commands: @@ -266,6 +271,8 @@ steps: - pytest -v -s tests/kernels/moe/test_flashinfer_moe.py - pytest -v -s tests/kernels/moe/test_trtllm_nvfp4_moe.py - pytest -v -s tests/kernels/moe/test_cutedsl_moe.py + - pytest -v -s tests/kernels/mamba/test_gdn_prefill_cutedsl.py + - pytest -v -s tests/kernels/test_bf16x3_router_gemm_cutedsl.py - pytest -v -s tests/kernels/test_ll_bf16_gemm.py # e2e - pytest -v -s tests/models/quantization/test_nvfp4.py diff --git a/vllm/cute_utils/_tcgen05.py b/vllm/cute_utils/_tcgen05.py index 9367fa12a11..325c30498c7 100644 --- a/vllm/cute_utils/_tcgen05.py +++ b/vllm/cute_utils/_tcgen05.py @@ -10,8 +10,8 @@ from cutlass.cutlass_dsl import dsl_user_op NVVM_CTA_GROUP_MAP = [ None, - nvvm.Tcgen05GroupKind.CTA_1, - nvvm.Tcgen05GroupKind.CTA_2, + nvvm.CTAGroupKind.CTA_1, + nvvm.CTAGroupKind.CTA_2, ] LDST_MAP = { "32x32b": (nvvm.Tcgen05LdStShape.SHAPE_32X32B, 1), @@ -136,7 +136,7 @@ def commit(mbar, cta_mask=None, cta_group: int = 1, *, loc=None, ip=None): group = NVVM_CTA_GROUP_MAP[cta_group] if cutlass.const_expr(cta_mask is not None): with cute.arch.elect_one(): - nvvm.tcgen05_commit_arrive( + nvvm.tcgen05_commit( mbar_llvm, multicast_mask=cta_mask.ir_value(loc=loc, ip=ip), group=group, @@ -145,7 +145,7 @@ def commit(mbar, cta_mask=None, cta_group: int = 1, *, loc=None, ip=None): ) else: with cute.arch.elect_one(): - nvvm.tcgen05_commit_arrive(mbar_llvm, group=group, loc=loc, ip=ip) + nvvm.tcgen05_commit(mbar_llvm, group=group, loc=loc, ip=ip) @dsl_user_op From 752bd106477f1a72385fb38616abfe198dd67815 Mon Sep 17 00:00:00 2001 From: Andreas Karatzas Date: Sun, 19 Jul 2026 23:02:03 -0500 Subject: [PATCH 35/51] [ROCm][CI] Fix sparse MLA metadata sync fixture (#49128) Signed-off-by: Andreas Karatzas Co-authored-by: OpenAI Codex --- .../attention/test_rocm_aiter_mla_sparse_metadata_sync.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/kernels/attention/test_rocm_aiter_mla_sparse_metadata_sync.py b/tests/kernels/attention/test_rocm_aiter_mla_sparse_metadata_sync.py index 4cb50e6abb6..cc9ac6b8d71 100644 --- a/tests/kernels/attention/test_rocm_aiter_mla_sparse_metadata_sync.py +++ b/tests/kernels/attention/test_rocm_aiter_mla_sparse_metadata_sync.py @@ -58,6 +58,7 @@ def _make_builder(): max_num_batched_tokens + 1, dtype=torch.int32, device="cpu" ) builder._num_attention_heads = 16 + builder._num_compute_units = current_platform.num_compute_units() builder._mla_work_meta_data = torch.empty(1, dtype=torch.int32, device="cpu") builder._mla_work_indptr = torch.empty(1, dtype=torch.int32, device="cpu") builder._mla_work_info_set = torch.empty(1, dtype=torch.int32, device="cpu") @@ -116,6 +117,7 @@ def test_sparse_persistent_metadata_syncs_only_after_recompute(monkeypatch): assert events == ["metadata", "sync"] assert fake_get_mla_metadata_v1_mock.call_count == 1 + assert fake_get_mla_metadata_v1_mock.call_args.kwargs["max_split_per_batch"] == 1 events.clear() From dcfebf93f4eccf30f71872283331eee757915daf Mon Sep 17 00:00:00 2001 From: aoshen02 Date: Mon, 20 Jul 2026 12:17:18 +0800 Subject: [PATCH 36/51] =?UTF-8?q?[Bugfix]=20Fix=20logprobs=20token-string?= =?UTF-8?q?=20collision=20from=20SentencePiece=20space=E2=80=A6=20(#48674)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Allen Shen Co-authored-by: mvanhorn Co-authored-by: Claude Opus 4.6 (1M context) Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> --- tests/tokenizers_/test_detokenize.py | 120 +++++++++++++++++++++++++++ vllm/tokenizers/detokenizer_utils.py | 82 ++++++++++++++++-- 2 files changed, 194 insertions(+), 8 deletions(-) diff --git a/tests/tokenizers_/test_detokenize.py b/tests/tokenizers_/test_detokenize.py index 23eaca9fc36..ad399ba36dc 100644 --- a/tests/tokenizers_/test_detokenize.py +++ b/tests/tokenizers_/test_detokenize.py @@ -8,6 +8,7 @@ import pytest from transformers import AutoTokenizer, PythonBackend, TokenizersBackend from vllm.sampling_params import SamplingParams +from vllm.tokenizers.detokenizer_utils import convert_ids_list_to_tokens from vllm.tokenizers.mistral import MistralTokenizer from vllm.v1.engine import EngineCoreRequest from vllm.v1.engine.detokenizer import ( @@ -239,3 +240,122 @@ def test_oov_decode(tokenizer, fast): assert decoded_text == "" assert out_ids == [len(tokenizer)] + + +# ---------- convert_ids_list_to_tokens collision tests ---------- + + +class _MockBackend: + """Fake backend_tokenizer that exposes pre_tokenizer config.""" + + def __init__(self, pre_tokenizer_type, replacement=None): + import json + + pre: dict = {"type": pre_tokenizer_type} + if replacement is not None: + pre["replacement"] = replacement + self._config = json.dumps({"pre_tokenizer": pre}) + + def to_str(self): + return self._config + + +class _MockTokenizer: + """Minimal tokenizer mock for testing convert_ids_list_to_tokens.""" + + def __init__( + self, + raw_tokens: dict[int, str], + decoded_tokens: dict[int, str], + pre_tokenizer_type: str = "Metaspace", + replacement: str | None = "▁", + ): + self._raw = raw_tokens + self._decoded = decoded_tokens + self.backend_tokenizer = _MockBackend(pre_tokenizer_type, replacement) + + def convert_ids_to_tokens( + self, ids: list[int], skip_special_tokens: bool = False + ) -> list[str]: + return [self._raw[tid] for tid in ids] + + def decode(self, ids: list[int], skip_special_tokens: bool = False) -> str: + return "".join(self._decoded[tid] for tid in ids) + + +def test_sentencepiece_leading_space_preserved(): + """▁true and true must produce distinct strings.""" + tok = _MockTokenizer( + raw_tokens={0: "▁true", 1: "true", 2: "▁false", 3: "false"}, + decoded_tokens={0: "true", 1: "true", 2: "false", 3: "false"}, + ) + result = convert_ids_list_to_tokens(tok, [0, 1, 2, 3]) + assert result == [" true", "true", " false", "false"] + + # No dict collision when used as top_logprobs keys + logprobs = dict(zip(result, [-0.1, -0.2, -0.3, -0.4])) + assert len(logprobs) == 4 + + +def test_whitespace_run_tokens_stay_distinct(): + """▁, ▁▁, ▁▁▁ must produce different-length space strings.""" + tok = _MockTokenizer( + raw_tokens={0: "▁", 1: "▁▁", 2: "▁▁▁"}, + decoded_tokens={0: "", 1: " ", 2: " "}, + ) + result = convert_ids_list_to_tokens(tok, [0, 1, 2]) + assert result == [" ", " ", " "] + + +def test_bpe_leading_space_already_preserved(): + """GPT-2 BPE: Ġtrue already decodes to ' true', no fix needed.""" + tok = _MockTokenizer( + raw_tokens={0: "Ġtrue", 1: "true"}, + decoded_tokens={0: " true", 1: "true"}, + pre_tokenizer_type="ByteLevel", + replacement=None, + ) + result = convert_ids_list_to_tokens(tok, [0, 1]) + assert result == [" true", "true"] + + +def test_logprobs_count_stable_across_k(): + """logprobs=4 and logprobs=10 must return 4 and 10 entries.""" + tok = _MockTokenizer( + raw_tokens={ + 0: "▁true", + 1: "a", + 2: "b", + 3: "c", + 4: "true", + 5: "d", + 6: "e", + 7: "f", + 8: "g", + 9: "h", + }, + decoded_tokens={ + 0: "true", + 1: "a", + 2: "b", + 3: "c", + 4: "true", + 5: "d", + 6: "e", + 7: "f", + 8: "g", + 9: "h", + }, + ) + ids = list(range(10)) + lps = [-0.1 * (i + 1) for i in range(10)] + + tokens4 = convert_ids_list_to_tokens(tok, ids[:4]) + top4 = dict(zip(tokens4, lps[:4])) + + tokens10 = convert_ids_list_to_tokens(tok, ids) + top10 = dict(zip(tokens10, lps)) + + assert len(top4) == 4 + assert len(top10) == 10 + assert top4[" true"] == top10[" true"] diff --git a/vllm/tokenizers/detokenizer_utils.py b/vllm/tokenizers/detokenizer_utils.py index 8e73d5dc537..516f1b70e73 100644 --- a/vllm/tokenizers/detokenizer_utils.py +++ b/vllm/tokenizers/detokenizer_utils.py @@ -1,6 +1,9 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from __future__ import annotations + +import json from vllm.tokenizers import TokenizerLike @@ -56,6 +59,63 @@ def _convert_tokens_to_string_with_added_encoders( INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET = 5 +_CACHED_MARKER_KEY = "_vllm_space_marker_cache" +_NOT_CACHED = "__not_computed__" + + +def _get_leading_space_marker(tokenizer: TokenizerLike) -> str | None: + """Read the space marker from the tokenizer's pre_tokenizer config. + + Only Metaspace pre_tokenizers (used by SentencePiece-based models like + Llama, Mistral, T5) have a replacement character whose leading instance + gets stripped by decode(). ByteLevel (GPT-2), BertPreTokenizer (BERT), + and others do not have this issue. + + Returns the marker character, or None if decode() is safe for single + tokens. + """ + cached = getattr(tokenizer, _CACHED_MARKER_KEY, _NOT_CACHED) + if cached is not _NOT_CACHED: + return cached # type: ignore[return-value] + + backend = getattr(tokenizer, "backend_tokenizer", None) + if backend is None: + result = None + else: + result = None + try: + config = json.loads(backend.to_str()) + except Exception: + pass + else: + pre = config.get("pre_tokenizer") or {} + pre_type = pre.get("type") + if pre_type == "Metaspace": + result = pre.get("replacement", "▁") + elif pre_type == "Sequence": + for sub in pre.get("pretokenizers", []): + if sub.get("type") == "Metaspace": + result = sub.get("replacement", "▁") + break + + setattr(tokenizer, _CACHED_MARKER_KEY, result) + return result + + +def _restore_leading_spaces(raw_token: str, token_str: str, marker: str) -> str: + """Restore leading spaces that decode() stripped from a raw vocab piece.""" + num_markers = 0 + for ch in raw_token: + if ch != marker: + break + num_markers += 1 + if num_markers == 0: + return token_str + existing = len(token_str) - len(token_str.lstrip(" ")) + missing = num_markers - existing + return " " * missing + token_str if missing > 0 else token_str + + def convert_prompt_ids_to_tokens( tokenizer: TokenizerLike, prompt_ids: list[int], @@ -86,6 +146,10 @@ def convert_ids_list_to_tokens( ) -> list[str]: """Detokenize the input ids individually. + Uses decode() for human-readable output, then checks the raw vocab + piece via convert_ids_to_tokens() to restore any leading spaces that + decode() stripped (SentencePiece add_dummy_prefix inverse). + Args: tokenizer: tokenizer used by model under test token_ids: convert these tokens (Python list form) @@ -94,14 +158,16 @@ def convert_ids_list_to_tokens( Python list of token string representations """ - token_str_lst = [] - for token_id in token_ids: - # use default skip_special_tokens. - token_str = tokenizer.decode([token_id]) - if token_str is None: - token_str = "" - token_str_lst.append(token_str) - return token_str_lst + if not token_ids: + return [] + marker = _get_leading_space_marker(tokenizer) + if marker is None: + return [tokenizer.decode([tid]) or "" for tid in token_ids] + raw_tokens = tokenizer.convert_ids_to_tokens(token_ids) + return [ + _restore_leading_spaces(raw, tokenizer.decode([tid]) or "", marker) + for tid, raw in zip(token_ids, raw_tokens) + ] # Based on From 5c9f6557d7845b52917ef013011c83618439b17f Mon Sep 17 00:00:00 2001 From: Akash kaothalkar <61960177+Akashcodes732@users.noreply.github.com> Date: Mon, 20 Jul 2026 11:45:16 +0530 Subject: [PATCH 37/51] [Hardware][CPU] Enable granite-4 model on cpu (#47641) Signed-off-by: Akash Kaothalkar Signed-off-by: Akash Kaothalkar Signed-off-by: Akash Kaothalkar Signed-off-by: Akash kaothalkar Co-authored-by: Akash Kaothalkar Co-authored-by: Akash Kaothalkar Co-authored-by: Akash Kaothalkar Co-authored-by: Akash kaothalkar Co-authored-by: Li, Jiang --- .buildkite/hardware_tests/cpu.yaml | 6 +- .../scripts/hardware_ci/run-cpu-test-arm.sh | 5 +- cmake/cpu_extension.cmake | 3 + csrc/cpu/cpu_types_vsx.hpp | 15 +- csrc/cpu/mamba_cpu.cpp | 285 +++++++++++++ csrc/cpu/mamba_kernels.hpp | 382 ++++++++++++++++++ csrc/cpu/torch_bindings.cpp | 50 +++ tests/kernels/mamba/cpu/test_cpu_gdn_ops.py | 2 +- tests/kernels/mamba/test_causal_conv1d.py | 11 +- tests/kernels/mamba/test_mamba_ssm.py | 50 ++- vllm/_custom_ops.py | 87 ++++ vllm/config/mamba.py | 1 + .../layers/mamba/ops/causal_conv1d.py | 12 + .../layers/mamba/ops/cpu/causal_conv1d.py | 76 +++- .../layers/mamba/ops/cpu/gdn_attention.py | 68 ++-- .../layers/mamba/ops/cpu/mamba_ssm.py | 144 +++++++ .../layers/mamba/ops/mamba_ssm.py | 10 + .../layers/mamba/ops/ssd_combined.py | 8 + .../layers/mamba/ops/ssu_dispatch.py | 85 +++- .../model_executor/layers/mamba/short_conv.py | 37 +- vllm/model_executor/layers/utils.py | 8 +- 21 files changed, 1283 insertions(+), 62 deletions(-) create mode 100644 csrc/cpu/mamba_cpu.cpp create mode 100644 csrc/cpu/mamba_kernels.hpp create mode 100644 vllm/model_executor/layers/mamba/ops/cpu/mamba_ssm.py diff --git a/.buildkite/hardware_tests/cpu.yaml b/.buildkite/hardware_tests/cpu.yaml index ebfd1c7524a..fd2d45be5a0 100644 --- a/.buildkite/hardware_tests/cpu.yaml +++ b/.buildkite/hardware_tests/cpu.yaml @@ -18,6 +18,8 @@ steps: - tests/kernels/quantization/test_cpu_fp8_scaled_mm.py - tests/kernels/mamba/cpu/test_cpu_gdn_ops.py - tests/kernels/mamba/test_cpu_short_conv.py + - tests/kernels/mamba/test_causal_conv1d.py + - tests/kernels/mamba/test_mamba_ssm.py commands: - | bash .buildkite/scripts/hardware_ci/run-cpu-test.sh 30m " @@ -28,7 +30,9 @@ steps: pytest -x -v -s tests/kernels/test_onednn.py pytest -x -v -s tests/kernels/test_awq_int4_to_int8.py pytest -x -v -s tests/kernels/quantization/test_cpu_fp8_scaled_mm.py - pytest -x -v -s tests/kernels/mamba/cpu/test_cpu_gdn_ops.py" + pytest -x -v -s tests/kernels/mamba/cpu/test_cpu_gdn_ops.py + pytest -x -v -s tests/kernels/mamba/test_causal_conv1d.py + pytest -x -v -s tests/kernels/mamba/test_mamba_ssm.py" # Note: SDE can't be downloaded from CI host because of AWS WAF # - label: CPU-Compatibility Tests diff --git a/.buildkite/scripts/hardware_ci/run-cpu-test-arm.sh b/.buildkite/scripts/hardware_ci/run-cpu-test-arm.sh index 2d11dd477ea..09396a9697a 100755 --- a/.buildkite/scripts/hardware_ci/run-cpu-test-arm.sh +++ b/.buildkite/scripts/hardware_ci/run-cpu-test-arm.sh @@ -40,7 +40,9 @@ function cpu_tests() { pytest -x -v -s tests/kernels/moe/test_cpu_fused_moe.py pytest -x -v -s tests/kernels/mamba/cpu/test_cpu_gdn_ops.py pytest -x -v -s tests/kernels/moe/test_cpu_int4_moe.py - pytest -x -v -s tests/kernels/mamba/test_cpu_short_conv.py" + pytest -x -v -s tests/kernels/mamba/test_cpu_short_conv.py + pytest -x -v -s tests/kernels/mamba/test_causal_conv1d.py + pytest -x -v -s tests/kernels/mamba/test_mamba_ssm.py" # skip tests requiring model downloads if HF_TOKEN is not set # due to rate-limits @@ -97,3 +99,4 @@ function cpu_tests() { # All of CPU tests are expected to be finished less than 40 mins. export -f cpu_tests timeout 2h bash -c cpu_tests + diff --git a/cmake/cpu_extension.cmake b/cmake/cpu_extension.cmake index 3aca9bcea91..4e7837df0d7 100644 --- a/cmake/cpu_extension.cmake +++ b/cmake/cpu_extension.cmake @@ -430,6 +430,7 @@ set(VLLM_EXT_SRC "csrc/cpu/layernorm.cpp" "csrc/cpu/mla_decode.cpp" "csrc/cpu/pos_encoding.cpp" + "csrc/cpu/mamba_cpu.cpp" "csrc/moe/dynamic_4bit_int_moe_cpu.cpp" "csrc/cpu/cpu_attn.cpp" "csrc/cpu/torch_bindings.cpp") @@ -489,6 +490,7 @@ if (ENABLE_X86_ISA) "csrc/cpu/spec_decode_utils.cpp" "csrc/cpu/cpu_attn.cpp" "csrc/cpu/dnnl_kernels.cpp" + "csrc/cpu/mamba_cpu.cpp" "csrc/cpu/torch_bindings.cpp" # TODO: Remove these files "csrc/cpu/activation.cpp" @@ -502,6 +504,7 @@ if (ENABLE_X86_ISA) "csrc/cpu/utils.cpp" "csrc/cpu/spec_decode_utils.cpp" "csrc/cpu/cpu_attn.cpp" + "csrc/cpu/mamba_cpu.cpp" "csrc/cpu/dnnl_kernels.cpp" "csrc/cpu/torch_bindings.cpp" # TODO: Remove these files diff --git a/csrc/cpu/cpu_types_vsx.hpp b/csrc/cpu/cpu_types_vsx.hpp index 250c870dbe4..64fe961da22 100644 --- a/csrc/cpu/cpu_types_vsx.hpp +++ b/csrc/cpu/cpu_types_vsx.hpp @@ -336,13 +336,14 @@ struct FP32Vec8 : public Vec { reg.val[1] = fp16_to_fp32_bits(raw_lo); } float reduce_sum() const { - AliasReg ar; - ar.reg = reg; - float result = 0; - unroll_loop( - [&result, &ar](int i) { result += ar.values[i]; }); - - return result; + // VSX horizontal reduction: 3 vector ops instead of 8 scalar adds. + // Step 1: pairwise sum of the two 4-wide halves + __vector float s = vec_add(reg.val[0], reg.val[1]); + // Step 2: rotate by 8 bytes (2 floats) and add + s = vec_add(s, vec_sld(s, s, 8)); + // Step 3: rotate by 4 bytes (1 float) and add => all lanes hold total + s = vec_add(s, vec_sld(s, s, 4)); + return vec_extract(s, 0); } FP32Vec8 exp() const { f32x4x2_t out; diff --git a/csrc/cpu/mamba_cpu.cpp b/csrc/cpu/mamba_cpu.cpp new file mode 100644 index 00000000000..54e4f99c2d6 --- /dev/null +++ b/csrc/cpu/mamba_cpu.cpp @@ -0,0 +1,285 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright contributors to the vLLM project +// +// CPU at::Tensor wrappers for Mamba decode-step kernels defined in +// mamba_kernels.hpp. + +#include "cpu/mamba_kernels.hpp" + +#include +#include +#include + +#include "cpu_types.hpp" + +// --------------------------------------------------------------------------- +// causal_conv1d_update +// --------------------------------------------------------------------------- +at::Tensor causal_conv1d_update_cpu_impl( + at::Tensor& x, at::Tensor& conv_state, const at::Tensor& weight, + const c10::optional& bias, + const c10::optional& activation, + const c10::optional& conv_state_indices, + const c10::optional& query_start_loc, int64_t pad_slot_id) { + bool do_silu = false; + if (activation.has_value()) { + const std::string& act = activation.value(); + do_silu = (act == "silu" || act == "swish"); + } + + at::ScalarType dtype = x.scalar_type(); + + // Input x: contiguous in native dtype. + at::Tensor x_c = x.is_contiguous() ? x : x.contiguous(); + + // conv_state: NEVER copy the full paged tensor just for layout reasons. + // If the dtype matches we work directly on conv_state (contiguous or not) + // by extracting strides and passing them to the kernel. + // Only a dtype-conversion copy is made when types differ (rare for BF16). + bool state_type_ok = (conv_state.scalar_type() == dtype); + at::Tensor state_c = state_type_ok ? conv_state : conv_state.to(dtype); + // state_c and conv_state may be non-contiguous — that is intentional. + + // Weight: coerce to same dtype if needed (should match in practice) + at::Tensor w_c = + (weight.scalar_type() != dtype) + ? weight.to(dtype).contiguous() + : (weight.is_contiguous() ? weight : weight.contiguous()); + + // Bias stays float32 (small scalar, used only for fp32 accumulation) + at::Tensor bias_f32; + if (bias.has_value() && bias.value().defined()) + bias_f32 = bias.value().to(at::kFloat).contiguous(); + + int64_t batch = x_c.size(0); + int64_t dim = x_c.size(1); + int64_t seqlen = (x_c.dim() == 3) ? x_c.size(2) : 1; + int64_t width = w_c.size(1); + int64_t state_len = state_c.size(2); + + // Extract strides — works for contiguous AND non-contiguous (transposed) + // state. stride(0): between cache slots (e.g. num_slots × dim × width-1 in + // contiguous) stride(1): between conv channels (dim stride) stride(2): + // between state elements (=1 when contiguous, =dim when transposed) + int64_t stride_s_slot = state_c.stride(0); + int64_t stride_s_dim = state_c.stride(1); + int64_t stride_s_state = state_c.stride(2); + + at::Tensor out = x_c.clone(); // native dtype, no float32 alloc + + const int32_t* cache_idx_ptr = nullptr; + at::Tensor cache_idx_int; + if (conv_state_indices.has_value()) { + cache_idx_int = conv_state_indices.value().to(at::kInt).contiguous(); + cache_idx_ptr = cache_idx_int.data_ptr(); + } + + VLLM_DISPATCH_FLOATING_TYPES(dtype, "causal_conv1d_update", [&] { + mamba_cpu::causal_conv1d_update_kernel( + x_c.data_ptr(), state_c.data_ptr(), stride_s_slot, + stride_s_dim, stride_s_state, w_c.data_ptr(), + bias_f32.defined() ? bias_f32.data_ptr() : nullptr, + out.data_ptr(), cache_idx_ptr, + static_cast(pad_slot_id), batch, dim, seqlen, width, state_len, + do_silu); + }); + + // Write back only when a type-conversion copy was made. + // Layout-only non-contiguity is handled via strides above — no copy needed. + if (!state_type_ok) conv_state.copy_(state_c); + + return out; +} + +// --------------------------------------------------------------------------- +// selective_state_update +// --------------------------------------------------------------------------- +void selective_state_update_cpu_impl( + at::Tensor& state, // (nstates, nheads, dim, dstate) + const at::Tensor& x, // (N, nheads, dim) + const at::Tensor& dt, const at::Tensor& A, const at::Tensor& B, + const at::Tensor& C, const c10::optional& D, + const c10::optional& z, + const c10::optional& dt_bias, bool dt_softplus, + const c10::optional& state_batch_indices, + const c10::optional& dst_state_batch_indices, + int64_t null_block_id, at::Tensor& out, + const c10::optional& num_accepted_tokens, + const c10::optional& cu_seqlens) { + at::ScalarType state_type = state.scalar_type(); + at::ScalarType input_type = x.scalar_type(); + + // x, B, C must be contiguous and match input_type + auto ensure_input = [input_type](const at::Tensor& t) -> at::Tensor { + at::Tensor r = (t.scalar_type() != input_type) ? t.to(input_type) : t; + return r.is_contiguous() ? r : r.contiguous(); + }; + at::Tensor x_in = ensure_input(x); + at::Tensor B_in = ensure_input(B); + at::Tensor C_in = ensure_input(C); + at::Tensor z_in; + if (z.has_value() && z.value().defined()) z_in = ensure_input(z.value()); + + // A, D, dt_bias are float32 model parameters that arrive here as expanded + // tensors, e.g. A is (nheads, head_dim, dstate) with strides (1, 0, 0). + // We need just the scalar value per head as a (nheads,) 1-D array so that + // A_ptr[h] in the kernel correctly reads head h's value. + // + // Strategy: peel trailing expanded (stride=0) dims via .select(), which is + // a zero-copy view. For A: (nheads, head_dim, dstate) strides (1,0,0) + // → .select(2,0) → (nheads, head_dim) strides (1,0) + // → .select(1,0) → (nheads,) stride (1,) ← contiguous, free. + // No allocation, no type conversion (A is already float32). + auto to_per_head_1d_f32 = [](const at::Tensor& t) -> at::Tensor { + at::Tensor r = t; + // Peel trailing dimensions that are broadcast (stride=0 or size=1) + while (r.dim() > 1) r = r.select(r.dim() - 1, 0); + if (r.scalar_type() != at::kFloat) r = r.to(at::kFloat); + return r.is_contiguous() ? r : r.contiguous(); + }; + + at::Tensor A_f32 = to_per_head_1d_f32(A); // (nheads,) float32 + at::Tensor D_f32, dt_bias_f32; + if (D.has_value() && D.value().defined()) + D_f32 = to_per_head_1d_f32(D.value()); + if (dt_bias.has_value() && dt_bias.value().defined()) + dt_bias_f32 = to_per_head_1d_f32(dt_bias.value()); + + // dt: reduce (N, nheads, head_dim) expanded tensor → (N, nheads) BEFORE + // the type conversion so we convert head_dim x fewer elements. + at::Tensor dt_f32; + { + // If dt was expanded to (N, nheads, head_dim) with stride-0 in dim 2, + // take a zero-copy view of index 0 along that dim first. + at::Tensor t2 = (dt.dim() == 3) ? dt.select(2, 0) : dt; // (N, nheads) + at::Tensor t3 = (t2.scalar_type() != at::kFloat) ? t2.to(at::kFloat) : t2; + dt_f32 = t3.is_contiguous() ? t3 : t3.contiguous(); + } + + int64_t nheads = state.size(1); + int64_t dim = state.size(2); + int64_t dstate = state.size(3); + int64_t N = (cu_seqlens.has_value() && cu_seqlens.value().defined()) + ? cu_seqlens.value().size(0) - 1 + : x_in.size(0); + int64_t ngroups = B_in.size(1); + + // Strides + int64_t stride_state_n = state.stride(0); + int64_t stride_state_h = state.stride(1); + int64_t stride_state_d = state.stride(2); + int64_t stride_x_n = x_in.stride(0); + int64_t stride_x_h = x_in.stride(1); + int64_t stride_dt_n = dt_f32.stride(0); // dt is (N, nheads) + int64_t stride_BC_n = B_in.stride(0); + int64_t stride_BC_g = B_in.stride(1); + int64_t stride_out_n = out.stride(0); + int64_t stride_out_h = out.stride(1); + + // Optional index pointers + auto get_int32_ptr = + [](const c10::optional& opt) -> const int32_t* { + return (opt.has_value() && opt.value().defined()) + ? opt.value().data_ptr() + : nullptr; + }; + const int32_t* sbi_ptr = get_int32_ptr(state_batch_indices); + const int32_t* dsbi_ptr = get_int32_ptr(dst_state_batch_indices); + const int32_t* nat_ptr = get_int32_ptr(num_accepted_tokens); + const int32_t* csl_ptr = get_int32_ptr(cu_seqlens); + + // Dispatch on (state_t, input_t, out_t): write directly into `out` + // without any intermediate float32 buffer. + VLLM_DISPATCH_FLOATING_TYPES(state_type, "ssu_state", [&] { + using state_t = scalar_t; + VLLM_DISPATCH_FLOATING_TYPES(input_type, "ssu_input", [&] { + using input_t = scalar_t; + VLLM_DISPATCH_FLOATING_TYPES(out.scalar_type(), "ssu_out", [&] { + using out_t = scalar_t; + mamba_cpu::selective_state_update_kernel( + state.data_ptr(), stride_state_n, stride_state_h, + stride_state_d, x_in.data_ptr(), stride_x_n, stride_x_h, + dt_f32.data_ptr(), stride_dt_n, A_f32.data_ptr(), + B_in.data_ptr(), C_in.data_ptr(), stride_BC_n, + stride_BC_g, D_f32.defined() ? D_f32.data_ptr() : nullptr, + z_in.defined() ? z_in.data_ptr() : nullptr, + dt_bias_f32.defined() ? dt_bias_f32.data_ptr() : nullptr, + out.data_ptr(), stride_out_n, stride_out_h, sbi_ptr, + dsbi_ptr, static_cast(null_block_id), nat_ptr, csl_ptr, N, + nheads, ngroups, dim, dstate, dt_softplus); + }); + }); + }); +} + +// --------------------------------------------------------------------------- +// mamba_chunk_scan_fwd_cpu +// --------------------------------------------------------------------------- +void mamba_chunk_scan_fwd_cpu_impl( + at::Tensor& out, // [seqlen, nheads, headdim] — pre-allocated by caller + at::Tensor& + final_states, // [batch, nheads, headdim, dstate] float32 contiguous + const at::Tensor& x, // [seqlen, nheads, headdim] + const at::Tensor& + dt, // [seqlen, nheads] float32 (preprocessed: bias+softplus+clamp) + const at::Tensor& A, // [nheads] float32 + const at::Tensor& B, // [seqlen, ngroups, dstate] + const at::Tensor& C, // [seqlen, ngroups, dstate] + const c10::optional& D, // [nheads] float32 (optional) + const c10::optional& z, // [seqlen, nheads, headdim] (optional) + const at::Tensor& cu_seqlens // [batch+1] int32 +) { + const at::ScalarType input_type = x.scalar_type(); + + auto ensure_contig = [input_type](const at::Tensor& t) -> at::Tensor { + at::Tensor r = (t.scalar_type() != input_type) ? t.to(input_type) : t; + return r.is_contiguous() ? r : r.contiguous(); + }; + at::Tensor x_in = ensure_contig(x); + at::Tensor B_in = ensure_contig(B); + at::Tensor C_in = ensure_contig(C); + at::Tensor z_in; + if (z.has_value() && z.value().defined()) z_in = ensure_contig(z.value()); + + // A and D are float32 model parameters, potentially broadcast-expanded. + // Strip trailing broadcast dims to get a contiguous (nheads,) array. + auto to_per_head_f32 = [](const at::Tensor& t) -> at::Tensor { + at::Tensor r = t; + while (r.dim() > 1) r = r.select(r.dim() - 1, 0); + if (r.scalar_type() != at::kFloat) r = r.to(at::kFloat); + return r.is_contiguous() ? r : r.contiguous(); + }; + at::Tensor A_f32 = to_per_head_f32(A); + at::Tensor D_f32; + if (D.has_value() && D.value().defined()) D_f32 = to_per_head_f32(D.value()); + + // dt: [seqlen, nheads] float32 — caller has applied bias+softplus+clamp in + // Python. + at::Tensor dt_c = dt.is_contiguous() ? dt : dt.contiguous(); + if (dt_c.scalar_type() != at::kFloat) dt_c = dt_c.to(at::kFloat); + + at::Tensor cu_int = cu_seqlens.to(at::kInt).contiguous(); + + const int64_t batch = final_states.size(0); + const int64_t nheads = final_states.size(1); + const int64_t headdim = final_states.size(2); + const int64_t dstate = final_states.size(3); + const int64_t ngroups = B_in.size(1); + + TORCH_CHECK(final_states.is_contiguous(), + "mamba_chunk_scan_fwd_cpu: final_states must be contiguous"); + TORCH_CHECK(out.is_contiguous(), + "mamba_chunk_scan_fwd_cpu: out must be contiguous (writes via " + "raw data_ptr)"); + + VLLM_DISPATCH_FLOATING_TYPES(input_type, "mamba_chunk_scan_fwd_cpu", [&] { + mamba_cpu::mamba_chunk_scan_fwd_kernel( + final_states.data_ptr(), x_in.data_ptr(), + dt_c.data_ptr(), A_f32.data_ptr(), + B_in.data_ptr(), C_in.data_ptr(), + D_f32.defined() ? D_f32.data_ptr() : nullptr, + z_in.defined() ? z_in.data_ptr() : nullptr, + out.data_ptr(), cu_int.data_ptr(), batch, nheads, + ngroups, headdim, dstate); + }); +} diff --git a/csrc/cpu/mamba_kernels.hpp b/csrc/cpu/mamba_kernels.hpp new file mode 100644 index 00000000000..722dad97e3c --- /dev/null +++ b/csrc/cpu/mamba_kernels.hpp @@ -0,0 +1,382 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright contributors to the vLLM project +// +// Fused CPU vector kernels for Mamba decode-step hotspots: +// - causal_conv1d_update (depthwise 1-D conv state roll + compute) +// - selective_state_update (SSM recurrence, single-step) + +#pragma once + +#include "cpu_types.hpp" +#include +#include +#include +#include + +namespace mamba_cpu { + +// --------------------------------------------------------------------------- +// causal_conv1d_update — templated for native BF16/FP32 +// +// state_ptr may point to a NON-CONTIGUOUS paged KV cache tensor. +// Explicit strides are passed so the kernel writes directly into the +// correct memory locations without making a contiguous copy of the full +// paged tensor (which was the source of the 34-41% direct_copy_kernel). +// +// stride_s_slot = state.stride(0) — between cache slots +// stride_s_dim = state.stride(1) — between conv_dim channels +// stride_s_state = state.stride(2) — between state elements +// +// When stride_s_state == 1 (contiguous), the memmove fast path is used. +// --------------------------------------------------------------------------- +template +inline void causal_conv1d_update_kernel( + const scalar_t* __restrict__ x_ptr, scalar_t* __restrict__ state_ptr, + int64_t stride_s_slot, int64_t stride_s_dim, int64_t stride_s_state, + const scalar_t* __restrict__ weight_ptr, const float* __restrict__ bias_ptr, + scalar_t* __restrict__ out_ptr, const int32_t* __restrict__ cache_idxs, + int32_t pad_slot_id, int64_t batch, int64_t dim, int64_t seqlen, + int64_t width, int64_t state_len, bool do_silu) { +#pragma omp parallel for + for (int64_t b = 0; b < batch; ++b) { + int64_t cache_idx = (cache_idxs != nullptr) ? cache_idxs[b] : b; + if (cache_idx == pad_slot_id) continue; + + for (int64_t t = 0; t < seqlen; ++t) { + const scalar_t* x_b = x_ptr + (b * dim * seqlen + t); + scalar_t* out_b = out_ptr + (b * dim * seqlen + t); + // Base of this slot in the (possibly non-contiguous) paged state + scalar_t* s_base = state_ptr + cache_idx * stride_s_slot; + + for (int64_t d = 0; d < dim; ++d) { + float x_val = static_cast(x_b[d * seqlen]); + scalar_t* sd = s_base + d * stride_s_dim; // start of this dim's state + const scalar_t* w = weight_ptr + d * width; + + // Accumulate in float32 for precision + float acc = (bias_ptr != nullptr) ? bias_ptr[d] : 0.0f; + for (int64_t k = 0; k < state_len; ++k) { + acc += static_cast(w[k]) * + static_cast(sd[k * stride_s_state]); + } + acc += static_cast(w[state_len]) * x_val; + + // Shift state left and append new input. + // Use memmove when contiguous (stride==1); element loop otherwise. + if (stride_s_state == 1) { + if (state_len > 1) + std::memmove(sd, sd + 1, (state_len - 1) * sizeof(scalar_t)); + if (state_len > 0) sd[state_len - 1] = static_cast(x_val); + } else { + for (int64_t k = 0; k < state_len - 1; ++k) + sd[k * stride_s_state] = sd[(k + 1) * stride_s_state]; + if (state_len > 0) + sd[(state_len - 1) * stride_s_state] = static_cast(x_val); + } + + if (do_silu) { + float sigmoid = (acc >= 0) ? 1.0f / (1.0f + std::exp(-acc)) + : std::exp(acc) / (1.0f + std::exp(acc)); + acc *= sigmoid; + } + out_b[d * seqlen] = static_cast(acc); + } + } + } +} + +// --------------------------------------------------------------------------- +// selective_state_update +// +// Template parameters: +// state_t - dtype of ssm_state cache (typically BFloat16) +// input_t - dtype of x, B, C (typically BFloat16) +// out_t - dtype of output tensor (typically BFloat16) +// Write directly — no float32 intermediate buffer needed. +// +// A, D, dt_bias are accepted as const float* (they are always float32 +// model parameters in Mamba2). This eliminates the per-call float32→BF16 +// conversion and the .contiguous() materialisation of the broadcast-expand. +// +// dt is accepted as a (N, nheads) scalar-per-head tensor, not as the +// (N, nheads, head_dim) expansion, so no .contiguous() copy is needed. +// --------------------------------------------------------------------------- +template +inline void selective_state_update_kernel( + state_t* __restrict__ state_ptr, int64_t stride_state_n, + int64_t stride_state_h, int64_t stride_state_d, + const input_t* __restrict__ x_ptr, int64_t stride_x_n, int64_t stride_x_h, + // dt: (N, nheads) — scalar per head, NOT expanded to head_dim + const float* __restrict__ dt_ptr, int64_t stride_dt_n, + // A: (nheads,) float32 — scalar per head + const float* __restrict__ A_ptr, const input_t* __restrict__ B_ptr, + const input_t* __restrict__ C_ptr, int64_t stride_BC_n, int64_t stride_BC_g, + // D: (nheads,) float32 — scalar per head (nullptr if not used) + const float* __restrict__ D_ptr, + // z: same shape as x (optional) + const input_t* __restrict__ z_ptr, + // dt_bias: (nheads,) float32 — scalar per head (nullptr if not used) + const float* __restrict__ dt_bias_ptr, out_t* __restrict__ out_ptr, + int64_t stride_out_n, int64_t stride_out_h, + const int32_t* __restrict__ state_batch_indices, + const int32_t* __restrict__ dst_state_batch_indices, int32_t null_block_id, + const int32_t* __restrict__ num_accepted_tokens, + const int32_t* __restrict__ cu_seqlens, int64_t N, int64_t nheads, + int64_t ngroups, int64_t dim, int64_t dstate, bool dt_softplus) { + using state_vec_t = vec_op::vec_t; + using input_vec_t = vec_op::vec_t; + constexpr int VEC_ELEM_NUM = 8; + + int64_t nheads_per_group = nheads / ngroups; + + for (int64_t seq_idx = 0; seq_idx < N; ++seq_idx) { + int64_t bos, seq_len; + if (cu_seqlens != nullptr) { + bos = cu_seqlens[seq_idx]; + seq_len = cu_seqlens[seq_idx + 1] - bos; + } else { + bos = seq_idx; + seq_len = 1; + } + + int64_t state_read_idx = (state_batch_indices != nullptr) + ? state_batch_indices[seq_idx] + : seq_idx; + if (state_read_idx == null_block_id) continue; + + int64_t state_write_idx = (num_accepted_tokens == nullptr) + ? ((dst_state_batch_indices != nullptr) + ? dst_state_batch_indices[seq_idx] + : state_read_idx) + : -1; + + state_t* s = state_ptr + state_read_idx * stride_state_n; + + for (int64_t t = 0; t < seq_len; ++t) { + int64_t token_idx = bos + t; + const input_t* x_tok = x_ptr + token_idx * stride_x_n; + // dt: (N, nheads) — one float per head per token + const float* dt_tok = dt_ptr + token_idx * stride_dt_n; + const input_t* B_tok = B_ptr + token_idx * stride_BC_n; + const input_t* C_tok = C_ptr + token_idx * stride_BC_n; + out_t* out_tok = out_ptr + token_idx * stride_out_n; + +#pragma omp parallel for + for (int64_t h = 0; h < nheads; ++h) { + int64_t g = h / nheads_per_group; + const input_t* x_h = x_tok + h * stride_x_h; + const input_t* B_g = B_tok + g * stride_BC_g; + const input_t* C_g = C_tok + g * stride_BC_g; + out_t* out_h = out_tok + h * stride_out_h; + state_t* s_h = s + h * stride_state_h; + + // Read scalars-per-head (A, dt, dt_bias, D) — no per-dim indexing + float dt_val = dt_tok[h]; + if (dt_bias_ptr != nullptr) dt_val += dt_bias_ptr[h]; + if (dt_softplus) { + dt_val = (dt_val <= 20.0f) ? std::log1p(std::exp(dt_val)) : dt_val; + } + const float A_val = A_ptr[h]; // scalar: same for all dim, dstate + const float D_val = (D_ptr != nullptr) ? D_ptr[h] : 0.0f; + + const input_t* z_h = + (z_ptr != nullptr) ? z_ptr + token_idx * stride_x_n + h * stride_x_h + : nullptr; + + vec_op::FP32Vec8 dt_vec(dt_val); + // dA = exp(A * dt): A and dt are SCALARS per head, so compute once + // and broadcast. This saves 7 redundant std::exp() calls that + // FP32Vec8::exp() would otherwise make on the broadcast vector. + const float dA_scalar = std::exp(A_val * dt_val); + vec_op::FP32Vec8 dA(dA_scalar); // broadcast + + for (int64_t d = 0; d < dim; ++d) { + float x_val = static_cast(x_h[d]); + + vec_op::FP32Vec8 out_vec(0.0f); + state_t* s_hd = s_h + d * stride_state_d; + const input_t* B_g_base = B_g; + const input_t* C_g_base = C_g; + + vec_op::FP32Vec8 x_vec(x_val); + // dBx = B * x * dt — same dA for all dstate (A is scalar) + // s_new = s * dA + B * x * dt + + int64_t n = 0; + for (; n <= dstate - VEC_ELEM_NUM; n += VEC_ELEM_NUM) { + vec_op::FP32Vec8 B_v((input_vec_t(B_g_base + n))); + vec_op::FP32Vec8 C_v((input_vec_t(C_g_base + n))); + vec_op::FP32Vec8 s_v((state_vec_t(s_hd + n))); + + vec_op::FP32Vec8 dBx = B_v * x_vec * dt_vec; + vec_op::FP32Vec8 s_new = s_v * dA + dBx; + + state_vec_t(s_new).save(s_hd + n); + out_vec = out_vec + s_new * C_v; + } + + float out_val = out_vec.reduce_sum(); + for (; n < dstate; ++n) { + // Reuse dA_scalar computed once per head — no exp() re-call + float dBx = static_cast(B_g[n]) * x_val * dt_val; + float s_new = static_cast(s_hd[n]) * dA_scalar + dBx; + s_hd[n] = static_cast(s_new); + out_val += s_new * static_cast(C_g[n]); + } + + if (D_ptr != nullptr) out_val += x_val * D_val; + if (z_h != nullptr) { + float z_val = static_cast(z_h[d]); + float sigmoid = (z_val >= 0) + ? 1.0f / (1.0f + std::exp(-z_val)) + : std::exp(z_val) / (1.0f + std::exp(z_val)); + out_val *= z_val * sigmoid; + } + out_h[d] = static_cast(out_val); + } + } + + if (num_accepted_tokens != nullptr && + dst_state_batch_indices != nullptr) { + int64_t token_dst_idx = dst_state_batch_indices[seq_idx * seq_len + t]; + if (token_dst_idx != null_block_id && token_dst_idx != state_read_idx) { + state_t* dst_s = state_ptr + token_dst_idx * stride_state_n; + std::memmove(dst_s, s, nheads * stride_state_h * sizeof(state_t)); + } + } + } + + if (num_accepted_tokens == nullptr && state_write_idx != null_block_id && + state_write_idx != state_read_idx) { + state_t* dst_s = state_ptr + state_write_idx * stride_state_n; + std::memmove(dst_s, s, nheads * stride_state_h * sizeof(state_t)); + } + } +} + +// --------------------------------------------------------------------------- +// mamba_chunk_scan_fwd +// +// Prefill SSM recurrence for Mamba2 / SSD models. +// +// Key difference from selective_state_update_kernel (decode path): +// - #pragma omp parallel for collapse(2) is OUTSIDE the time loop. +// Each thread owns a (batch, head) slice and runs the entire token +// sequence without any per-token OpenMP synchronisation overhead. +// For seqlen=256, this eliminates 256 thread-barrier launches per batch. +// +// `dt` arrives already processed (float32, after bias + softplus + clamp) +// to keep this kernel simple. Preprocessing is done in the Python wrapper. +// +// `states_ptr` points to the [batch, nheads, headdim, dstate] float32 output +// tensor, pre-initialised by the caller (zero or from initial_states). +// Each (b, h) slice is private to exactly one thread via collapse(2), so +// there are no write conflicts. +// +// D is treated as a scalar per head ([nheads] float32). +// --------------------------------------------------------------------------- +template +inline void mamba_chunk_scan_fwd_kernel( + float* __restrict__ states_ptr, // [batch, nheads, headdim, dstate] f32 + const input_t* __restrict__ x_ptr, // [seqlen, nheads, headdim] + const float* __restrict__ dt_ptr, // [seqlen, nheads] f32 (preprocessed) + const float* __restrict__ A_ptr, // [nheads] f32 + const input_t* __restrict__ B_ptr, // [seqlen, ngroups, dstate] + const input_t* __restrict__ C_ptr, // [seqlen, ngroups, dstate] + const float* __restrict__ D_ptr, // [nheads] f32 (nullable) + const input_t* __restrict__ z_ptr, // [seqlen, nheads, headdim] (nullable) + input_t* __restrict__ out_ptr, // [seqlen, nheads, headdim] + const int32_t* __restrict__ cu_seqlens, // [batch+1] int32 + int64_t batch, int64_t nheads, int64_t ngroups, int64_t headdim, + int64_t dstate) { + using input_vec_t = vec_op::vec_t; + constexpr int VEC_ELEM_NUM = 8; + + const int64_t nheads_per_group = nheads / ngroups; + // states layout: [batch, nheads, headdim, dstate] contiguous (caller + // guarantee) + const int64_t stride_s_b = nheads * headdim * dstate; + const int64_t stride_s_h = headdim * dstate; + // stride_s_d = dstate, stride_s_n = 1 + +#pragma omp parallel for collapse(2) schedule(static) + for (int64_t b = 0; b < batch; ++b) { + for (int64_t h = 0; h < nheads; ++h) { + const int64_t seq_start = cu_seqlens[b]; + const int64_t seq_end = cu_seqlens[b + 1]; + const int64_t g = h / nheads_per_group; + + const float A_val = A_ptr[h]; + const float D_val = (D_ptr != nullptr) ? D_ptr[h] : 0.0f; + + // Working state slice: states[b, h, :, :] — float32, headdim * dstate. + // Fits in L1/L2 for typical dims (e.g. 64*128*4 = 32 KB). + float* s_bh = states_ptr + b * stride_s_b + h * stride_s_h; + + for (int64_t t = seq_start; t < seq_end; ++t) { + const input_t* x_h = x_ptr + t * nheads * headdim + h * headdim; + const float* dt_h = dt_ptr + t * nheads + h; + const input_t* B_g = B_ptr + t * ngroups * dstate + g * dstate; + const input_t* C_g = C_ptr + t * ngroups * dstate + g * dstate; + const input_t* z_h = (z_ptr != nullptr) + ? z_ptr + t * nheads * headdim + h * headdim + : nullptr; + input_t* out_h = out_ptr + t * nheads * headdim + h * headdim; + + const float dt_val = *dt_h; + const float dA_val = std::exp(A_val * dt_val); + const vec_op::FP32Vec8 dA_vec(dA_val); // broadcast scalar + const vec_op::FP32Vec8 dt_vec(dt_val); + + for (int64_t d = 0; d < headdim; ++d) { + const float x_val = static_cast(x_h[d]); + float* s_bhd = s_bh + d * dstate; // [dstate] contiguous float32 + + // Vectorised SSM update + readout over dstate: + // s_new = s * dA + x * dt * B + // y += s_new * C + int64_t n = 0; + vec_op::FP32Vec8 y_vec(0.0f); + const vec_op::FP32Vec8 x_vec(x_val); + + for (; n <= dstate - VEC_ELEM_NUM; n += VEC_ELEM_NUM) { + const vec_op::FP32Vec8 B_v((input_vec_t(B_g + n))); + const vec_op::FP32Vec8 C_v((input_vec_t(C_g + n))); + const vec_op::FP32Vec8 s_v(s_bhd + n); + + const vec_op::FP32Vec8 s_new = s_v * dA_vec + x_vec * dt_vec * B_v; + s_new.save(s_bhd + n); + y_vec = y_vec + s_new * C_v; + } + + float y_val = y_vec.reduce_sum(); + + // Scalar tail for remaining dstate elements + for (; n < dstate; ++n) { + const float B_n = static_cast(B_g[n]); + const float C_n = static_cast(C_g[n]); + const float s_new = s_bhd[n] * dA_val + x_val * dt_val * B_n; + s_bhd[n] = s_new; + y_val += s_new * C_n; + } + + // D skip connection (scalar per head) + if (D_ptr != nullptr) y_val += x_val * D_val; + + // z gating: out = y * z * sigmoid(z) (SiLU) + if (z_h != nullptr) { + const float z_val = static_cast(z_h[d]); + const float sigmoid = + (z_val >= 0.0f) ? 1.0f / (1.0f + std::exp(-z_val)) + : std::exp(z_val) / (1.0f + std::exp(z_val)); + y_val *= z_val * sigmoid; + } + + out_h[d] = static_cast(y_val); + } + } + } + } +} + +} // namespace mamba_cpu diff --git a/csrc/cpu/torch_bindings.cpp b/csrc/cpu/torch_bindings.cpp index 88a593725ee..8b7d924dace 100644 --- a/csrc/cpu/torch_bindings.cpp +++ b/csrc/cpu/torch_bindings.cpp @@ -213,6 +213,32 @@ void compute_slot_mapping_kernel_impl(const torch::Tensor query_start_loc, torch::Tensor slot_mapping, const int64_t block_size); +at::Tensor causal_conv1d_update_cpu_impl( + at::Tensor& x, at::Tensor& conv_state, const at::Tensor& weight, + const c10::optional& bias, + const c10::optional& activation, + const c10::optional& conv_state_indices, + const c10::optional& query_start_loc, int64_t pad_slot_id); + +void selective_state_update_cpu_impl( + at::Tensor& state, const at::Tensor& x, const at::Tensor& dt, + const at::Tensor& A, const at::Tensor& B, const at::Tensor& C, + const c10::optional& D, const c10::optional& z, + const c10::optional& dt_bias, bool dt_softplus, + const c10::optional& state_batch_indices, + const c10::optional& dst_state_batch_indices, + int64_t null_block_id, at::Tensor& out, + const c10::optional& num_accepted_tokens, + const c10::optional& cu_seqlens); + +void mamba_chunk_scan_fwd_cpu_impl(at::Tensor& out, at::Tensor& final_states, + const at::Tensor& x, const at::Tensor& dt, + const at::Tensor& A, const at::Tensor& B, + const at::Tensor& C, + const c10::optional& D, + const c10::optional& z, + const at::Tensor& cu_seqlens); + void init_cpu_memory_env(std::vector node_ids); namespace cpu_utils { @@ -595,6 +621,30 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) { "block_size) -> ()", &compute_slot_mapping_kernel_impl); + // Mamba CPU kernels + ops.def( + "causal_conv1d_update_cpu_vec(" + "Tensor(a0!) x, Tensor(a1!) conv_state, Tensor weight, " + "Tensor? bias, str? activation, Tensor? conv_state_indices, " + "Tensor? query_start_loc, SymInt pad_slot_id) -> Tensor", + &causal_conv1d_update_cpu_impl); + + ops.def( + "selective_state_update_cpu(" + "Tensor(a0!) state, Tensor x, Tensor dt, Tensor A, Tensor B, Tensor C, " + "Tensor? D, Tensor? z, Tensor? dt_bias, bool dt_softplus, " + "Tensor? state_batch_indices, Tensor? dst_state_batch_indices, " + "SymInt null_block_id, Tensor(a13!) out, " + "Tensor? num_accepted_tokens, Tensor? cu_seqlens) -> ()", + &selective_state_update_cpu_impl); + + ops.def( + "mamba_chunk_scan_fwd_cpu(" + "Tensor(a0!) out, Tensor(a1!) final_states, " + "Tensor x, Tensor dt, Tensor A, Tensor B, Tensor C, " + "Tensor? D, Tensor? z, Tensor cu_seqlens) -> ()", + &mamba_chunk_scan_fwd_cpu_impl); + ops.def("init_cpu_memory_env(SymInt[] node_ids) -> ()", &init_cpu_memory_env); // Speculative decoding kernels diff --git a/tests/kernels/mamba/cpu/test_cpu_gdn_ops.py b/tests/kernels/mamba/cpu/test_cpu_gdn_ops.py index bd30bc4f1ce..083e9ff22aa 100644 --- a/tests/kernels/mamba/cpu/test_cpu_gdn_ops.py +++ b/tests/kernels/mamba/cpu/test_cpu_gdn_ops.py @@ -425,7 +425,7 @@ def test_causal_conv1d_torch_two_call_split(total_tokens: int, split: int) -> No match the single-call result. """ from vllm.model_executor.layers.mamba.ops.cpu.causal_conv1d import ( - causal_conv1d_torch, + causal_conv1d_fn_cpu as causal_conv1d_torch, ) x, weight, bias = _conv_inputs(total_tokens) diff --git a/tests/kernels/mamba/test_causal_conv1d.py b/tests/kernels/mamba/test_causal_conv1d.py index c6554f131fe..d2a2981edf2 100644 --- a/tests/kernels/mamba/test_causal_conv1d.py +++ b/tests/kernels/mamba/test_causal_conv1d.py @@ -18,8 +18,12 @@ from vllm.v1.attention.backends.utils import NULL_BLOCK_ID DEVICE = current_platform.device_type pytestmark = pytest.mark.skipif( - not (current_platform.is_cuda_alike() or current_platform.is_xpu()), - reason="causal_conv1d Triton kernels require CUDA-alike or XPU", + not ( + current_platform.is_cuda_alike() + or current_platform.is_xpu() + or current_platform.is_cpu() + ), + reason="causal_conv1d Triton kernels require CUDA-alike, XPU, or CPU", ) @@ -284,7 +288,8 @@ def test_causal_conv1d_varlen( batch, with_padding, dim, seqlen, width, has_bias, silu_activation, itype ): device = DEVICE - torch.accelerator.empty_cache() + if not current_platform.is_cpu(): + torch.accelerator.empty_cache() rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (3e-3, 5e-3) if itype == torch.bfloat16: rtol, atol = 1e-2, 5e-2 diff --git a/tests/kernels/mamba/test_mamba_ssm.py b/tests/kernels/mamba/test_mamba_ssm.py index 81b57d4b16b..f8af3db56be 100644 --- a/tests/kernels/mamba/test_mamba_ssm.py +++ b/tests/kernels/mamba/test_mamba_ssm.py @@ -20,8 +20,12 @@ from vllm.v1.attention.backends.utils import NULL_BLOCK_ID DEVICE = current_platform.device_type pytestmark = pytest.mark.skipif( - not (current_platform.is_cuda_alike() or current_platform.is_xpu()), - reason="mamba_ssm kernels require CUDA-alike or XPU", + not ( + current_platform.is_cuda_alike() + or current_platform.is_xpu() + or current_platform.is_cpu() + ), + reason="mamba_ssm kernels require CUDA-alike, XPU, or CPU", ) # selective_scan_fn is backed by the CUDA-only `ops.selective_scan_fwd` C++ op, @@ -342,6 +346,13 @@ def test_selective_scan( @pytest.mark.parametrize("has_z", [False, True]) @pytest.mark.parametrize("dstate", [16, 64]) @pytest.mark.parametrize("dim", [2048, 2048 + 16, 4096]) +@pytest.mark.skipif( + current_platform.is_cpu(), + reason=( + "CPU kernel for selective_state_update only supports " + "Mamba 2 (scalar A/dt), not Mamba 1." + ), +) def test_selective_state_update(dim, dstate, has_z, itype): device = DEVICE rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2) @@ -436,6 +447,13 @@ def test_selective_state_update_stochastic_rounding(dim, dstate, has_z, philox_r @pytest.mark.parametrize("dstate", [16, 64]) @pytest.mark.parametrize("dim", [2048, 2048 + 16, 4096]) @pytest.mark.parametrize("max_seq_len", [1, 2, 4]) +@pytest.mark.skipif( + current_platform.is_cpu(), + reason=( + "CPU kernel for selective_state_update only supports " + "Mamba 2 (scalar A/dt), not Mamba 1." + ), +) def test_selective_state_update_varlen(dim, dstate, has_z, itype, max_seq_len): device = DEVICE rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2) @@ -697,6 +715,13 @@ def test_selective_scan_varlen( @pytest.mark.parametrize("dim", [2048, 2048 + 16, 4096]) # tests correctness in case subset of the sequences are padded @pytest.mark.parametrize("with_padding", [True, False]) +@pytest.mark.skipif( + current_platform.is_cpu(), + reason=( + "CPU kernel for selective_state_update only supports " + "Mamba 2 (scalar A/dt), not Mamba 1." + ), +) def test_selective_state_update_with_batch_indices( with_padding, dim, dstate, has_z, itype ): @@ -789,6 +814,13 @@ def test_selective_state_update_with_batch_indices( @pytest.mark.parametrize("ngroups", [1, 4]) @pytest.mark.parametrize("dstate", [16, 64]) @pytest.mark.parametrize("dim", [2048, 4096]) +@pytest.mark.skipif( + current_platform.is_cpu(), + reason=( + "CPU kernel for selective_state_update only supports " + "Mamba 2 (scalar A/dt), not Mamba 1." + ), +) def test_selective_state_update_with_heads_with_batch_indices( dim, dstate, ngroups, has_z, tie_hdim, itype ): @@ -862,6 +894,13 @@ def test_selective_state_update_with_heads_with_batch_indices( @pytest.mark.parametrize("dstate", [16, 64]) @pytest.mark.parametrize("dim", [2048, 4096]) @pytest.mark.parametrize("max_seq_len", [2, 4]) +@pytest.mark.skipif( + current_platform.is_cpu(), + reason=( + "CPU kernel for selective_state_update only supports " + "Mamba 2 (scalar A/dt), not Mamba 1." + ), +) def test_selective_state_update_with_num_accepted_tokens( dim, dstate, has_z, itype, max_seq_len ): @@ -988,6 +1027,13 @@ def test_selective_state_update_with_num_accepted_tokens( @pytest.mark.parametrize("dstate", [16, 64]) @pytest.mark.parametrize("dim", [2048, 4096]) @pytest.mark.parametrize("max_seq_len", [2, 4]) +@pytest.mark.skipif( + current_platform.is_cpu(), + reason=( + "CPU kernel for selective_state_update only supports " + "Mamba 2 (scalar A/dt), not Mamba 1." + ), +) def test_selective_state_update_varlen_with_num_accepted( dim, dstate, has_z, itype, max_seq_len ): diff --git a/vllm/_custom_ops.py b/vllm/_custom_ops.py index 599cac0ed6f..588e0fc654f 100644 --- a/vllm/_custom_ops.py +++ b/vllm/_custom_ops.py @@ -2070,6 +2070,93 @@ def selective_scan_fwd( ) +def causal_conv1d_update_cpu_vec( + x: torch.Tensor, + conv_state: torch.Tensor, + weight: torch.Tensor, + bias: torch.Tensor | None = None, + activation: str | None = None, + conv_state_indices: torch.Tensor | None = None, + query_start_loc: torch.Tensor | None = None, + pad_slot_id: int = 0, +) -> torch.Tensor: + return torch.ops._C.causal_conv1d_update_cpu_vec( + x, + conv_state, + weight, + bias, + activation, + conv_state_indices, + query_start_loc, + pad_slot_id, + ) + + +def selective_state_update_cpu( + state: torch.Tensor, + x: torch.Tensor, + dt: torch.Tensor, + A: torch.Tensor, + B: torch.Tensor, + C: torch.Tensor, + D: torch.Tensor | None, + z: torch.Tensor | None, + dt_bias: torch.Tensor | None, + dt_softplus: bool, + state_batch_indices: torch.Tensor | None, + dst_state_batch_indices: torch.Tensor | None, + null_block_id: int, + out: torch.Tensor, + num_accepted_tokens: torch.Tensor | None, + cu_seqlens: torch.Tensor | None, +): + torch.ops._C.selective_state_update_cpu( + state, + x, + dt, + A, + B, + C, + D, + z, + dt_bias, + dt_softplus, + state_batch_indices, + dst_state_batch_indices, + null_block_id, + out, + num_accepted_tokens, + cu_seqlens, + ) + + +def mamba_chunk_scan_fwd_cpu( + out: torch.Tensor, + final_states: torch.Tensor, + x: torch.Tensor, + dt: torch.Tensor, + A: torch.Tensor, + B: torch.Tensor, + C: torch.Tensor, + D: torch.Tensor | None, + z: torch.Tensor | None, + cu_seqlens: torch.Tensor, +) -> None: + """Prefill SSM scan kernel. out and final_states are written in-place.""" + torch.ops._C.mamba_chunk_scan_fwd_cpu( + out, + final_states, + x, + dt, + A, + B, + C, + D, + z, + cu_seqlens, + ) + + # ROCm skinny gemms def LLMM1(a: torch.Tensor, b: torch.Tensor, rows_per_block: int) -> torch.Tensor: return torch.ops._rocm_C.LLMM1(a, b, rows_per_block) diff --git a/vllm/config/mamba.py b/vllm/config/mamba.py index 996478c3676..a68842f4bfe 100644 --- a/vllm/config/mamba.py +++ b/vllm/config/mamba.py @@ -27,6 +27,7 @@ class MambaBackendEnum(Enum, metaclass=_MambaBackendEnumMeta): TRITON = "triton" FLASHINFER = "flashinfer" + CPU = "cpu" @config diff --git a/vllm/model_executor/layers/mamba/ops/causal_conv1d.py b/vllm/model_executor/layers/mamba/ops/causal_conv1d.py index f7c237ca2db..15f08f26550 100644 --- a/vllm/model_executor/layers/mamba/ops/causal_conv1d.py +++ b/vllm/model_executor/layers/mamba/ops/causal_conv1d.py @@ -1237,3 +1237,15 @@ def causal_conv1d_update( if unsqueeze: out = out.squeeze(-1) return out.to(original_x_dtype) + + +from vllm.platforms import current_platform # noqa: E402 + +if current_platform.is_cpu(): + from vllm.model_executor.layers.mamba.ops.cpu.causal_conv1d import ( + causal_conv1d_fn_cpu, + causal_conv1d_update_cpu, + ) + + causal_conv1d_fn = causal_conv1d_fn_cpu # type: ignore + causal_conv1d_update = causal_conv1d_update_cpu # type: ignore diff --git a/vllm/model_executor/layers/mamba/ops/cpu/causal_conv1d.py b/vllm/model_executor/layers/mamba/ops/cpu/causal_conv1d.py index b047ca6d616..c552d93445e 100644 --- a/vllm/model_executor/layers/mamba/ops/cpu/causal_conv1d.py +++ b/vllm/model_executor/layers/mamba/ops/cpu/causal_conv1d.py @@ -6,18 +6,31 @@ from __future__ import annotations import torch import torch.nn.functional as F +from vllm._custom_ops import causal_conv1d_update_cpu_vec +from vllm.v1.attention.backends.utils import NULL_BLOCK_ID, PAD_SLOT_ID -# for prefill -def causal_conv1d_torch( + +def causal_conv1d_fn_cpu( x: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor | None, conv_states: torch.Tensor, query_start_loc: torch.Tensor, - cache_indices: torch.Tensor, - has_initial_state: torch.Tensor, + cache_indices: torch.Tensor | None = None, + has_initial_state: torch.Tensor | None = None, activation: str | None = "silu", + pad_slot_id: int = PAD_SLOT_ID, + **kwargs, ) -> torch.Tensor: + """CPU implementation for causal_conv1d_fwd.""" + if isinstance(activation, bool) and activation: + activation = "silu" + elif isinstance(activation, bool): + activation = None + + original_x_dtype = x.dtype + x = x.to(conv_states.dtype) + out = torch.empty_like(x) state_len = weight.shape[1] - 1 assert activation in {None, "silu", "swish"} @@ -27,11 +40,21 @@ def causal_conv1d_torch( for idx in range(query_start_loc.shape[0] - 1) ] weight = weight.unsqueeze(1) + for seq_idx, (bos, eos) in enumerate(seq_begin_end_idx): - slot = int(cache_indices[seq_idx].item()) + if bos == eos: + continue + + slot = ( + int(cache_indices[seq_idx].item()) if cache_indices is not None else seq_idx + ) + + if slot == pad_slot_id: + continue seq_x = x[:, bos:eos].unsqueeze(0) - if bool(has_initial_state[seq_idx].item()): + + if has_initial_state is not None and bool(has_initial_state[seq_idx].item()): initial_state = conv_states[slot, :, :state_len].unsqueeze(0) else: initial_state = torch.zeros( @@ -51,16 +74,48 @@ def causal_conv1d_torch( groups=weight.shape[0], ) seq_out = seq_out[..., -seq_x.shape[-1] :].to(dtype=x.dtype) + if activation in ("silu", "swish"): seq_out = F.silu(seq_out) out[:, bos:eos] = seq_out.squeeze(0) conv_states[slot, :, :state_len].copy_(conv_input[..., -state_len:].squeeze(0)) - return out + return out.to(original_x_dtype) + + +def causal_conv1d_update_cpu( + x: torch.Tensor, + conv_state: torch.Tensor, + weight: torch.Tensor, + bias: torch.Tensor | None = None, + activation: bool | str | None = None, + conv_state_indices: torch.Tensor | None = None, + query_start_loc: torch.Tensor | None = None, + pad_slot_id: int | None = None, + **kwargs, +) -> torch.Tensor: + """CPU implementation for causal_conv1d_update.""" + if isinstance(activation, bool): + activation = "silu" if activation else None + + if pad_slot_id is None: + pad_slot_id = kwargs.get("null_block_id", NULL_BLOCK_ID) + if pad_slot_id is None: + pad_slot_id = NULL_BLOCK_ID + + return causal_conv1d_update_cpu_vec( + x, + conv_state, + weight, + bias, + activation, + conv_state_indices, + query_start_loc, + pad_slot_id, + ) -# for decode def causal_conv1d_update_torch( x: torch.Tensor, conv_state: torch.Tensor, @@ -68,6 +123,11 @@ def causal_conv1d_update_torch( bias: torch.Tensor | None = None, activation: str | None = None, ) -> torch.Tensor: + """ + Pure PyTorch fallback for causal_conv1d_update. + Currently used as a fallback for Arm (aarch64) to leverage + oneDNN/ACL F.conv1d kernels for batched decoding. + """ assert activation in {None, "silu", "swish"} _, dim, seq_len = x.shape diff --git a/vllm/model_executor/layers/mamba/ops/cpu/gdn_attention.py b/vllm/model_executor/layers/mamba/ops/cpu/gdn_attention.py index 9f3aa12c8d1..f93d2e18db6 100644 --- a/vllm/model_executor/layers/mamba/ops/cpu/gdn_attention.py +++ b/vllm/model_executor/layers/mamba/ops/cpu/gdn_attention.py @@ -10,9 +10,13 @@ import vllm._custom_ops as ops from vllm.forward_context import ForwardContext, get_forward_context from vllm.model_executor.layers.mamba.mamba_utils import is_conv_state_dim_first from vllm.model_executor.layers.mamba.ops.cpu.causal_conv1d import ( - causal_conv1d_torch, + causal_conv1d_fn_cpu as causal_conv1d_torch, +) +from vllm.model_executor.layers.mamba.ops.cpu.causal_conv1d import ( + causal_conv1d_update_cpu, causal_conv1d_update_torch, ) +from vllm.platforms import CpuArchEnum, current_platform from vllm.utils.torch_utils import ( LayerNameType, _resolve_layer_name, @@ -140,21 +144,30 @@ def _cpu_gdn_attention_nonspec( conv_states=conv_state, weight=layer.conv1d.weight, bias=layer.conv1d.bias, - silu_activation=layer.activation == "silu", + silu_activation=(layer.activation == "silu"), conv_state_indices=decode_state_indices, is_vnni=True, ) else: - decode_conv_state = conv_state[decode_state_indices].contiguous() - decode_mixed_qkv = causal_conv1d_update_torch( - # [B, dim] -> [B, dim, 1] - x=decode_mixed_qkv.unsqueeze(-1), - conv_state=decode_conv_state, - weight=conv_weights, - bias=layer.conv1d.bias, - activation=layer.activation, - ).squeeze(-1) - conv_state[decode_state_indices] = decode_conv_state + if current_platform.get_cpu_architecture() == CpuArchEnum.ARM: + decode_conv_state = conv_state[decode_state_indices].contiguous() + decode_mixed_qkv = causal_conv1d_update_torch( + x=decode_mixed_qkv.unsqueeze(-1), + conv_state=decode_conv_state, + weight=conv_weights, + bias=layer.conv1d.bias, + activation=layer.activation, + ).squeeze(-1) + conv_state[decode_state_indices] = decode_conv_state + else: + decode_mixed_qkv = causal_conv1d_update_cpu( + x=decode_mixed_qkv, + conv_state=conv_state, + weight=conv_weights, + bias=layer.conv1d.bias, + activation=layer.activation, + conv_state_indices=decode_state_indices, + ) query, key, value = layer.rearrange_mixed_qkv(decode_mixed_qkv) @@ -495,17 +508,26 @@ def _spec_aware_nonspec( decode_a = a[:num_decode_tokens] decode_state_indices = state_indices_tensor[:num_decodes] # Only the first ``width-1`` columns hold the real conv state. - decode_conv_state = conv_buf[decode_state_indices][ - :, :, : width - 1 - ].contiguous() - decode_mixed_qkv = causal_conv1d_update_torch( - x=decode_mixed_qkv.unsqueeze(-1), - conv_state=decode_conv_state, - weight=conv_weights, - bias=layer.conv1d.bias, - activation=layer.activation, - ).squeeze(-1) - conv_buf[decode_state_indices, :, : width - 1] = decode_conv_state + if current_platform.get_cpu_architecture() == CpuArchEnum.ARM: + conv_state_view = conv_buf[:, :, : width - 1] + decode_conv_state = conv_state_view[decode_state_indices].contiguous() + decode_mixed_qkv = causal_conv1d_update_torch( + x=decode_mixed_qkv.unsqueeze(-1), + conv_state=decode_conv_state, + weight=conv_weights, + bias=layer.conv1d.bias, + activation=layer.activation, + ).squeeze(-1) + conv_state_view[decode_state_indices] = decode_conv_state + else: + decode_mixed_qkv = causal_conv1d_update_cpu( + x=decode_mixed_qkv, + conv_state=conv_buf[:, :, : width - 1], + weight=conv_weights, + bias=layer.conv1d.bias, + activation=layer.activation, + conv_state_indices=decode_state_indices, + ) query, key, value = layer.rearrange_mixed_qkv(decode_mixed_qkv) # rearrange_mixed_qkv can return views whose last dim is not diff --git a/vllm/model_executor/layers/mamba/ops/cpu/mamba_ssm.py b/vllm/model_executor/layers/mamba/ops/cpu/mamba_ssm.py new file mode 100644 index 00000000000..a65793d7924 --- /dev/null +++ b/vllm/model_executor/layers/mamba/ops/cpu/mamba_ssm.py @@ -0,0 +1,144 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import torch + +import vllm._custom_ops as ops +from vllm.v1.attention.backends.utils import NULL_BLOCK_ID + + +def _mamba_chunk_scan_combined_fwd_cpu( + x, + dt, + A, + B, + C, + chunk_size, + out, + D=None, + z=None, + dt_bias=None, + initial_states=None, + return_intermediate_states=False, + seq_idx=None, + cu_seqlens=None, + cu_chunk_seqlens=None, + last_chunk_indices=None, + dt_softplus=False, + dt_limit=(0.0, float("inf")), + state_dtype=None, + **kwargs, +): + seqlen, nheads, headdim = x.shape + _, ngroups, dstate = B.shape + + assert cu_seqlens is not None + batch = cu_seqlens.size(0) - 1 + + dt_f = dt.float() + if dt_bias is not None: + dt_f = dt_f + dt_bias.float().unsqueeze(0) + if dt_softplus: + dt_f = torch.nn.functional.softplus(dt_f) + if dt_limit[0] > 0.0 or dt_limit[1] < float("inf"): + dt_f = dt_f.clamp(min=dt_limit[0], max=dt_limit[1]) + + all_states = torch.zeros( + batch, nheads, headdim, dstate, dtype=torch.float32, device=x.device + ) + if initial_states is not None: + all_states.copy_(initial_states.float()) + + assert out.is_contiguous(), ( + "_mamba_chunk_scan_combined_fwd_cpu: `out` must be " + "pre-allocated as a contiguous tensor" + ) + + D_1d = None + if D is not None: + d = D.float() + while d.dim() > 1 and d.stride(-1) == 0: + d = d.squeeze(-1) + D_1d = d.contiguous() + + ops.mamba_chunk_scan_fwd_cpu( + out, + all_states, + x, + dt_f, + A, + B, + C, + D_1d, + z, + cu_seqlens.to(torch.int32), + ) + + out_dtype = state_dtype if state_dtype is not None else x.dtype + all_states = all_states.to(out_dtype) + + return all_states + + +def selective_state_update( + state, + x, + dt, + A, + B, + C, + D=None, + z=None, + dt_bias=None, + dt_softplus=False, + state_batch_indices=None, + dst_state_batch_indices=None, + null_block_id=NULL_BLOCK_ID, + out=None, + num_accepted_tokens=None, + cu_seqlens=None, + is_blackwell=False, + enable_stochastic_rounding=False, + cache_philox_rounds=0, +): + """CPU implementation for selective_state_update.""" + # Ensure out tensor exists + if out is None: + out = torch.empty_like(x if x.dim() == 2 else x) + + _state = state.unsqueeze(1) if state.dim() == 3 else state + _x = x.unsqueeze(1) if x.dim() == 2 else x + _dt = dt.unsqueeze(1) if dt.dim() == 2 else dt + _A = A.unsqueeze(0) if A.dim() == 2 else A + _B = B.unsqueeze(1) if B.dim() == 2 else B + _C = C.unsqueeze(1) if C.dim() == 2 else C + _D = D.unsqueeze(0) if (D is not None and D.dim() == 1) else D + _z = z.unsqueeze(1) if (z is not None and z.dim() == 2) else z + _dt_bias = ( + dt_bias.unsqueeze(0) + if (dt_bias is not None and dt_bias.dim() == 1) + else dt_bias + ) + _out = out.unsqueeze(1) if out.dim() == 2 else out + + _sbi = state_batch_indices + _dsbi = dst_state_batch_indices + ops.selective_state_update_cpu( + _state, + _x, + _dt, + _A, + _B, + _C, + _D, + _z, + _dt_bias, + dt_softplus, + _sbi, + _dsbi, + null_block_id, + _out, + num_accepted_tokens, + cu_seqlens, + ) + return _out.squeeze(1) if out.dim() == 2 else _out diff --git a/vllm/model_executor/layers/mamba/ops/mamba_ssm.py b/vllm/model_executor/layers/mamba/ops/mamba_ssm.py index d348defcc76..af45467886e 100644 --- a/vllm/model_executor/layers/mamba/ops/mamba_ssm.py +++ b/vllm/model_executor/layers/mamba/ops/mamba_ssm.py @@ -845,3 +845,13 @@ def selective_scan_fn( return delta # output written inplace to delta else: return z # output written inplace to z + + +from vllm.platforms import current_platform # noqa: E402 + +if current_platform.is_cpu(): + from vllm.model_executor.layers.mamba.ops.cpu.mamba_ssm import ( + selective_state_update as selective_state_update_cpu, + ) + + selective_state_update = selective_state_update_cpu # type: ignore diff --git a/vllm/model_executor/layers/mamba/ops/ssd_combined.py b/vllm/model_executor/layers/mamba/ops/ssd_combined.py index 4c93a768b62..8c645574b9e 100644 --- a/vllm/model_executor/layers/mamba/ops/ssd_combined.py +++ b/vllm/model_executor/layers/mamba/ops/ssd_combined.py @@ -225,3 +225,11 @@ def mamba_chunk_scan_combined_varlen( ) return varlen_states + + +from vllm.platforms import current_platform # noqa: E402 + +if current_platform.is_cpu(): + import vllm.model_executor.layers.mamba.ops.cpu.mamba_ssm as cpu_mamba_ssm + + _mamba_chunk_scan_combined_fwd = cpu_mamba_ssm._mamba_chunk_scan_combined_fwd_cpu # type: ignore diff --git a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py index 92258ef204b..1795a036b12 100644 --- a/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py +++ b/vllm/model_executor/layers/mamba/ops/ssu_dispatch.py @@ -4,8 +4,9 @@ Dispatch module for Mamba selective state update (SSU) backends. Provides a unified `selective_state_update` function that dispatches to -either the Triton or FlashInfer backend based on the configured -`MambaBackendEnum`. Follows SGLang's dispatch pattern adapted for vLLM. +the Triton, FlashInfer, or CPU backend based on the configured +`MambaBackendEnum`. On CPU-only platforms (PowerPC, x86 without CUDA) +the backend defaults to 'cpu'. """ from abc import ABC, abstractmethod @@ -182,9 +183,75 @@ class FlashInferSSUBackend(MambaSSUBackend): ) +class CPUSSUBackend(MambaSSUBackend): + """CPU SSU backend using the compiled C++ VSX/scalar kernel. + + On CPU-only platforms (PowerPC, x86 without CUDA) this dispatches to + the vectorized C++ kernel registered as ``torch.ops._C.selective_state_update_cpu``. + That kernel uses vec_op SIMD intrinsics (VSX on ppc64le, AVX2 on x86, + scalar fallback elsewhere) and is parallelised with OpenMP across heads. + + Falls back to the pure-PyTorch implementation only if the C++ op is + unavailable (e.g. a CPU-less build). + """ + + def __init__(self, mamba_config: MambaConfig): + super().__init__(mamba_config) + from vllm import _custom_ops as ops + + self._cpp_kernel = ops.selective_state_update_cpu + logger.info("CPUSSUBackend: using compiled C++ selective_state_update kernel.") + + @property + def name(self) -> str: + return "cpu" + + def __call__( + self, + state: torch.Tensor, + x: torch.Tensor, + dt: torch.Tensor, + A: torch.Tensor, + B: torch.Tensor, + C: torch.Tensor, + D: torch.Tensor, + dt_bias: torch.Tensor, + z: torch.Tensor | None = None, + dt_softplus: bool = False, + state_batch_indices: torch.Tensor | None = None, + dst_state_batch_indices: torch.Tensor | None = None, + null_block_id: int = NULL_BLOCK_ID, + out: torch.Tensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + cu_seqlens: torch.Tensor | None = None, + is_blackwell: bool = False, + ) -> None: + # C++ kernel: state shape expected as (nstates, nheads, dim, dstate) + # The kernel writes in-place into `out` and updates `state`. + self._cpp_kernel( + state, + x, + dt, + A, + B, + C, + D, + z, + dt_bias, + dt_softplus, + state_batch_indices, + dst_state_batch_indices, + null_block_id, + out, + num_accepted_tokens, + cu_seqlens, + ) + + _BACKEND_REGISTRY: dict[MambaBackendEnum, type[MambaSSUBackend]] = { MambaBackendEnum.TRITON: TritonSSUBackend, MambaBackendEnum.FLASHINFER: FlashInferSSUBackend, + MambaBackendEnum.CPU: CPUSSUBackend, } _mamba_ssu_backend: MambaSSUBackend | None = None @@ -210,6 +277,20 @@ def initialize_mamba_ssu_backend( global _mamba_ssu_backend backend = mamba_config.backend + + # On CPU-only platforms (PowerPC, x86 without CUDA) Triton JIT is + # unstable or unavailable. Silently fall back to the CPU + # backend unless the user explicitly chose something other than "triton". + if backend == MambaBackendEnum.TRITON: + from vllm.platforms import current_platform + + if current_platform.is_cpu(): + logger.info( + "CPU platform detected: overriding Mamba SSU backend " + "from 'triton' to 'cpu'." + ) + backend = MambaBackendEnum.CPU + if backend not in _BACKEND_REGISTRY: raise ValueError( f"Unknown Mamba SSU backend: {backend}. " diff --git a/vllm/model_executor/layers/mamba/short_conv.py b/vllm/model_executor/layers/mamba/short_conv.py index e7e36f2fc53..7a11a4f5475 100644 --- a/vllm/model_executor/layers/mamba/short_conv.py +++ b/vllm/model_executor/layers/mamba/short_conv.py @@ -94,9 +94,13 @@ class ShortConv(MambaBase, CustomOp): # Reference torch causal conv1d; runs on all CPU platforms. AMX kernels # for causal conv can be plugged in here later. from vllm.model_executor.layers.mamba.ops.cpu.causal_conv1d import ( - causal_conv1d_torch, + causal_conv1d_fn_cpu as causal_conv1d_torch, + ) + from vllm.model_executor.layers.mamba.ops.cpu.causal_conv1d import ( + causal_conv1d_update_cpu, causal_conv1d_update_torch, ) + from vllm.platforms import CpuArchEnum, current_platform forward_context = get_forward_context() attn_metadata_raw = forward_context.attn_metadata @@ -164,17 +168,26 @@ class ShortConv(MambaBase, CustomOp): if has_decode: assert attn_metadata.state_indices_tensor_d is not None state_indices_d = attn_metadata.state_indices_tensor_d.flatten() - Bx_d = (B_d * x_d).unsqueeze(-1) # (num_decodes, dim, 1) - # Advanced indexing returns a copy; update in-place then scatter back - gathered = conv_state[state_indices_d] # (num_decodes, dim, state_len) - out_d = causal_conv1d_update_torch( - Bx_d, - gathered, - conv_weights, - self.conv.bias, - activation=None, - ).squeeze(-1) # (num_decodes, dim) - conv_state[state_indices_d] = gathered + Bx_d = B_d * x_d # (num_decodes, dim) + if current_platform.get_cpu_architecture() == CpuArchEnum.ARM: + conv_state_view = conv_state[state_indices_d].contiguous() + out_d = causal_conv1d_update_torch( + Bx_d.unsqueeze(-1), + conv_state_view, + conv_weights, + self.conv.bias, + activation=None, + ).squeeze(-1) + conv_state[state_indices_d] = conv_state_view + else: + out_d = causal_conv1d_update_cpu( + Bx_d, + conv_state, + conv_weights, + self.conv.bias, + activation=None, + conv_state_indices=state_indices_d, + ) conv_output_list.insert(0, C_d * out_d) hidden_states_out = torch.vstack(conv_output_list) diff --git a/vllm/model_executor/layers/utils.py b/vllm/model_executor/layers/utils.py index 4d2e50420f4..2b19b3d3c4e 100644 --- a/vllm/model_executor/layers/utils.py +++ b/vllm/model_executor/layers/utils.py @@ -234,10 +234,14 @@ def dispatch_cpu_unquantized_gemm( layer.cpu_linear = torch.nn.functional.linear return + # Skip CPU GEMM dispatch for non-2D weights (e.g. MoE 3D expert weights). + # These layers are handled by their own specialized methods. if layer.weight.ndim != 2: # this is not a linear layer - # For now it should be a causal_conv1d op - if torch.cpu._is_amx_tile_supported(): + # For now it should be a causal_conv1d op or MoE 3D expert weights + if torch.cpu._is_amx_tile_supported() and hasattr( + ops, "causal_conv1d_weight_pack" + ): # prepack conv weight unpacked = ( layer.weight.view( From 9bc266d923560f1c8bc1cd6f50e841f0fabacc71 Mon Sep 17 00:00:00 2001 From: Yifan Qiao Date: Sun, 19 Jul 2026 23:39:11 -0700 Subject: [PATCH 38/51] [Bugfix][KV Offload] Propagate EAGLE mode to SimpleCPU coordinator (#49071) Signed-off-by: Yifan Qiao Co-authored-by: OpenAI Codex --- vllm/v1/simple_kv_offload/manager.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/vllm/v1/simple_kv_offload/manager.py b/vllm/v1/simple_kv_offload/manager.py index 515a2f19b0e..67c6f664630 100644 --- a/vllm/v1/simple_kv_offload/manager.py +++ b/vllm/v1/simple_kv_offload/manager.py @@ -117,12 +117,13 @@ class SimpleCPUOffloadScheduler: "lazy" if lazy_offload else "eager", ) - # TODO (yifan): maybe need to enable kv_cache_events and metrics_collector here. + spec_config = vllm_config.speculative_config + use_eagle = spec_config is not None and spec_config.use_eagle() self.cpu_coordinator: KVCacheCoordinator = get_kv_cache_coordinator( kv_cache_config=self.cpu_kv_cache_config, max_model_len=vllm_config.model_config.max_model_len, max_in_flight_tokens=vllm_config.max_in_flight_tokens, - use_eagle=False, + use_eagle=use_eagle, enable_caching=True, enable_kv_cache_events=self.enable_kv_cache_events, dcp_world_size=dcp_world_size, @@ -131,7 +132,6 @@ class SimpleCPUOffloadScheduler: hash_block_size=self.hash_block_size, ) self.cpu_block_pool: BlockPool = self.cpu_coordinator.block_pool - # GPU block pool reference - bound after scheduler builds kv_cache_manager self._gpu_block_pool: BlockPool | None = None From 5245c80564aff26bd8ee1594c053e668c5920222 Mon Sep 17 00:00:00 2001 From: Itay Etelis <92247226+Etelis@users.noreply.github.com> Date: Mon, 20 Jul 2026 09:48:43 +0300 Subject: [PATCH 39/51] [Doc] Document blocks_per_chunk in the KV offloading guide (#49100) Signed-off-by: Itay Etelis Co-authored-by: Itay Etelis --- docs/features/kv_offloading_usage.md | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/docs/features/kv_offloading_usage.md b/docs/features/kv_offloading_usage.md index 6349b632915..0ae7c23f403 100644 --- a/docs/features/kv_offloading_usage.md +++ b/docs/features/kv_offloading_usage.md @@ -68,13 +68,14 @@ vllm serve \ | --- | --- | --- | --- | --- | | `spec_name` | no | `CPUOffloadingSpec` | both | Set to `TieringOffloadingSpec` for multi-tier. | | `cpu_bytes_to_use` | yes | — | both | Total bytes of host memory reserved for the CPU tier across all workers (not per-worker). | -| `block_size` | no | GPU block size | both | Offloaded block size in tokens; must be a multiple of the GPU block size. | +| `block_size` | no | GPU block size | both | Offloaded block size in tokens; must be a multiple of the GPU block size. Mutually exclusive with `blocks_per_chunk`. | +| `blocks_per_chunk` | no | `1` | both | Offloaded chunk size in GPU blocks; must be > 0. Alternative to `block_size` for models whose KV cache groups have different block sizes. | | `eviction_policy` | no | `lru` | both | Primary tier policy: `lru` or `arc`. | | `store_threshold` | no | `0` | single-tier | Min lookups before a block is offloaded. Values ≥ 2 are rejected by `TieringOffloadingSpec`. | | `max_tracker_size` | no | `64000` | single-tier | Max entries in the lookup tracker. | | `secondary_tiers` | no | `[]` | multi-tier | List of secondary tier configs (see below). | | `offload_prompt_only` | no | `true` | both | If `true`, only prompt (prefill) blocks are offloaded; decode blocks are skipped. | -| `self_describing_kv_events` | no | `false` | single-tier | Opt-in. When `true` *and* KV cache events are enabled (`--kv-events-config` with `enable_kv_cache_events`), the connector emits self-describing block-granular `BlockStored`/`BlockRemoved` payloads (constituent block hashes, whole-chunk `token_ids`, per-block `block_size`, parent hash, LoRA + group/cache-spec metadata) instead of the placeholder fallback, so external KV-event consumers can index offloaded blocks. Inert unless events are enabled. Currently rejected by `TieringOffloadingSpec`. Full-attention groups only; sliding-window/SSM groups keep the placeholder fallback. In chunk mode (`block_size` > GPU block size), overlapping chunks re-announce shared per-block hashes, so consumers must reference-count (deduplicate) repeated store/remove announcements. | +| `self_describing_kv_events` | no | `false` | single-tier | Opt-in. When `true` *and* KV cache events are enabled (`--kv-events-config` with `enable_kv_cache_events`), the connector emits self-describing block-granular `BlockStored`/`BlockRemoved` payloads (constituent block hashes, whole-chunk `token_ids`, per-block `block_size`, parent hash, LoRA + group/cache-spec metadata) instead of the placeholder fallback, so external KV-event consumers can index offloaded blocks. Inert unless events are enabled. Currently rejected by `TieringOffloadingSpec`. Full-attention groups only; sliding-window/SSM groups keep the placeholder fallback. In chunk mode (`block_size` > GPU block size, or `blocks_per_chunk` > 1), overlapping chunks re-announce shared per-block hashes, so consumers must reference-count (deduplicate) repeated store/remove announcements. | | `spec_module_path` | no | — | both | Python import path for a custom `OffloadingSpec` not in the built-in registry. Required only when `spec_name` is not built-in (advanced). | ## Secondary Tiers @@ -179,7 +180,7 @@ Rather than embedding `host`/`port` in each `secondary_tiers` entry, set them on - `cpu_bytes_to_use`: a bigger CPU tier means fewer trips to slower secondary tiers and a higher hit rate. The value is total across all workers, not per-worker. Leave headroom for the rest of the host workload. - For single-tier (CPU-only) setups, set `cpu_bytes_to_use` larger than the aggregate GPU KV cache. Because offloading is immediate, a smaller CPU tier just mirrors what the GPU already holds and adds no hit rate. -- `block_size`: larger offloaded blocks reduce per-block bookkeeping overhead but increase the granularity of lookups. Must be a multiple of the GPU block size. +- `block_size` / `blocks_per_chunk`: larger offloaded chunks reduce per-block bookkeeping overhead but increase the granularity of lookups. - FS thread counts: tune `n_read_threads` and `n_write_threads` to the parallelism your storage can sustain. Reads are latency-sensitive on the prefill path, so prefer more read threads when prefill hit rates are high. - Sharing `root_dir` across runs: runs with the same model, `block_size`, parallelism layout, and dtype share files under the same `` subdirectory. Changing any of these produces a new subdirectory; old ones are orphaned but harmless. Delete them to reclaim disk. From 9459fc647105f10f754697b3bf136d194564d603 Mon Sep 17 00:00:00 2001 From: aoshen02 Date: Mon, 20 Jul 2026 15:02:56 +0800 Subject: [PATCH 40/51] [Bugfix][RL] Set vLLM config during weight reload (#45989) Signed-off-by: aoshen02 --- .../worker/test_gpu_worker_weight_transfer.py | 28 ++++++++++++++++++ vllm/v1/worker/gpu_worker.py | 29 +++++++++++-------- 2 files changed, 45 insertions(+), 12 deletions(-) diff --git a/tests/v1/worker/test_gpu_worker_weight_transfer.py b/tests/v1/worker/test_gpu_worker_weight_transfer.py index aeb727d9ce3..6a97d64c6be 100644 --- a/tests/v1/worker/test_gpu_worker_weight_transfer.py +++ b/tests/v1/worker/test_gpu_worker_weight_transfer.py @@ -9,6 +9,7 @@ session is active. These tests verify that delegation and the session guard. import pytest +from vllm.config import VllmConfig, get_current_vllm_config from vllm.v1.worker.gpu_worker import Worker @@ -21,29 +22,55 @@ class _RecordingEngine: self.finished = False self.reset_count = 0 self.update_calls: list[dict] = [] + self.seen_configs: list[VllmConfig] = [] + + def _record_config(self) -> None: + self.seen_configs.append(get_current_vllm_config()) def start_weight_update(self) -> None: + self._record_config() self.started = True def update_weights(self, update_info: dict) -> None: + self._record_config() self.update_calls.append(update_info) if self.raise_on_update: raise ValueError("boom") def finish_weight_update(self) -> None: + self._record_config() self.finished = True def reset_weight_update_target(self) -> None: self.reset_count += 1 +class _RecordingModelRunner: + def __init__(self) -> None: + self.seen_config: VllmConfig | None = None + + def reload_weights(self) -> None: + self.seen_config = get_current_vllm_config() + + def _make_worker(engine: _RecordingEngine | None) -> Worker: worker = object.__new__(Worker) + worker.vllm_config = VllmConfig() worker.weight_transfer_engine = engine worker._weight_update_active = False return worker +def test_reload_weights_sets_current_config(): + worker = _make_worker(None) + model_runner = _RecordingModelRunner() + worker.model_runner = model_runner # type: ignore[assignment] + + Worker.reload_weights(worker) + + assert model_runner.seen_config is worker.vllm_config + + def test_start_update_finish_delegates_to_engine(): engine = _RecordingEngine() worker = _make_worker(engine) @@ -60,6 +87,7 @@ def test_start_update_finish_delegates_to_engine(): assert engine.finished is True assert engine.reset_count == 1 assert worker._weight_update_active is False + assert engine.seen_configs == [worker.vllm_config] * 3 def test_double_start_raises(): diff --git a/vllm/v1/worker/gpu_worker.py b/vllm/v1/worker/gpu_worker.py index 9c20df0d18c..7ac35ad2329 100644 --- a/vllm/v1/worker/gpu_worker.py +++ b/vllm/v1/worker/gpu_worker.py @@ -442,7 +442,8 @@ class Worker(WorkerBase): self.model_runner.update_config(overrides) def reload_weights(self, *args, **kwargs) -> None: - self.model_runner.reload_weights(*args, **kwargs) + with set_current_vllm_config(self.vllm_config): + self.model_runner.reload_weights(*args, **kwargs) @torch.inference_mode() def determine_available_memory(self) -> int: @@ -1301,14 +1302,16 @@ class Worker(WorkerBase): the configured weight transfer engine. The worker only tracks that a session is active. """ - self._start_weight_update() + with set_current_vllm_config(self.vllm_config): + self._start_weight_update() def start_draft_weight_update(self) -> None: """ Like start_weight_update, but retargets the engine at the speculative draft model for this session. """ - self._start_weight_update(is_draft=True) + with set_current_vllm_config(self.vllm_config): + self._start_weight_update(is_draft=True) def _start_weight_update(self, is_draft: bool = False) -> None: self._check_weight_transfer_engine() @@ -1355,12 +1358,13 @@ class Worker(WorkerBase): "start_weight_update must be called before update_weights." ) - try: - self.weight_transfer_engine.update_weights(update_info) - except BaseException: - self._weight_update_active = False - self.weight_transfer_engine.reset_weight_update_target() - raise + with set_current_vllm_config(self.vllm_config): + try: + self.weight_transfer_engine.update_weights(update_info) + except BaseException: + self._weight_update_active = False + self.weight_transfer_engine.reset_weight_update_target() + raise def finish_weight_update(self) -> None: """Finish the current weight update session.""" @@ -1372,9 +1376,10 @@ class Worker(WorkerBase): "finish_weight_update called without a matching start_weight_update." ) - self.weight_transfer_engine.finish_weight_update() - self.weight_transfer_engine.reset_weight_update_target() - self._weight_update_active = False + with set_current_vllm_config(self.vllm_config): + self.weight_transfer_engine.finish_weight_update() + self.weight_transfer_engine.reset_weight_update_target() + self._weight_update_active = False def shutdown(self) -> None: gc.unfreeze() From 37bf988c2f8de10165e279fff652e9d818556fe8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Micha=C5=82=20Ganczarenko?= Date: Mon, 20 Jul 2026 09:25:56 +0200 Subject: [PATCH 41/51] [XPU][Bugfix] Fix GroupCoordinator device_index (#47295) Signed-off-by: Michal Ganczarenko Co-authored-by: Kunshang Ji --- vllm/distributed/parallel_state.py | 11 ++++------- 1 file changed, 4 insertions(+), 7 deletions(-) diff --git a/vllm/distributed/parallel_state.py b/vllm/distributed/parallel_state.py index f0647323e61..b3545d54b4d 100644 --- a/vllm/distributed/parallel_state.py +++ b/vllm/distributed/parallel_state.py @@ -400,13 +400,10 @@ class GroupCoordinator: self.rank = torch.distributed.get_rank() self.local_rank = local_rank self.device_index: int - if _WORLD is not None: - self.device_index = _WORLD.device_index - else: - assert local_rank >= 0, ( - "local_rank must be provided when creating the world group" - ) - self.device_index = local_rank + assert local_rank >= 0, ( + "local_rank must be provided when creating the world group" + ) + self.device_index = local_rank self_device_group = None self_cpu_group = None From 4938d44a3b818fd443cf45aea92e670f16973ba0 Mon Sep 17 00:00:00 2001 From: Sihan Chen <757407490@qq.com> Date: Mon, 20 Jul 2026 15:33:13 +0800 Subject: [PATCH 42/51] [CPU] fixes heterogeneous NIXL KV transfer into CPU_ATTN decode workers (#47871) Signed-off-by: Spycsh Co-authored-by: Li, Jiang --- .../kv_connector/v1/nixl/base_worker.py | 7 +------ vllm/platforms/cpu.py | 21 ++++++++++++------- 2 files changed, 15 insertions(+), 13 deletions(-) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py index 834c5fccc7a..4cdc54f7e76 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py @@ -1935,13 +1935,8 @@ class NixlBaseConnectorWorker: indices = torch.tensor(block_ids, device=self.device_type, dtype=torch.long) for _, cache_or_caches in self.device_kv_caches.items(): - blocks_to_update = cache_or_caches.index_select(1, indices) current_platform.pack_kv_cache( - key=blocks_to_update[0], - value=blocks_to_update[1], - key_cache=cache_or_caches[0], - value_cache=cache_or_caches[1], - block_ids=block_ids, + kv_cache=cache_or_caches, indices=indices, ) diff --git a/vllm/platforms/cpu.py b/vllm/platforms/cpu.py index 0e805794f5f..698e83a8aba 100644 --- a/vllm/platforms/cpu.py +++ b/vllm/platforms/cpu.py @@ -461,11 +461,7 @@ class CpuPlatform(Platform): @classmethod def pack_kv_cache( cls, - key: torch.Tensor, - value: torch.Tensor, - key_cache: torch.Tensor, - value_cache: torch.Tensor, - block_ids: list[int], + kv_cache: torch.Tensor, indices: torch.Tensor, ) -> None: """ @@ -476,15 +472,26 @@ class CpuPlatform(Platform): from vllm._custom_ops import cpu_attn_reshape_and_cache from vllm.v1.attention.backends.cpu_attn import _get_attn_isa + num_blocks, num_kv_heads, block_size, fused_head_size = kv_cache.shape + head_size = fused_head_size // 2 + + # Fused path used by heterogeneous NIXL CPU_ATTN post-processing. + blocks_to_update = kv_cache.index_select(0, indices) + key = blocks_to_update[..., :head_size] + value = blocks_to_update[..., head_size:] + + key_cache, value_cache = kv_cache.view( + num_blocks, num_kv_heads, block_size * 2, head_size + ).chunk(2, dim=2) + dtype = key.dtype # For CPU_ATTN, the shape is [N, num_kv_heads, block_size, head_size] - _, _, block_size, head_size = key_cache.shape key = key.permute(0, 2, 1, 3).flatten(0, 1) value = value.permute(0, 2, 1, 3).flatten(0, 1) isa = _get_attn_isa(dtype, block_size, head_size) block_offsets = torch.arange(block_size, device="cpu", dtype=torch.long) - num_blocks = len(block_ids) + num_blocks = indices.numel() slot_mapping = ( block_offsets.reshape(1, block_size) + indices.reshape(num_blocks, 1) * block_size From df13b5aef55c5cf6a3111b28e8d48b4a71575d8d Mon Sep 17 00:00:00 2001 From: zofia <110436990+zufangzhu@users.noreply.github.com> Date: Mon, 20 Jul 2026 15:47:26 +0800 Subject: [PATCH 43/51] [XPU] [MoE] add quant input when prepare for fusedmoe (#47122) Signed-off-by: mayuyuace Signed-off-by: Zhu, Zufang Signed-off-by: zofia <110436990+zufangzhu@users.noreply.github.com> Co-authored-by: mayuyuace Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> Co-authored-by: Kunshang Ji --- .../layers/fused_moe/experts/xpu_moe.py | 29 ++++++++++++++++++- .../fused_moe/topk_weight_and_reduce.py | 2 +- vllm/model_executor/layers/fused_moe/utils.py | 10 +++++-- 3 files changed, 37 insertions(+), 4 deletions(-) diff --git a/vllm/model_executor/layers/fused_moe/experts/xpu_moe.py b/vllm/model_executor/layers/fused_moe/experts/xpu_moe.py index cde167e5d36..dfc578b4417 100644 --- a/vllm/model_executor/layers/fused_moe/experts/xpu_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/xpu_moe.py @@ -20,6 +20,7 @@ from vllm.model_executor.layers.quantization.utils.quant_utils import ( kFp8StaticTensorSym, kInt4Static, kInt4Static32, + kMxfp4Dynamic, kMxfp4Static, kMxfp8Dynamic, kMxfp8Static, @@ -64,10 +65,16 @@ class XPUExperts(mk.FusedMoEExpertsModular): ) self.gemm1_clamp_limit = quant_config.gemm1_clamp_limit self.fused_moe_impl: XpuFusedMoe | None = None + is_xe2_or_xe3 = torch.ops._xpu_C.is_xe2_arch() or torch.ops._xpu_C.is_xe3_arch() + if not is_xe2_or_xe3: + raise NotImplementedError( + "XPUExperts is only supported on Intel Xe2/Xe3 GPUs" + ) + self._expects_unquantized_inputs = is_xe2_or_xe3 @property def expects_unquantized_inputs(self) -> bool: - return True + return self._expects_unquantized_inputs @staticmethod def activation_format() -> mk.FusedMoEActivationFormat: @@ -172,6 +179,7 @@ class XPUExperts(mk.FusedMoEExpertsModular): hidden_states=hidden_states, topk_weights=topk_weights, topk_ids=topk_ids, + a1q_scale=a1q_scale, ) @@ -309,6 +317,24 @@ class XPUExpertsMxFp4(XPUExperts): num_dispatchers, ) + def workspace_shapes( + self, + M: int, + N: int, + K: int, + topk: int, + global_num_experts: int, + local_num_experts: int, + expert_tokens_meta: mk.ExpertTokensMetadata | None, + activation: MoEActivation, + ) -> tuple[tuple[int, ...], tuple[int, ...], tuple[int, ...]]: + # K = a1q.size(-1). When activations are pre-quantized packed mxfp4, + # K is the packed hidden_size (= logical / 2); the kernel output is at + # logical hidden_size (2 * K). When unquantized (bf16), K is already + # the logical size. + logical_K = K if self.expects_unquantized_inputs else 2 * K + return (0,), (0,), (M, logical_K) + @staticmethod def _supports_quant_scheme( weight_key: QuantKey | None, @@ -316,5 +342,6 @@ class XPUExpertsMxFp4(XPUExperts): ) -> bool: SUPPORTED_W_A = [ (kMxfp4Static, None), + (kMxfp4Static, kMxfp4Dynamic), ] return (weight_key, activation_key) in SUPPORTED_W_A diff --git a/vllm/model_executor/layers/fused_moe/topk_weight_and_reduce.py b/vllm/model_executor/layers/fused_moe/topk_weight_and_reduce.py index 837c1498622..81c84d64be0 100644 --- a/vllm/model_executor/layers/fused_moe/topk_weight_and_reduce.py +++ b/vllm/model_executor/layers/fused_moe/topk_weight_and_reduce.py @@ -71,7 +71,7 @@ class TopKWeightAndReduceNoOP(mk.TopKWeightAndReduce): assert output.size() == fused_expert_output.size(), ( "output shape is expected to match the fused_expert_output shape. " f"But got output={output.size()}, " - f"used_expert_output={fused_expert_output.size()}" + f"fused_expert_output={fused_expert_output.size()}" ) output.copy_(fused_expert_output, non_blocking=True) return output diff --git a/vllm/model_executor/layers/fused_moe/utils.py b/vllm/model_executor/layers/fused_moe/utils.py index 6aa4eb0cd2a..e3d6493dda2 100644 --- a/vllm/model_executor/layers/fused_moe/utils.py +++ b/vllm/model_executor/layers/fused_moe/utils.py @@ -17,12 +17,14 @@ from vllm.model_executor.layers.quantization.utils.int8_utils import ( ) from vllm.model_executor.layers.quantization.utils.mxfp4_utils import ( quant_dequant_mxfp4, + xpu_mxfp4_quantize, ) from vllm.model_executor.layers.quantization.utils.mxfp6_utils import ( quant_dequant_mxfp6, ) from vllm.model_executor.layers.quantization.utils.mxfp8_utils import ( mxfp8_e4m3_quantize, + xpu_mxfp8_quantize, ) from vllm.model_executor.layers.quantization.utils.nvfp4_emulation_utils import ( ref_nvfp4_quant_dequant, @@ -195,6 +197,8 @@ def _mxfp4_quantize( per_act_token_quant: bool, block_shape: list[int] | None = None, ) -> tuple[torch.Tensor, None]: + if current_platform.is_xpu(): + return xpu_mxfp4_quantize(A) assert block_shape is None # TODO: native mxfp4 is currently not integrated in vllm, # so simulating even on devices supporting this data type natively. @@ -223,6 +227,8 @@ def _mxfp8_e4m3_quantize( is_sf_swizzled_layout: bool = False, mx_alignment: int = 0, ) -> tuple[torch.Tensor, torch.Tensor]: + if current_platform.is_xpu(): + return xpu_mxfp8_quantize(A) assert A_scale is None assert not per_act_token_quant assert block_shape is None or block_shape == [1, 32] @@ -309,7 +315,7 @@ def moe_kernel_quantize_input( A = ref_nvfp4_quant_dequant(A, A_scale, block_size=16) return A, None elif quant_dtype == "mxfp4": - if not quantization_emulation: + if not current_platform.is_xpu() and not quantization_emulation: raise NotImplementedError( "moe_kernel_quantize_input should not be used for native" " quant_dtype='mxfp4' MOE. Please open an issue." @@ -318,7 +324,7 @@ def moe_kernel_quantize_input( elif quant_dtype == "mxfp8": # TODO: `quant_dtype == "mxfp8"` is ambiguous, # should be fp8_e4m3. OCP MX also defines `fp8_e5m2`. - if quantization_emulation: + if not current_platform.is_xpu() and quantization_emulation: raise NotImplementedError( "moe_kernel_quantize_input does not support quant_dtype='mxfp8' MOE " "quantization emulation. Please open an issue." From f1f1259692004c909f8b7f8f371d42c7a7871fa3 Mon Sep 17 00:00:00 2001 From: Sage <80211083+sagearc@users.noreply.github.com> Date: Mon, 20 Jul 2026 11:28:25 +0300 Subject: [PATCH 44/51] [Rust Frontend] Use zero-copy slicing for multimodal tensors (#48781) Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> Signed-off-by: Sage Ahrac --- rust/src/chat/src/multimodal/item.rs | 8 +- rust/src/chat/src/multimodal/tensor.rs | 149 ++++++++---------- rust/src/chat/src/multimodal/video.rs | 8 +- .../engine-core-client/src/protocol/tensor.rs | 25 +-- 4 files changed, 91 insertions(+), 99 deletions(-) diff --git a/rust/src/chat/src/multimodal/item.rs b/rust/src/chat/src/multimodal/item.rs index d5600c56880..69ab71af307 100644 --- a/rust/src/chat/src/multimodal/item.rs +++ b/rust/src/chat/src/multimodal/item.rs @@ -38,7 +38,7 @@ pub(super) fn build_batched_items( let keep_on_cpu = spec.keep_on_cpu_keys.contains(key); let (value, field) = match spec.field_layout_for(key) { Some(FieldLayout::Batched) => ( - tensor.batched_value_at(index)?, + tensor.batched_wire_value_at(index)?, MmField::Batched(MmBatchedField { keep_on_cpu }), ), Some(FieldLayout::Flat { sizes_key }) => { @@ -47,7 +47,7 @@ pub(super) fn build_batched_items( })?; let (start, end) = tensor::flat_range_for_index(sizes, sizes_key, index)?; ( - tensor.flat_value_range(start, end)?, + tensor.flat_wire_value_range(start, end)?, MmField::Flat(MmFlatField { slices: vec![MmSlice::Slice(SliceSpec { start: Some(0), @@ -60,7 +60,7 @@ pub(super) fn build_batched_items( ) } None => ( - tensor.clone(), + tensor.try_into()?, MmField::Shared(MmSharedField { batch_size: len, keep_on_cpu, @@ -71,7 +71,7 @@ pub(super) fn build_batched_items( data.insert( key.clone(), MmFieldElem { - data: Some(value.try_into()?), + data: Some(value), field, }, ); diff --git a/rust/src/chat/src/multimodal/tensor.rs b/rust/src/chat/src/multimodal/tensor.rs index f701e0646c6..35b57dd5fe6 100644 --- a/rust/src/chat/src/multimodal/tensor.rs +++ b/rust/src/chat/src/multimodal/tensor.rs @@ -12,7 +12,7 @@ use vllm_engine_core_client::protocol::tensor::{ShapeExt as _, WireTensor}; use crate::error::{Error, Result, bail_multimodal, multimodal}; /// Representation for multimodal kwarg values for transformation. -#[derive(Debug, Clone)] +#[derive(Debug)] pub(super) enum KwargValue { /// Float tensor with row-major flat data and shape. F32Tensor { data: Vec, shape: Vec }, @@ -107,28 +107,19 @@ impl KwargValue { } } -impl TryFrom for ProtocolKwargValue { +impl TryFrom<&KwargValue> for ProtocolKwargValue { type Error = Error; - fn try_from(value: KwargValue) -> Result { - match value { - KwargValue::F32Tensor { data, shape } => Ok(Self::Tensor( - WireTensor::from_f32(shape, data).map_err(Error::Multimodal)?, - )), - KwargValue::F16Tensor { data, shape } => Ok(Self::Tensor( - WireTensor::from_f16(shape, data).map_err(Error::Multimodal)?, - )), - KwargValue::Bf16Tensor { data, shape } => Ok(Self::Tensor( - WireTensor::from_bf16(shape, data).map_err(Error::Multimodal)?, - )), - KwargValue::I64Tensor { data, shape } => Ok(Self::Tensor( - WireTensor::from_i64(shape, data).map_err(Error::Multimodal)?, - )), - KwargValue::U32Tensor { data, shape } => Ok(Self::Tensor( - WireTensor::from_u32(shape, data).map_err(Error::Multimodal)?, - )), - KwargValue::Passthrough(value) => Ok(value), - } + fn try_from(value: &KwargValue) -> Result { + let tensor = match value { + KwargValue::F32Tensor { data, shape } => WireTensor::from_f32(shape.clone(), data), + KwargValue::F16Tensor { data, shape } => WireTensor::from_f16(shape.clone(), data), + KwargValue::Bf16Tensor { data, shape } => WireTensor::from_bf16(shape.clone(), data), + KwargValue::I64Tensor { data, shape } => WireTensor::from_i64(shape.clone(), data), + KwargValue::U32Tensor { data, shape } => WireTensor::from_u32(shape.clone(), data), + KwargValue::Passthrough(value) => return Ok(value.clone()), + }; + tensor.map(ProtocolKwargValue::Tensor).map_err(Error::Multimodal) } } @@ -145,63 +136,55 @@ impl KwargValue { } } - /// Extract one media item from a batched tensor field. + /// Convert one media item from a batched tensor field to wire bytes. /// /// Batched fields use their first axis as media-item index and drop that /// axis in the per-feature value, matching vLLM's batched-field semantics. - pub(super) fn batched_value_at(&self, index: usize) -> Result { - match self { - Self::F32Tensor { data, shape } => { - let (shape, data) = slice_first_axis_range(shape, data, index, index + 1, true)?; - Ok(Self::F32Tensor { data, shape }) - } - Self::F16Tensor { data, shape } => { - let (shape, data) = slice_first_axis_range(shape, data, index, index + 1, true)?; - Ok(Self::F16Tensor { data, shape }) - } - Self::Bf16Tensor { data, shape } => { - let (shape, data) = slice_first_axis_range(shape, data, index, index + 1, true)?; - Ok(Self::Bf16Tensor { data, shape }) - } - Self::I64Tensor { data, shape } => { - let (shape, data) = slice_first_axis_range(shape, data, index, index + 1, true)?; - Ok(Self::I64Tensor { data, shape }) - } - Self::U32Tensor { data, shape } => { - let (shape, data) = slice_first_axis_range(shape, data, index, index + 1, true)?; - Ok(Self::U32Tensor { data, shape }) - } - Self::Passthrough(value) => Ok(Self::Passthrough(value.clone())), - } + pub(super) fn batched_wire_value_at(&self, index: usize) -> Result { + self.wire_value_range(index, index + 1, true) } - /// Extract one media item's variable-length range from a flat tensor field. + /// Convert one media item's flat tensor range directly to wire bytes. /// /// Flat fields keep the first axis as the sliced length for this item. - pub(super) fn flat_value_range(&self, start: usize, end: usize) -> Result { - match self { + pub(super) fn flat_wire_value_range( + &self, + start: usize, + end: usize, + ) -> Result { + self.wire_value_range(start, end, false) + } + + fn wire_value_range( + &self, + start: usize, + end: usize, + drop_axis: bool, + ) -> Result { + let tensor = match self { Self::F32Tensor { data, shape } => { - let (shape, data) = slice_first_axis_range(shape, data, start, end, false)?; - Ok(Self::F32Tensor { data, shape }) + let (shape, data) = slice_first_axis_range(shape, data, start, end, drop_axis)?; + WireTensor::from_f32(shape, data) } Self::F16Tensor { data, shape } => { - let (shape, data) = slice_first_axis_range(shape, data, start, end, false)?; - Ok(Self::F16Tensor { data, shape }) + let (shape, data) = slice_first_axis_range(shape, data, start, end, drop_axis)?; + WireTensor::from_f16(shape, data) } Self::Bf16Tensor { data, shape } => { - let (shape, data) = slice_first_axis_range(shape, data, start, end, false)?; - Ok(Self::Bf16Tensor { data, shape }) + let (shape, data) = slice_first_axis_range(shape, data, start, end, drop_axis)?; + WireTensor::from_bf16(shape, data) } Self::I64Tensor { data, shape } => { - let (shape, data) = slice_first_axis_range(shape, data, start, end, false)?; - Ok(Self::I64Tensor { data, shape }) + let (shape, data) = slice_first_axis_range(shape, data, start, end, drop_axis)?; + WireTensor::from_i64(shape, data) } Self::U32Tensor { data, shape } => { - let (shape, data) = slice_first_axis_range(shape, data, start, end, false)?; - Ok(Self::U32Tensor { data, shape }) + let (shape, data) = slice_first_axis_range(shape, data, start, end, drop_axis)?; + WireTensor::from_u32(shape, data) } - Self::Passthrough(value) => Ok(Self::Passthrough(value.clone())), - } + Self::Passthrough(value) => return Ok(value.clone()), + }; + tensor.map(ProtocolKwargValue::Tensor).map_err(Error::Multimodal) } } @@ -240,13 +223,13 @@ fn tensor_as_usize_vec(tensor: &KwargValue) -> Result> { } /// Slice a flat row-major tensor along its first axis. -fn slice_first_axis_range( +fn slice_first_axis_range<'a, T>( shape: &[usize], - data: &[T], + data: &'a [T], start: usize, end: usize, drop_axis: bool, -) -> Result<(Vec, Vec)> { +) -> Result<(Vec, &'a [T])> { let first_dim = *shape.first().ok_or_else(|| multimodal!("tensor has no first dimension"))?; if start > end || end > first_dim { bail_multimodal!("invalid tensor slice {start}..{end} for first dimension {first_dim}"); @@ -270,7 +253,7 @@ fn slice_first_axis_range( shape[0] = end - start; shape }; - Ok((out_shape, data[data_start..data_end].to_vec())) + Ok((out_shape, &data[data_start..data_end])) } #[cfg(test)] @@ -278,35 +261,39 @@ mod tests { use super::*; #[test] - fn batched_value_at_drops_first_axis() { + fn batched_wire_value_at_drops_first_axis() { let value = KwargValue::F32Tensor { data: vec![1.0, 2.0, 3.0, 4.0], shape: vec![2, 2], }; - let value = value.batched_value_at(1).unwrap(); + let ProtocolKwargValue::Tensor(tensor) = value.batched_wire_value_at(1).unwrap() else { + panic!("expected tensor"); + }; - assert!(matches!( - value, - KwargValue::F32Tensor { data, shape } - if shape == vec![2] && data == vec![3.0, 4.0] - )); + assert_eq!(tensor.shape, vec![2]); + assert_eq!( + tensor.data.into_raw_view().unwrap(), + [3.0_f32, 4.0].into_iter().flat_map(f32::to_ne_bytes).collect::>() + ); } #[test] - fn flat_value_range_keeps_first_axis() { + fn flat_wire_value_range_keeps_first_axis() { let value = KwargValue::U32Tensor { data: (0..10).collect(), shape: vec![5, 2], }; - let value = value.flat_value_range(1, 3).unwrap(); + let ProtocolKwargValue::Tensor(tensor) = value.flat_wire_value_range(1, 3).unwrap() else { + panic!("expected tensor"); + }; - assert!(matches!( - value, - KwargValue::U32Tensor { data, shape } - if shape == vec![2, 2] && data == vec![2, 3, 4, 5] - )); + assert_eq!(tensor.shape, vec![2, 2]); + assert_eq!( + tensor.data.into_raw_view().unwrap(), + [2_u32, 3, 4, 5].into_iter().flat_map(u32::to_ne_bytes).collect::>() + ); } #[test] @@ -336,7 +323,7 @@ mod tests { let value = KwargValue::from_f32_tensor(vec![1.0, -1.0], vec![2], ModelDtype::BFloat16).unwrap(); - let ProtocolKwargValue::Tensor(tensor) = ProtocolKwargValue::try_from(value).unwrap() + let ProtocolKwargValue::Tensor(tensor) = ProtocolKwargValue::try_from(&value).unwrap() else { panic!("expected tensor"); }; @@ -351,7 +338,7 @@ mod tests { let value = KwargValue::from_f32_tensor(vec![1.0, -1.0], vec![2], ModelDtype::Float16).unwrap(); - let ProtocolKwargValue::Tensor(tensor) = ProtocolKwargValue::try_from(value).unwrap() + let ProtocolKwargValue::Tensor(tensor) = ProtocolKwargValue::try_from(&value).unwrap() else { panic!("expected tensor"); }; diff --git a/rust/src/chat/src/multimodal/video.rs b/rust/src/chat/src/multimodal/video.rs index 9cc3238ac60..6c9b961e7c2 100644 --- a/rust/src/chat/src/multimodal/video.rs +++ b/rust/src/chat/src/multimodal/video.rs @@ -130,7 +130,7 @@ fn build_video_item( let keep_on_cpu = support.spec.keep_on_cpu_keys.contains(&key); let (value, field) = match support.spec.field_layout_for(&key) { Some(FieldLayout::Batched) => ( - tensor.batched_value_at(0)?, + tensor.batched_wire_value_at(0)?, MmField::Batched(MmBatchedField { keep_on_cpu }), ), Some(FieldLayout::Flat { .. }) => { @@ -138,7 +138,7 @@ fn build_video_item( .first_dim() .ok_or_else(|| multimodal!("flat video input `{key}` is not a tensor"))?; ( - tensor, + (&tensor).try_into()?, MmField::Flat(MmFlatField { slices: vec![MmSlice::Slice(SliceSpec { start: Some(0), @@ -151,7 +151,7 @@ fn build_video_item( ) } None => ( - tensor, + (&tensor).try_into()?, MmField::Shared(MmSharedField { batch_size: 1, keep_on_cpu, @@ -162,7 +162,7 @@ fn build_video_item( data.insert( key, MmFieldElem { - data: Some(value.try_into()?), + data: Some(value), field, }, ); diff --git a/rust/src/engine-core-client/src/protocol/tensor.rs b/rust/src/engine-core-client/src/protocol/tensor.rs index 5eb1668e752..519bc238dec 100644 --- a/rust/src/engine-core-client/src/protocol/tensor.rs +++ b/rust/src/engine-core-client/src/protocol/tensor.rs @@ -55,52 +55,57 @@ pub struct WireNdArray { impl WireNdArray { /// Build a float32 tensor/ndarray backed by native-endian raw-view bytes. - pub fn from_f32(shape: Vec, data: Vec) -> Result { + pub fn from_f32(shape: Vec, data: impl AsRef<[f32]>) -> Result { + let data = data.as_ref(); validate_element_count(&shape, data.len())?; Ok(Self { dtype: "float32".to_string(), shape, - data: WireArrayData::RawView(pod_collect_to_vec::(&data)), + data: WireArrayData::RawView(pod_collect_to_vec::(data)), }) } /// Build a float16 tensor/ndarray backed by native-endian raw-view bytes. - pub fn from_f16(shape: Vec, data: Vec) -> Result { + pub fn from_f16(shape: Vec, data: impl AsRef<[f16]>) -> Result { + let data = data.as_ref(); validate_element_count(&shape, data.len())?; Ok(Self { dtype: "float16".to_string(), shape, - data: WireArrayData::RawView(pod_collect_to_vec::(&data)), + data: WireArrayData::RawView(pod_collect_to_vec::(data)), }) } /// Build a bfloat16 tensor/ndarray backed by native-endian raw-view bytes. - pub fn from_bf16(shape: Vec, data: Vec) -> Result { + pub fn from_bf16(shape: Vec, data: impl AsRef<[bf16]>) -> Result { + let data = data.as_ref(); validate_element_count(&shape, data.len())?; Ok(Self { dtype: "bfloat16".to_string(), shape, - data: WireArrayData::RawView(pod_collect_to_vec::(&data)), + data: WireArrayData::RawView(pod_collect_to_vec::(data)), }) } /// Build an int64 tensor/ndarray backed by native-endian raw-view bytes. - pub fn from_i64(shape: Vec, data: Vec) -> Result { + pub fn from_i64(shape: Vec, data: impl AsRef<[i64]>) -> Result { + let data = data.as_ref(); validate_element_count(&shape, data.len())?; Ok(Self { dtype: "int64".to_string(), shape, - data: WireArrayData::RawView(pod_collect_to_vec::(&data)), + data: WireArrayData::RawView(pod_collect_to_vec::(data)), }) } /// Build a uint32 tensor/ndarray backed by native-endian raw-view bytes. - pub fn from_u32(shape: Vec, data: Vec) -> Result { + pub fn from_u32(shape: Vec, data: impl AsRef<[u32]>) -> Result { + let data = data.as_ref(); validate_element_count(&shape, data.len())?; Ok(Self { dtype: "uint32".to_string(), shape, - data: WireArrayData::RawView(pod_collect_to_vec::(&data)), + data: WireArrayData::RawView(pod_collect_to_vec::(data)), }) } From 823eaf667d33125259d2d280b0ed00747fa92440 Mon Sep 17 00:00:00 2001 From: Xiaochang Wu Date: Mon, 20 Jul 2026 16:32:03 +0800 Subject: [PATCH 45/51] [XPU] FP8 o_proj with fp8_bmm and load-time scale transpose (#48334) Signed-off-by: Wu, Xiaochang Co-authored-by: Kunshang Ji Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> --- vllm/_xpu_ops.py | 63 +++++++++++++++++++ .../kernels/linear/scaled_mm/xpu.py | 35 ++++++++++- .../common/ops/fused_inv_rope_fp8_quant.py | 3 +- vllm/models/deepseek_v4/xpu/model.py | 2 - vllm/models/deepseek_v4/xpu/xpu_sparse.py | 36 ++++++++--- 5 files changed, 124 insertions(+), 15 deletions(-) diff --git a/vllm/_xpu_ops.py b/vllm/_xpu_ops.py index 8875ed49f6e..a131d5e4d31 100644 --- a/vllm/_xpu_ops.py +++ b/vllm/_xpu_ops.py @@ -219,6 +219,63 @@ def _xpu_ops_deepseek_scaling_rope_fake( return query, key +def _xpu_fp8_bmm_impl( + a: torch.Tensor, + b: torch.Tensor, + out_dtype: torch.dtype, + a_scale: torch.Tensor, + b_scale: torch.Tensor, + bias: torch.Tensor | None, +) -> torch.Tensor: + """XPU FP8 batched GEMM implementation for ``torch.ops.vllm.xpu_fp8_bmm``. + + Computes batched matrix multiplication over the leading group dimension: + ``[G, M, K] @ [G, K, N] -> [G, M, N]``. + + Args: + a: FP8 activation tensor with shape ``[G, M, K]``. + Does not need to be contiguous. + b: FP8 weight tensor with shape ``[G, K, N]``. + Does not need to be contiguous. + out_dtype: Output dtype accepted by the kernel (typically + ``torch.bfloat16`` for the DeepSeek-V4 O-proj path). + a_scale: Activation scale tensor for ``a``. + In current DeepSeek-V4 XPU usage it is block-scaled with shape + ``[G, M, K/bs]`` (``bs`` is the quant block size, e.g. 128). + Must be contiguous. + b_scale: Weight scale tensor for ``b``. + In current DeepSeek-V4 XPU usage it is block-scaled with shape + ``[G, K/bs, N/bs]`` (``bs`` is the quant block size, e.g. 128). + Must be contiguous. + bias: Optional bias tensor. Pass ``None`` when no bias is required. + + Returns: + Output tensor with shape ``[G, M, N]`` and dtype ``out_dtype``. + + Notes: + This implementation centralizes access to + ``torch.ops._xpu_C.fp8_bmm``. Both scales must be contiguous, while + ``a`` and ``b`` may be non-contiguous views. + """ + return torch.ops._xpu_C.fp8_bmm(a, b, out_dtype, a_scale, b_scale, bias) + + +def _xpu_fp8_bmm_fake( + a: torch.Tensor, + b: torch.Tensor, + out_dtype: torch.dtype, + a_scale: torch.Tensor, + b_scale: torch.Tensor, + bias: torch.Tensor | None, +) -> torch.Tensor: + # [G, M, K] @ [G, K, N] => [G, M, N] + return torch.empty( + (a.shape[0], a.shape[1], b.shape[2]), + dtype=out_dtype, + device=a.device, + ) + + def _xpu_fp8_mqa_logits_impl( q: torch.Tensor, k_quant: torch.Tensor, @@ -1053,6 +1110,12 @@ class xpu_ops: fake_impl=_xpu_mxfp4_quantize_fake, ) + direct_register_custom_op( + op_name="xpu_fp8_bmm", + op_func=_xpu_fp8_bmm_impl, + fake_impl=_xpu_fp8_bmm_fake, + ) + direct_register_custom_op( op_name="xpu_fp8_mqa_logits", op_func=_xpu_fp8_mqa_logits_impl, diff --git a/vllm/model_executor/kernels/linear/scaled_mm/xpu.py b/vllm/model_executor/kernels/linear/scaled_mm/xpu.py index f30d6ced1d3..1441c118e7a 100644 --- a/vllm/model_executor/kernels/linear/scaled_mm/xpu.py +++ b/vllm/model_executor/kernels/linear/scaled_mm/xpu.py @@ -197,6 +197,37 @@ class XPUFp8BlockScaledMMKernel(Fp8BlockScaledMMLinearKernel): return False, "XPUFp8BlockScaledMM only support on XPU" return True, None + def process_weights_after_loading(self, layer: torch.nn.Module): + super().process_weights_after_loading(layer) + scale_attr = ( + "weight_scale_inv" if hasattr(layer, "weight_scale_inv") else "weight_scale" + ) + scale = getattr(layer, scale_attr) + # Transpose scale from checkpoint layout [N/128, K/128] to + # oneDNN expected layout [K/128, N/128] at load time (one-time cost). + scale_t = scale.data.t().contiguous() + replace_parameter(layer, scale_attr, scale_t) + + # For BMM layers (e.g. wo_a), precompute 3D scale and weight: + # [K/bs, N/bs] -> [batch, K/bs, N_per_batch/bs] + if getattr(layer, "is_bmm", False): + batch = layer.bmm_batch_size + k_blocks = scale_t.shape[0] + n_per_batch_blocks = scale_t.shape[1] // batch + layer.bmm_scale = ( + scale_t.reshape(k_blocks, batch, n_per_batch_blocks) + .permute(1, 0, 2) + .contiguous() + ) + # Precompute [G, K, N] weight for fp8_bmm. + # Original weight is [N_total, K] where N_total = G * N_per_group. + w = layer.weight.data + N_total, K = w.shape + N_per_group = N_total // batch + layer.bmm_weight = w.reshape(batch, N_per_group, K).permute( + 0, 2, 1 + ) # [G, K, N] + def apply_block_scaled_mm( self, A: torch.Tensor, @@ -205,12 +236,12 @@ class XPUFp8BlockScaledMMKernel(Fp8BlockScaledMMLinearKernel): Bs: torch.Tensor, ) -> torch.Tensor: # Weight is [N, K]. Use .t() to create a [K, N] view without copying. - # Bs is [N/128, K/128] — transpose to [K/128, N/128] for oneDNN. + # Bs is already [K/128, N/128] from process_weights_after_loading. return torch.ops._xpu_C.fp8_gemm( A, B.t(), self.config.out_dtype, As, - Bs.t().contiguous(), + Bs, torch.Tensor(), ) diff --git a/vllm/models/deepseek_v4/common/ops/fused_inv_rope_fp8_quant.py b/vllm/models/deepseek_v4/common/ops/fused_inv_rope_fp8_quant.py index b667b87679c..aa8750580d0 100644 --- a/vllm/models/deepseek_v4/common/ops/fused_inv_rope_fp8_quant.py +++ b/vllm/models/deepseek_v4/common/ops/fused_inv_rope_fp8_quant.py @@ -257,7 +257,6 @@ def _fused_inv_rope_fp8_quant_kernel_impl( ) grid = (tma_aligned_T, n_groups * heads_per_group) use_gdc = current_platform.is_arch_support_pdl() - pdl_kwargs = {"launch_pdl": True} if use_gdc else {} _fused_inv_rope_fp8_quant_per_head[grid]( o, positions, @@ -281,8 +280,8 @@ def _fused_inv_rope_fp8_quant_kernel_impl( HALF_ROPE=half_rope, TMA_ALIGNED_SCALES=tma_aligned_scales, USE_GDC=use_gdc, + launch_pdl=use_gdc, num_stages=1, - **pdl_kwargs, num_warps=1, ) return fp8_buf, scale_buf diff --git a/vllm/models/deepseek_v4/xpu/model.py b/vllm/models/deepseek_v4/xpu/model.py index e8449b9c058..f70cac1bce6 100644 --- a/vllm/models/deepseek_v4/xpu/model.py +++ b/vllm/models/deepseek_v4/xpu/model.py @@ -8,7 +8,6 @@ import regex as re import torch import torch.nn as nn -from vllm.compilation.decorators import support_torch_compile from vllm.config import VllmConfig from vllm.distributed import ( get_ep_group, @@ -978,7 +977,6 @@ class DeepseekV4DecoderLayer(nn.Module): return x, residual, post_mix, res_mix -@support_torch_compile class DeepseekV4Model(nn.Module, EagleModelMixin): def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): super().__init__() diff --git a/vllm/models/deepseek_v4/xpu/xpu_sparse.py b/vllm/models/deepseek_v4/xpu/xpu_sparse.py index 77cc35cf492..73439d86461 100644 --- a/vllm/models/deepseek_v4/xpu/xpu_sparse.py +++ b/vllm/models/deepseek_v4/xpu/xpu_sparse.py @@ -89,19 +89,37 @@ class DeepseekV4XPUAttention(DeepseekV4Attention): return num_heads def _o_proj(self, o: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: - # XPU uses BF16 reference wo_a path (same as ROCm). - from vllm.models.deepseek_v4.amd.rocm import rocm_inv_rope_einsum + from vllm.models.deepseek_v4.common.ops.fused_inv_rope_fp8_quant import ( + fused_inv_rope_fp8_quant, + ) - z = rocm_inv_rope_einsum( - self.rotary_emb, + o_fp8, o_scale = fused_inv_rope_fp8_quant( o, positions, - self.rope_head_dim, - self.n_local_groups, - self.o_lora_rank, - self.wo_a, + self.rotary_emb.cos_sin_cache, + n_groups=self.n_local_groups, + heads_per_group=self.n_local_heads // self.n_local_groups, + nope_dim=self.nope_head_dim, + rope_dim=self.rope_head_dim, + tma_aligned_scales=False, ) - return self.wo_b(z.flatten(1)) + + # Precomputed contiguous [G, K, N] weight and [G, K/bs, N/bs] scale. + wo_a_weight = self.wo_a.bmm_weight + wo_a_scale = self.wo_a.bmm_scale + + # TODO: optimize fused_inv_rope_fp8_quant for xpu bmm to + # eliminate o_scale transpose + contiguous + z = torch.ops.vllm.xpu_fp8_bmm( + o_fp8.transpose(0, 1), + wo_a_weight, + torch.bfloat16, + o_scale.transpose(0, 1).contiguous(), + wo_a_scale, + None, + ) + + return self.wo_b(z.transpose(0, 1).flatten(1)) def forward_mqa( self, From c01618fdc81524ed4e91652bbf85ed23d6ddc448 Mon Sep 17 00:00:00 2001 From: Bugen Zhao Date: Mon, 20 Jul 2026 17:31:25 +0800 Subject: [PATCH 46/51] [Rust][Benchmark] Integrate `vllm-bench` to `vllm-rs` & `vllm` CLI (#48930) Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> Signed-off-by: Bugen Zhao --- rust/Cargo.lock | 1 + rust/Cargo.toml | 1 + rust/src/bench/src/cli.rs | 15 +- rust/src/bench/src/config.rs | 400 +++++++++++++----------- rust/src/bench/src/hub.rs | 7 +- rust/src/bench/src/lib.rs | 86 +++++ rust/src/bench/src/main.rs | 88 +----- rust/src/cmd/Cargo.toml | 1 + rust/src/cmd/src/cli.rs | 12 +- rust/src/cmd/src/cli/tests.rs | 22 +- rust/src/cmd/src/main.rs | 6 +- vllm/entrypoints/cli/benchmark/serve.py | 62 +++- 12 files changed, 421 insertions(+), 280 deletions(-) create mode 100644 rust/src/bench/src/lib.rs diff --git a/rust/Cargo.lock b/rust/Cargo.lock index 0709ef8fc3a..e7a157923f8 100644 --- a/rust/Cargo.lock +++ b/rust/Cargo.lock @@ -5560,6 +5560,7 @@ dependencies = [ "tracing", "tracing-subscriber", "uuid", + "vllm-bench", "vllm-chat", "vllm-engine-core-client", "vllm-managed-engine", diff --git a/rust/Cargo.toml b/rust/Cargo.toml index 492b1c5e260..5fc9ea6edb7 100644 --- a/rust/Cargo.toml +++ b/rust/Cargo.toml @@ -132,6 +132,7 @@ trait-set = "0.3.0" url = "2.5.7" uuid = { version = "1.22.0", features = ["v4"] } validator = { version = "0.20.0", features = ["derive"] } +vllm-bench = { path = "src/bench" } vllm-chat = { path = "src/chat" } vllm-engine-core-client = { path = "src/engine-core-client" } vllm-llm = { path = "src/llm" } diff --git a/rust/src/bench/src/cli.rs b/rust/src/bench/src/cli.rs index 9aa8f388633..06e32e14ce3 100644 --- a/rust/src/bench/src/cli.rs +++ b/rust/src/bench/src/cli.rs @@ -3,8 +3,6 @@ use std::fmt; -use clap::Parser; - /// Backend type for the benchmark endpoint. #[derive(clap::ValueEnum, Debug, Clone, Copy, PartialEq, Eq)] pub enum BackendKind { @@ -77,7 +75,7 @@ pub enum DatasetName { ShareGpt, #[value(name = "sonnet")] Sonnet, - #[value(name = "speed-bench")] + #[value(name = "speed-bench", alias = "speed_bench")] SpeedBench, #[value(name = "hf")] Hf, @@ -144,13 +142,8 @@ impl fmt::Display for SpeedBenchConfig { } /// High-performance benchmark client for vLLM serving endpoints. -#[derive(Parser, Debug, Clone)] -#[command( - name = "vllm-bench", - about = "Benchmark online serving throughput", - version -)] -pub struct Cli { +#[derive(clap::Args, Debug, Clone)] +pub struct BenchServeArgs { /// The type of backend or endpoint to use for the benchmark. #[arg(long, default_value = "openai")] pub backend: BackendKind, @@ -659,7 +652,7 @@ pub struct Cli { pub lora_assignment: LoraAssignment, } -impl Cli { +impl BenchServeArgs { /// Resolve the base URL from explicit --base-url or from --host/--port. pub fn resolve_base_url(&self) -> String { if let Some(ref base) = self.base_url { diff --git a/rust/src/bench/src/config.rs b/rust/src/bench/src/config.rs index 9696612c30f..ec132d11969 100644 --- a/rust/src/bench/src/config.rs +++ b/rust/src/bench/src/config.rs @@ -4,7 +4,9 @@ use std::collections::HashMap; use std::sync::Arc; -use crate::cli::{BackendKind, Cli, DatasetName, LoraAssignment, RampUpStrategy, SpeedBenchConfig}; +use crate::cli::{ + BackendKind, BenchServeArgs, DatasetName, LoraAssignment, RampUpStrategy, SpeedBenchConfig, +}; use crate::datasets::random_mm::{MmBucketKey, MmLimitPerPrompt}; use crate::error::{BenchError, Result}; @@ -215,63 +217,63 @@ pub struct BenchConfig { } impl BenchConfig { - pub fn from_cli(cli: &Cli) -> Result { - if cli.burstiness <= 0.0 { + pub fn from_args(args: &BenchServeArgs) -> Result { + if args.burstiness <= 0.0 { return Err(BenchError::Config("Burstiness must be positive".into())); } - if cli.num_prompts == 0 { + if args.num_prompts == 0 { return Err(BenchError::Config( "--num-prompts must be at least 1".into(), )); } - if cli.request_rate <= 0.0 && !cli.request_rate.is_infinite() { + if args.request_rate <= 0.0 && !args.request_rate.is_infinite() { return Err(BenchError::Config( "--request-rate must be positive (or inf)".into(), )); } - if cli.max_model_len == Some(0) { + if args.max_model_len == Some(0) { return Err(BenchError::Config( "--max-model-len must be at least 1".into(), )); } - let base_url = cli.resolve_base_url(); - let api_url = cli.resolve_api_url(); + let base_url = args.resolve_base_url(); + let api_url = args.resolve_api_url(); - let extra_headers = cli.parse_headers()?; - let mut extra_body = cli.parse_extra_body()?; + let extra_headers = args.parse_headers()?; + let mut extra_body = args.parse_extra_body()?; // Merge sampling parameters into extra_body (matches Python behavior). // Python collects non-None sampling params and merges them UNDER extra_body, // meaning extra_body keys take precedence over sampling params. { let mut sampling_params = serde_json::Map::new(); - if let Some(v) = cli.top_p { + if let Some(v) = args.top_p { sampling_params.insert("top_p".into(), serde_json::json!(v)); } - if let Some(v) = cli.top_k { + if let Some(v) = args.top_k { sampling_params.insert("top_k".into(), serde_json::json!(v)); } - if let Some(v) = cli.min_p { + if let Some(v) = args.min_p { sampling_params.insert("min_p".into(), serde_json::json!(v)); } - if let Some(v) = cli.temperature { + if let Some(v) = args.temperature { sampling_params.insert("temperature".into(), serde_json::json!(v)); } - if let Some(v) = cli.frequency_penalty { + if let Some(v) = args.frequency_penalty { sampling_params.insert("frequency_penalty".into(), serde_json::json!(v)); } - if let Some(v) = cli.presence_penalty { + if let Some(v) = args.presence_penalty { sampling_params.insert("presence_penalty".into(), serde_json::json!(v)); } - if let Some(v) = cli.repetition_penalty { + if let Some(v) = args.repetition_penalty { sampling_params.insert("repetition_penalty".into(), serde_json::json!(v)); } if !sampling_params.is_empty() { - if !cli.backend.is_openai_compatible() { + if !args.backend.is_openai_compatible() { return Err(BenchError::Config( "Sampling parameters are only supported by openai-compatible backends." .into(), @@ -299,7 +301,7 @@ impl BenchConfig { } // Parse metadata - let metadata = match &cli.metadata { + let metadata = match &args.metadata { None => None, Some(items) => { let mut pairs = Vec::new(); @@ -314,24 +316,24 @@ impl BenchConfig { }; // Parse goodput SLOs - let goodput = parse_goodput(&cli.goodput)?; + let goodput = parse_goodput(&args.goodput)?; // Parse ramp-up config - let ramp_up = parse_ramp_up(cli)?; + let ramp_up = parse_ramp_up(args)?; // Default percentile metrics based on backend type - let default_percentile_metrics = if cli.backend.is_pooling() { + let default_percentile_metrics = if args.backend.is_pooling() { "e2el" } else { "ttft,tpot,itl,e2el" }; let percentile_metrics_str = - cli.percentile_metrics.as_deref().unwrap_or(default_percentile_metrics); + args.percentile_metrics.as_deref().unwrap_or(default_percentile_metrics); let selected_percentile_metrics: Vec = percentile_metrics_str.split(',').map(|s| s.trim().to_string()).collect(); - let metric_percentiles = parse_percentiles(&cli.metric_percentiles, false)?; - let sweep_summary_percentiles = cli + let metric_percentiles = parse_percentiles(&args.metric_percentiles, false)?; + let sweep_summary_percentiles = args .sweep_summary_percentiles .as_deref() .map(|raw| parse_percentiles(raw, true)) @@ -344,38 +346,38 @@ impl BenchConfig { selected_percentiles.push(90.0); } - let tokenizer_id = if cli.skip_tokenizer_init { + let tokenizer_id = if args.skip_tokenizer_init { None } else { - Some(cli.tokenizer.clone().or_else(|| cli.model.clone()).unwrap_or_default()) + args.tokenizer.clone().or_else(|| args.model.clone()) }; // Resolve input/output lengths - let random_input_len = cli.resolved_random_input_len(); - let random_output_len = cli.resolved_random_output_len(); - let per_turn_input_len = cli.resolved_per_turn_input_len(); + let random_input_len = args.resolved_random_input_len(); + let random_output_len = args.resolved_random_output_len(); + let per_turn_input_len = args.resolved_per_turn_input_len(); // Normalized multi-turn turn counts (computed in validation block below, defaults // to num_turns if multi-turn mode is not active) - let mut multi_turn_min_turns = cli.multi_turn_num_turns; - let mut multi_turn_max_turns = cli.multi_turn_num_turns; + let mut multi_turn_min_turns = args.multi_turn_num_turns; + let mut multi_turn_max_turns = args.multi_turn_num_turns; // For random datasets with openai-compatible backends, default to ignore_eos. // Exception: multi-turn mode, where ignore_eos causes unbounded context growth // across turns. Multi-turn uses min_tokens instead for output length control. // Pooling backends don't generate tokens, so ignore_eos is irrelevant. - let ignore_eos = if cli.backend.is_pooling() { + let ignore_eos = if args.backend.is_pooling() { false } else { - cli.ignore_eos - || ((cli.dataset_name == DatasetName::Random - || cli.dataset_name == DatasetName::RandomMm) - && cli.backend.is_openai_compatible() - && !cli.multi_turn) + args.ignore_eos + || ((args.dataset_name == DatasetName::Random + || args.dataset_name == DatasetName::RandomMm) + && args.backend.is_openai_compatible() + && !args.multi_turn) }; // Pooling backends don't support multi-turn - if cli.backend.is_pooling() && cli.multi_turn { + if args.backend.is_pooling() && args.multi_turn { return Err(BenchError::Config( "Pooling/embedding backends do not support --multi-turn".into(), )); @@ -383,7 +385,7 @@ impl BenchConfig { // LoRA validation. Adapter names must be non-empty after trim; pooling // backends are out of scope (vLLM LoRA routing is for generative paths). - let lora_modules = match cli.lora_modules.as_ref() { + let lora_modules = match args.lora_modules.as_ref() { None => None, Some(names) => { if names.is_empty() { @@ -391,7 +393,7 @@ impl BenchConfig { "--lora-modules requires at least one adapter name".into(), )); } - if cli.backend.is_pooling() { + if args.backend.is_pooling() { return Err(BenchError::Config( "--lora-modules is not supported for pooling/embedding backends".into(), )); @@ -411,18 +413,18 @@ impl BenchConfig { }; // Random-MM validation and config parsing - let (random_mm_limit, random_mm_buckets) = if cli.dataset_name == DatasetName::RandomMm { - if cli.backend != BackendKind::OpenaiChat { + let (random_mm_limit, random_mm_buckets) = if args.dataset_name == DatasetName::RandomMm { + if args.backend != BackendKind::OpenaiChat { return Err(BenchError::Config( "Multi-modal content (images) is only supported on 'openai-chat' backend." .into(), )); } let limit = crate::datasets::random_mm::parse_limit_mm_per_prompt( - &cli.random_mm_limit_mm_per_prompt, + &args.random_mm_limit_mm_per_prompt, )?; let buckets = - crate::datasets::random_mm::parse_bucket_config(&cli.random_mm_bucket_config)?; + crate::datasets::random_mm::parse_bucket_config(&args.random_mm_bucket_config)?; (limit, buckets) } else { (MmLimitPerPrompt::default(), Vec::new()) @@ -432,18 +434,18 @@ impl BenchConfig { // sonnet (uses built-in Shakespeare's sonnets). // Range ratio (Python semantics: [len*(1-r), len*(1+r)], each r in [0,1)) - let random_range_ratio = RangeRatio::parse(&cli.random_range_ratio)?; + let random_range_ratio = RangeRatio::parse(&args.random_range_ratio)?; // Batched inputs only make sense for pooling backends (the generation // backends send one prompt per request). - if cli.random_batch_size == 0 { + if args.random_batch_size == 0 { return Err(BenchError::Config( "--random-batch-size must be at least 1".into(), )); } - if cli.random_batch_size > 1 - && !cli.backend.is_pooling() - && cli.dataset_name != DatasetName::RandomRerank + if args.random_batch_size > 1 + && !args.backend.is_pooling() + && args.dataset_name != DatasetName::RandomRerank { return Err(BenchError::Config( "--random-batch-size > 1 is only supported with embeddings/pooling backends".into(), @@ -451,16 +453,16 @@ impl BenchConfig { } // random-rerank validation (mirrors Python RandomDatasetForReranking) - let is_reranker = !cli.no_reranker; - if cli.dataset_name == DatasetName::RandomRerank { - if !cli.backend.is_pooling() { + let is_reranker = !args.no_reranker; + if args.dataset_name == DatasetName::RandomRerank { + if !args.backend.is_pooling() { return Err(BenchError::Config( "--dataset-name random-rerank requires an embeddings/pooling backend \ (e.g. --backend vllm-rerank)" .into(), )); } - if !is_reranker && (cli.num_prompts < 2 || cli.random_batch_size < 2) { + if !is_reranker && (args.num_prompts < 2 || args.random_batch_size < 2) { return Err(BenchError::Config( "--no-reranker requires --num-prompts > 1 and --random-batch-size > 1 \ (the query is folded into the first batch slot)" @@ -470,8 +472,8 @@ impl BenchConfig { } // Custom dataset validation - if cli.dataset_name == DatasetName::Custom { - match cli.dataset_path.as_deref() { + if args.dataset_name == DatasetName::Custom { + match args.dataset_path.as_deref() { None => { return Err(BenchError::Config( "--dataset-path is required for --dataset-name custom \ @@ -486,7 +488,7 @@ impl BenchConfig { } _ => {} } - if !cli.skip_chat_template { + if !args.skip_chat_template { eprintln!( "NOTE: client-side chat template rendering is not supported; custom \ dataset prompts are sent raw (equivalent to --skip-chat-template)." @@ -495,29 +497,29 @@ impl BenchConfig { } // Prefix repetition validation - if cli.dataset_name == DatasetName::PrefixRepetition { - if cli.prefix_repetition_num_prefixes == 0 { + if args.dataset_name == DatasetName::PrefixRepetition { + if args.prefix_repetition_num_prefixes == 0 { return Err(BenchError::Config( "--prefix-repetition-num-prefixes must be at least 1".into(), )); } - if cli.num_prompts < cli.prefix_repetition_num_prefixes { + if args.num_prompts < args.prefix_repetition_num_prefixes { return Err(BenchError::Config(format!( "--num-prompts ({}) must be >= --prefix-repetition-num-prefixes ({})", - cli.num_prompts, cli.prefix_repetition_num_prefixes + args.num_prompts, args.prefix_repetition_num_prefixes ))); } } // HF dataset validation - if cli.dataset_name == DatasetName::Hf && cli.dataset_path.is_none() { + if args.dataset_name == DatasetName::Hf && args.dataset_path.is_none() { return Err(BenchError::Config( "--dataset-path is required for --dataset-name hf \ (set to a HuggingFace dataset ID, e.g. 'allenai/WildChat-4.8M')" .into(), )); } - if let Some(len) = cli.hf_output_len + if let Some(len) = args.hf_output_len && len == 0 { return Err(BenchError::Config( @@ -526,13 +528,13 @@ impl BenchConfig { } // Multi-turn validation - if cli.multi_turn { - if cli.backend != BackendKind::OpenaiChat { + if args.multi_turn { + if args.backend != BackendKind::OpenaiChat { return Err(BenchError::Config( "--multi-turn requires --backend openai-chat".into(), )); } - if cli.multi_turn_num_turns == 0 { + if args.multi_turn_num_turns == 0 { return Err(BenchError::Config( "--multi-turn-num-turns must be at least 1".into(), )); @@ -541,18 +543,18 @@ impl BenchConfig { // Normalize and validate min/max turns. ShareGPT only consumes max_turns // (the loader walks all available turns up to the cap), so the // min/num/max coupling used for synthetic generation does not apply. - if cli.dataset_name == DatasetName::ShareGpt { - if cli.multi_turn_max_turns == 1 { + if args.dataset_name == DatasetName::ShareGpt { + if args.multi_turn_max_turns == 1 { return Err(BenchError::Config( "--multi-turn-max-turns must be at least 2 for ShareGPT multi-turn".into(), )); } } else { (multi_turn_min_turns, multi_turn_max_turns) = - match (cli.multi_turn_min_turns, cli.multi_turn_max_turns) { - (0, 0) => (cli.multi_turn_num_turns, cli.multi_turn_num_turns), - (m, 0) => (m, cli.multi_turn_num_turns), - (0, x) => (cli.multi_turn_num_turns, x), + match (args.multi_turn_min_turns, args.multi_turn_max_turns) { + (0, 0) => (args.multi_turn_num_turns, args.multi_turn_num_turns), + (m, 0) => (m, args.multi_turn_num_turns), + (0, x) => (args.multi_turn_num_turns, x), (m, x) => (m, x), }; if multi_turn_min_turns < 1 { @@ -575,8 +577,8 @@ impl BenchConfig { } // Validate prefix sharing ratios - let pg = cli.multi_turn_prefix_global_ratio; - let pc = cli.multi_turn_prefix_conversation_ratio; + let pg = args.multi_turn_prefix_global_ratio; + let pc = args.multi_turn_prefix_conversation_ratio; if !(0.0..=1.0).contains(&pg) { return Err(BenchError::Config( "--multi-turn-prefix-global-ratio must be in [0.0, 1.0]".into(), @@ -592,20 +594,20 @@ impl BenchConfig { "--multi-turn-prefix-global-ratio + --multi-turn-prefix-conversation-ratio must be < 1.0 (unique suffix required)".into(), )); } - if (pg > 0.0 || pc > 0.0) && cli.dataset_name != DatasetName::Random { + if (pg > 0.0 || pc > 0.0) && args.dataset_name != DatasetName::Random { return Err(BenchError::Config( "Prefix sharing (--multi-turn-prefix-global-ratio / --multi-turn-prefix-conversation-ratio) only works with --dataset-name random".into(), )); } } - if !(cli.steady_state_threshold > 0.0 && cli.steady_state_threshold <= 1.0) { + if !(args.steady_state_threshold > 0.0 && args.steady_state_threshold <= 1.0) { return Err(BenchError::Config(format!( "--steady-state-threshold must be in (0.0, 1.0], got {}", - cli.steady_state_threshold + args.steady_state_threshold ))); } - if let Some(mw) = cli.steady_state_min_window + if let Some(mw) = args.steady_state_min_window && mw < 0.0 { return Err(BenchError::Config(format!( @@ -613,122 +615,122 @@ impl BenchConfig { ))); } - if cli.profile_batch_threshold.is_some() && !cli.profile { + if args.profile_batch_threshold.is_some() && !args.profile { return Err(BenchError::Config( "--profile-batch-threshold requires --profile".into(), )); } - if cli.profile_duration <= 0.0 { + if args.profile_duration <= 0.0 { return Err(BenchError::Config( "--profile-duration must be positive".into(), )); } - if cli.profile_batch_threshold.is_none() && cli.profile_duration != 5.0 { + if args.profile_batch_threshold.is_none() && args.profile_duration != 5.0 { return Err(BenchError::Config( "--profile-duration requires --profile-batch-threshold".into(), )); } Ok(BenchConfig { - backend: cli.backend, + backend: args.backend, base_url, api_url, - model: cli.model.clone(), - model_name: cli.served_model_name.clone(), + model: args.model.clone(), + model_name: args.served_model_name.clone(), tokenizer_id, - tokenizer_mode: cli.tokenizer_mode.clone(), - trust_remote_code: cli.trust_remote_code, - skip_tokenizer_init: cli.skip_tokenizer_init, - dataset_name: cli.dataset_name, - dataset_path: cli.dataset_path.clone(), - max_model_len: cli.max_model_len, + tokenizer_mode: args.tokenizer_mode.clone(), + trust_remote_code: args.trust_remote_code, + skip_tokenizer_init: args.skip_tokenizer_init, + dataset_name: args.dataset_name, + dataset_path: args.dataset_path.clone(), + max_model_len: args.max_model_len, random_input_len, random_output_len, - random_prefix_len: cli.random_prefix_len, + random_prefix_len: args.random_prefix_len, random_range_ratio, - random_batch_size: cli.random_batch_size, + random_batch_size: args.random_batch_size, is_reranker, - custom_output_len: cli.output_len.map(|v| v as i64).unwrap_or(cli.custom_output_len), - prefix_repetition_prefix_len: cli.prefix_repetition_prefix_len, - prefix_repetition_suffix_len: cli.prefix_repetition_suffix_len, - prefix_repetition_num_prefixes: cli.prefix_repetition_num_prefixes, - prefix_repetition_output_len: cli + custom_output_len: args.output_len.map(|v| v as i64).unwrap_or(args.custom_output_len), + prefix_repetition_prefix_len: args.prefix_repetition_prefix_len, + prefix_repetition_suffix_len: args.prefix_repetition_suffix_len, + prefix_repetition_num_prefixes: args.prefix_repetition_num_prefixes, + prefix_repetition_output_len: args .output_len - .unwrap_or(cli.prefix_repetition_output_len), - random_cache_hit_fraction: cli.random_cache_hit_fraction, - random_cache_ratio: cli.random_cache_ratio, - sharegpt_output_len: cli.sharegpt_output_len, - sonnet_input_len: cli.sonnet_input_len, - sonnet_output_len: cli.sonnet_output_len, - sonnet_prefix_len: cli.sonnet_prefix_len, - no_oversample: cli.no_oversample, - disable_shuffle: cli.disable_shuffle, - num_prompts: cli.num_prompts, - request_rate: cli.request_rate, - burstiness: cli.burstiness, - max_concurrency: cli.max_concurrency, - steady_state_threshold: cli.steady_state_threshold, - steady_state_min_window: cli.steady_state_min_window, - no_steady_state: cli.no_steady_state, - disable_tqdm: cli.disable_tqdm, - num_warmups: cli.num_warmups, - profile: cli.profile, - profile_batch_threshold: cli.profile_batch_threshold, - profile_duration: cli.profile_duration, - save_result: cli.save_result, - save_detailed: cli.save_detailed, - append_result: cli.append_result, - result_dir: cli.result_dir.clone(), - result_filename: cli.result_filename.clone(), - seed: cli.seed, + .unwrap_or(args.prefix_repetition_output_len), + random_cache_hit_fraction: args.random_cache_hit_fraction, + random_cache_ratio: args.random_cache_ratio, + sharegpt_output_len: args.sharegpt_output_len, + sonnet_input_len: args.sonnet_input_len, + sonnet_output_len: args.sonnet_output_len, + sonnet_prefix_len: args.sonnet_prefix_len, + no_oversample: args.no_oversample, + disable_shuffle: args.disable_shuffle, + num_prompts: args.num_prompts, + request_rate: args.request_rate, + burstiness: args.burstiness, + max_concurrency: args.max_concurrency, + steady_state_threshold: args.steady_state_threshold, + steady_state_min_window: args.steady_state_min_window, + no_steady_state: args.no_steady_state, + disable_tqdm: args.disable_tqdm, + num_warmups: args.num_warmups, + profile: args.profile, + profile_batch_threshold: args.profile_batch_threshold, + profile_duration: args.profile_duration, + save_result: args.save_result, + save_detailed: args.save_detailed, + append_result: args.append_result, + result_dir: args.result_dir.clone(), + result_filename: args.result_filename.clone(), + seed: args.seed, ignore_eos, - insecure: cli.insecure, + insecure: args.insecure, selected_percentile_metrics, selected_percentiles, sweep_summary_percentiles, - label: cli.label.clone(), - logprobs: cli.logprobs, - request_id_prefix: cli.get_request_id_prefix(), - ready_check_timeout_sec: cli.ready_check_timeout_sec, + label: args.label.clone(), + logprobs: args.logprobs, + request_id_prefix: args.get_request_id_prefix(), + ready_check_timeout_sec: args.ready_check_timeout_sec, extra_headers, extra_body, metadata, - dry_run: cli.dry_run, + dry_run: args.dry_run, goodput, ramp_up, - multi_turn: cli.multi_turn, - multi_turn_num_turns: cli.multi_turn_num_turns, + multi_turn: args.multi_turn, + multi_turn_num_turns: args.multi_turn_num_turns, multi_turn_min_turns, multi_turn_max_turns, - sharegpt_multi_turn_max_turns: if cli.multi_turn - && cli.dataset_name == DatasetName::ShareGpt - && cli.multi_turn_max_turns != 0 + sharegpt_multi_turn_max_turns: if args.multi_turn + && args.dataset_name == DatasetName::ShareGpt + && args.multi_turn_max_turns != 0 { - Some(cli.multi_turn_max_turns) + Some(args.multi_turn_max_turns) } else { None }, per_turn_input_len, - multi_turn_concurrency: cli.multi_turn_concurrency, - multi_turn_delay_ms: cli.multi_turn_delay_ms, - multi_turn_prefix_global_ratio: cli.multi_turn_prefix_global_ratio, - multi_turn_prefix_conversation_ratio: cli.multi_turn_prefix_conversation_ratio, - speed_bench_config: cli.speed_bench_config, - speed_bench_category: cli.speed_bench_category.clone(), - speed_bench_max_input_len: cli.speed_bench_max_input_len, - hf_split: cli.hf_split.clone(), - hf_subset: cli.hf_subset.clone(), - hf_output_len: cli.hf_output_len, - hf_text_column: cli.hf_text_column.clone(), - reset_prefix_cache: cli.reset_prefix_cache, - prompt_token_ids: cli.prompt_token_ids, - random_mm_base_items_per_request: cli.random_mm_base_items_per_request, - random_mm_num_mm_items_range_ratio: cli.random_mm_num_mm_items_range_ratio, + multi_turn_concurrency: args.multi_turn_concurrency, + multi_turn_delay_ms: args.multi_turn_delay_ms, + multi_turn_prefix_global_ratio: args.multi_turn_prefix_global_ratio, + multi_turn_prefix_conversation_ratio: args.multi_turn_prefix_conversation_ratio, + speed_bench_config: args.speed_bench_config, + speed_bench_category: args.speed_bench_category.clone(), + speed_bench_max_input_len: args.speed_bench_max_input_len, + hf_split: args.hf_split.clone(), + hf_subset: args.hf_subset.clone(), + hf_output_len: args.hf_output_len, + hf_text_column: args.hf_text_column.clone(), + reset_prefix_cache: args.reset_prefix_cache, + prompt_token_ids: args.prompt_token_ids, + random_mm_base_items_per_request: args.random_mm_base_items_per_request, + random_mm_num_mm_items_range_ratio: args.random_mm_num_mm_items_range_ratio, random_mm_limit, random_mm_buckets, - enable_multimodal_chat: cli.enable_multimodal_chat, + enable_multimodal_chat: args.enable_multimodal_chat, lora_modules, - lora_assignment: cli.lora_assignment, + lora_assignment: args.lora_assignment, }) } } @@ -811,17 +813,17 @@ fn parse_goodput(goodput_args: &Option>) -> Result { Ok(config) } -fn parse_ramp_up(cli: &Cli) -> Result> { - let strategy = match cli.ramp_up_strategy { +fn parse_ramp_up(args: &BenchServeArgs) -> Result> { + let strategy = match args.ramp_up_strategy { None => return Ok(None), Some(s) => s, }; - let start_rps = cli.ramp_up_start_rps.ok_or_else(|| { + let start_rps = args.ramp_up_start_rps.ok_or_else(|| { BenchError::Config("--ramp-up-start-rps is required when --ramp-up-strategy is set".into()) })?; - let end_rps = cli.ramp_up_end_rps.ok_or_else(|| { + let end_rps = args.ramp_up_end_rps.ok_or_else(|| { BenchError::Config("--ramp-up-end-rps is required when --ramp-up-strategy is set".into()) })?; @@ -843,7 +845,21 @@ mod tests { use clap::Parser; use super::*; - use crate::cli::Cli; + use crate::cli::BenchServeArgs; + + #[derive(Parser)] + struct TestCli { + #[command(flatten)] + args: BenchServeArgs, + } + + fn parse_args(args: I) -> BenchServeArgs + where + I: IntoIterator, + T: Into + Clone, + { + TestCli::parse_from(args).args + } fn base_multi_turn_args() -> Vec<&'static str> { vec![ @@ -859,8 +875,8 @@ mod tests { #[test] fn test_prefix_sharing_defaults_to_zero() { let args = base_multi_turn_args(); - let cli = Cli::parse_from(args); - let config = BenchConfig::from_cli(&cli).unwrap(); + let args = parse_args(args); + let config = BenchConfig::from_args(&args).unwrap(); assert_eq!(config.multi_turn_prefix_global_ratio, 0.0); assert_eq!(config.multi_turn_prefix_conversation_ratio, 0.0); } @@ -874,8 +890,8 @@ mod tests { "--multi-turn-prefix-conversation-ratio", "0.8", ]); - let cli = Cli::parse_from(args); - let config = BenchConfig::from_cli(&cli).unwrap(); + let args = parse_args(args); + let config = BenchConfig::from_args(&args).unwrap(); assert!((config.multi_turn_prefix_global_ratio - 0.1).abs() < 1e-10); assert!((config.multi_turn_prefix_conversation_ratio - 0.8).abs() < 1e-10); } @@ -889,8 +905,8 @@ mod tests { "--multi-turn-prefix-conversation-ratio", "0.6", ]); - let cli = Cli::parse_from(args); - assert!(BenchConfig::from_cli(&cli).is_err()); + let args = parse_args(args); + assert!(BenchConfig::from_args(&args).is_err()); } #[test] @@ -902,16 +918,16 @@ mod tests { "--multi-turn-prefix-conversation-ratio", "0.5", ]); - let cli = Cli::parse_from(args); - assert!(BenchConfig::from_cli(&cli).is_err()); + let args = parse_args(args); + assert!(BenchConfig::from_args(&args).is_err()); } #[test] fn test_prefix_sharing_out_of_range_fails() { let mut args = base_multi_turn_args(); args.extend(["--multi-turn-prefix-global-ratio", "1.5"]); - let cli = Cli::parse_from(args); - assert!(BenchConfig::from_cli(&cli).is_err()); + let args = parse_args(args); + assert!(BenchConfig::from_args(&args).is_err()); } #[test] @@ -928,8 +944,8 @@ mod tests { "--multi-turn-prefix-global-ratio", "0.1", ]; - let cli = Cli::parse_from(args); - assert!(BenchConfig::from_cli(&cli).is_err()); + let args = parse_args(args); + assert!(BenchConfig::from_args(&args).is_err()); } #[test] @@ -944,8 +960,8 @@ mod tests { "--dataset-name", "sharegpt", ]; - let cli = Cli::parse_from(args); - let config = BenchConfig::from_cli(&cli).unwrap(); + let args = parse_args(args); + let config = BenchConfig::from_args(&args).unwrap(); assert_eq!(config.multi_turn_max_turns, 3); assert_eq!(config.sharegpt_multi_turn_max_turns, None); @@ -968,8 +984,8 @@ mod tests { "--multi-turn-max-turns", "2", ]; - let cli = Cli::parse_from(args); - let config = BenchConfig::from_cli(&cli).unwrap(); + let args = parse_args(args); + let config = BenchConfig::from_args(&args).unwrap(); assert_eq!(config.sharegpt_multi_turn_max_turns, Some(2)); } @@ -987,8 +1003,8 @@ mod tests { "--multi-turn-max-turns", "1", ]; - let cli = Cli::parse_from(args); - let err = BenchConfig::from_cli(&cli).unwrap_err().to_string(); + let args = parse_args(args); + let err = BenchConfig::from_args(&args).unwrap_err().to_string(); assert!( err.contains("at least 2 for ShareGPT"), "expected ShareGPT-specific error, got: {err}" @@ -1009,8 +1025,8 @@ mod tests { "--multi-turn-max-turns", "20", ]; - let cli = Cli::parse_from(args); - let config = BenchConfig::from_cli(&cli).unwrap(); + let args = parse_args(args); + let config = BenchConfig::from_args(&args).unwrap(); assert_eq!(config.sharegpt_multi_turn_max_turns, Some(20)); } @@ -1018,8 +1034,8 @@ mod tests { #[test] fn test_sweep_summary_percentiles_default_empty() { let args = base_multi_turn_args(); - let cli = Cli::parse_from(args); - let config = BenchConfig::from_cli(&cli).unwrap(); + let args = parse_args(args); + let config = BenchConfig::from_args(&args).unwrap(); assert!(config.sweep_summary_percentiles.is_empty()); assert_eq!(config.selected_percentiles, vec![99.0, 90.0]); @@ -1034,8 +1050,8 @@ mod tests { "--sweep-summary-percentiles", "90,95,90", ]); - let cli = Cli::parse_from(args); - let config = BenchConfig::from_cli(&cli).unwrap(); + let args = parse_args(args); + let config = BenchConfig::from_args(&args).unwrap(); assert_eq!(config.sweep_summary_percentiles, vec![90.0, 95.0]); assert_eq!(config.selected_percentiles, vec![99.0, 95.0, 90.0]); @@ -1045,8 +1061,8 @@ mod tests { fn test_invalid_sweep_summary_percentile_fails() { let mut args = base_multi_turn_args(); args.extend(["--sweep-summary-percentiles", "101"]); - let cli = Cli::parse_from(args); - assert!(BenchConfig::from_cli(&cli).is_err()); + let args = parse_args(args); + assert!(BenchConfig::from_args(&args).is_err()); } #[test] @@ -1058,12 +1074,20 @@ mod tests { "--max-model-len", "4096", ]; - let cli = Cli::parse_from(args); - let config = BenchConfig::from_cli(&cli).unwrap(); + let args = parse_args(args); + let config = BenchConfig::from_args(&args).unwrap(); assert_eq!(config.max_model_len, Some(4096)); } + #[test] + fn test_tokenizer_id_deferred_when_model_is_unspecified() { + let args = parse_args(["vllm-bench"]); + let config = BenchConfig::from_args(&args).unwrap(); + + assert_eq!(config.tokenizer_id, None); + } + #[test] fn test_zero_max_model_len_fails() { let args = vec![ @@ -1073,9 +1097,9 @@ mod tests { "--max-model-len", "0", ]; - let cli = Cli::parse_from(args); + let args = parse_args(args); - assert!(BenchConfig::from_cli(&cli).is_err()); + assert!(BenchConfig::from_args(&args).is_err()); } #[test] fn test_range_ratio_parse_float() { diff --git a/rust/src/bench/src/hub.rs b/rust/src/bench/src/hub.rs index 0e1b3e962c8..ab51ccdcf08 100644 --- a/rust/src/bench/src/hub.rs +++ b/rust/src/bench/src/hub.rs @@ -40,8 +40,11 @@ impl HubRepo { .build() .map_err(|e| format!("Failed to build download runtime: {e}"))?; rt.block_on(async move { - let api = hf_hub::api::tokio::Api::new() - .map_err(|e| format!("Failed to init HF API: {e}"))?; + let mut builder = hf_hub::api::tokio::ApiBuilder::from_env(); + if let Ok(token) = std::env::var("HF_TOKEN") { + builder = builder.with_token(Some(token)); + } + let api = builder.build().map_err(|e| format!("Failed to init HF API: {e}"))?; api.repo(repo).get(&filename).await.map_err(|e| format!("{e}")) }) }) diff --git a/rust/src/bench/src/lib.rs b/rust/src/bench/src/lib.rs new file mode 100644 index 00000000000..13f252cee33 --- /dev/null +++ b/rust/src/bench/src/lib.rs @@ -0,0 +1,86 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +mod backends; +mod benchmark; +mod cli; +mod compare; +mod config; +mod datasets; +mod error; +mod hub; +mod metrics; +mod multi_run; +mod multi_turn; +mod output; +mod rate_control; +mod ready_checker; +mod sweep; +mod tiktoken; +mod tokenizer; + +use anyhow::Context; + +pub use cli::{ + BackendKind, BenchServeArgs, DatasetName, LoraAssignment, RampUpStrategy, SpeedBenchConfig, +}; +use config::BenchConfig; + +/// Prepare process-wide resources for a benchmark run. +pub fn prepare_process() { + // Raise the open-file soft limit to the hard limit. High-concurrency + // benchmarks (1024+ requests) easily exceed the default 1024 fd soft limit. + if let Ok(new) = rlimit::increase_nofile_limit(u64::MAX) + && new > 1024 + { + eprintln!("Open-file limit: {new}"); + } +} + +/// Run the online serving benchmark. +pub async fn run(args: BenchServeArgs) -> anyhow::Result<()> { + // --- Compare mode: no server needed, just diff two JSON files --- + if let Some(ref files) = args.compare { + return compare::compare_results(&files[0], &files[1]).context("Comparison failed"); + } + + let config = BenchConfig::from_args(&args).context("Configuration error")?; + + async { + if config.multi_turn { + if let Some(ref sweep_mc) = args.sweep_max_concurrency { + // --- Sweep over concurrency in multi-turn mode --- + let values = sweep::parse_concurrency_values(sweep_mc) + .context("Invalid --sweep-max-concurrency")?; + sweep::run_multi_turn_concurrency_sweep( + &config, + &values, + args.sweep_num_prompts_factor, + ) + .await?; + } else { + // --- Single multi-turn conversation benchmark --- + multi_turn::run_multi_turn_benchmark(&config).await?; + } + } else if let Some(ref sweep_mc) = args.sweep_max_concurrency { + // --- Sweep over max-concurrency --- + let values = sweep::parse_concurrency_values(sweep_mc) + .context("Invalid --sweep-max-concurrency")?; + sweep::run_concurrency_sweep(&config, &values, args.sweep_num_prompts_factor).await?; + } else if let Some(ref sweep_rate) = args.sweep_request_rate { + // --- Sweep over request-rate --- + let values = + sweep::parse_rate_values(sweep_rate).context("Invalid --sweep-request-rate")?; + sweep::run_rate_sweep(&config, &values).await?; + } else if args.num_runs > 1 { + // --- Multi-run with statistical aggregation --- + multi_run::run_multi(&config, args.num_runs).await?; + } else { + // --- Normal single benchmark --- + benchmark::run_benchmark(&config).await?; + } + anyhow::Ok(()) + } + .await + .context("Benchmark failed") +} diff --git a/rust/src/bench/src/main.rs b/rust/src/bench/src/main.rs index dd37ec56db3..919e1e0b369 100644 --- a/rust/src/bench/src/main.rs +++ b/rust/src/bench/src/main.rs @@ -1,92 +1,32 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright contributors to the vLLM project -mod backends; -mod benchmark; -mod cli; -mod compare; -mod config; -mod datasets; -mod error; -mod hub; -mod metrics; -mod multi_run; -mod multi_turn; -mod output; -mod rate_control; -mod ready_checker; -mod sweep; -mod tiktoken; -mod tokenizer; - #[cfg(not(target_env = "msvc"))] #[global_allocator] static GLOBAL: mimalloc::MiMalloc = mimalloc::MiMalloc; use anyhow::Context; use clap::Parser; -use cli::Cli; -use config::BenchConfig; + +#[derive(Parser)] +#[command( + name = "vllm-bench", + about = "Benchmark online serving throughput", + version +)] +struct Cli { + #[command(flatten)] + args: vllm_bench::BenchServeArgs, +} fn main() -> anyhow::Result<()> { - // Raise the open-file soft limit to the hard limit. High-concurrency - // benchmarks (1024+ requests) easily exceed the default 1024 fd soft limit. - if let Ok(new) = rlimit::increase_nofile_limit(u64::MAX) - && new > 1024 - { - eprintln!("Open-file limit: {new}"); - } - let cli = Cli::parse(); - - // --- Compare mode: no server needed, just diff two JSON files --- - if let Some(ref files) = cli.compare { - return compare::compare_results(&files[0], &files[1]).context("Comparison failed"); - } - - let config = BenchConfig::from_cli(&cli).context("Configuration error")?; + vllm_bench::prepare_process(); let runtime = tokio::runtime::Builder::new_multi_thread() .enable_all() .build() - .expect("Failed to build tokio runtime"); + .context("Failed to build tokio runtime")?; - runtime - .block_on(async { - if config.multi_turn { - if let Some(ref sweep_mc) = cli.sweep_max_concurrency { - // --- Sweep over concurrency in multi-turn mode --- - let values = sweep::parse_concurrency_values(sweep_mc) - .context("Invalid --sweep-max-concurrency")?; - sweep::run_multi_turn_concurrency_sweep( - &config, - &values, - cli.sweep_num_prompts_factor, - ) - .await?; - } else { - // --- Single multi-turn conversation benchmark --- - multi_turn::run_multi_turn_benchmark(&config).await?; - } - } else if let Some(ref sweep_mc) = cli.sweep_max_concurrency { - // --- Sweep over max-concurrency --- - let values = sweep::parse_concurrency_values(sweep_mc) - .context("Invalid --sweep-max-concurrency")?; - sweep::run_concurrency_sweep(&config, &values, cli.sweep_num_prompts_factor) - .await?; - } else if let Some(ref sweep_rate) = cli.sweep_request_rate { - // --- Sweep over request-rate --- - let values = - sweep::parse_rate_values(sweep_rate).context("Invalid --sweep-request-rate")?; - sweep::run_rate_sweep(&config, &values).await?; - } else if cli.num_runs > 1 { - // --- Multi-run with statistical aggregation --- - multi_run::run_multi(&config, cli.num_runs).await?; - } else { - // --- Normal single benchmark --- - benchmark::run_benchmark(&config).await?; - } - anyhow::Ok(()) - }) - .context("Benchmark failed") + runtime.block_on(vllm_bench::run(cli.args)) } diff --git a/rust/src/cmd/Cargo.toml b/rust/src/cmd/Cargo.toml index 030d4c6d116..a6955059c26 100644 --- a/rust/src/cmd/Cargo.toml +++ b/rust/src/cmd/Cargo.toml @@ -29,6 +29,7 @@ tokio-util.workspace = true tracing.workspace = true tracing-subscriber.workspace = true uuid.workspace = true +vllm-bench.workspace = true vllm-chat.workspace = true vllm-engine-core-client.workspace = true vllm-managed-engine.workspace = true diff --git a/rust/src/cmd/src/cli.rs b/rust/src/cmd/src/cli.rs index 314765bab26..454a397ad05 100644 --- a/rust/src/cmd/src/cli.rs +++ b/rust/src/cmd/src/cli.rs @@ -79,13 +79,23 @@ impl Cli { } /// Supported top-level CLI commands. -#[derive(Debug, Subcommand, PartialEq, Eq)] +#[derive(Debug, Subcommand)] pub enum Command { /// Run the Rust OpenAI frontend as a Python-supervised worker. Frontend(FrontendArgs), /// Launch a managed Python headless engine, then run the Rust OpenAI /// frontend. Serve(ServeArgs), + /// Run vLLM benchmarks. + #[command(subcommand)] + Bench(BenchCommand), +} + +/// Supported benchmark commands. +#[derive(Debug, Subcommand)] +pub enum BenchCommand { + /// Benchmark online serving throughput. + Serve(vllm_bench::BenchServeArgs), } /// A JSON-encoded list of strings, matching Python's `json.loads` CLI type for diff --git a/rust/src/cmd/src/cli/tests.rs b/rust/src/cmd/src/cli/tests.rs index ee0d3ebf048..d8ec9792b7e 100644 --- a/rust/src/cmd/src/cli/tests.rs +++ b/rust/src/cmd/src/cli/tests.rs @@ -5,7 +5,27 @@ use expect_test::expect; use vllm_engine_core_client::TransportMode; use vllm_server::{Config, HttpListenerMode, ParserSelection, RendererSelection}; -use super::{Cli, Command}; +use super::{BenchCommand, Cli, Command}; + +#[test] +fn bench_serve_args_parse_without_managed_engine_repartition() { + let cli = Cli::try_parse_from([ + "vllm-rs", + "bench", + "serve", + "--backend", + "openai-chat", + "--request-rate", + "inf", + ]) + .unwrap(); + + let Command::Bench(BenchCommand::Serve(args)) = cli.command else { + panic!("expected bench serve args"); + }; + assert_eq!(args.backend, vllm_bench::BackendKind::OpenaiChat); + assert!(args.request_rate.is_infinite()); +} #[test] fn serve_args_forward_python_flags_with_separator() { diff --git a/rust/src/cmd/src/main.rs b/rust/src/cmd/src/main.rs index 9637a3be578..69314678103 100644 --- a/rust/src/cmd/src/main.rs +++ b/rust/src/cmd/src/main.rs @@ -12,7 +12,7 @@ use tokio_util::sync::CancellationToken; use tracing::{info, warn}; use vllm_managed_engine::ManagedEngineHandle; -use crate::cli::{Cli, Command}; +use crate::cli::{BenchCommand, Cli, Command}; #[global_allocator] static GLOBAL: mimalloc::MiMalloc = mimalloc::MiMalloc; @@ -100,6 +100,10 @@ fn main() -> Result<()> { async fn async_main(cli: Cli) -> Result<()> { match cli.command { Command::Frontend(args) => vllm_server::serve(args.into_config(), shutdown_signal()).await, + Command::Bench(BenchCommand::Serve(bench_args)) => { + vllm_bench::prepare_process(); + vllm_bench::run(bench_args).await + } Command::Serve(args) => { let handshake_port = args.managed_engine.resolve_handshake_port()?; diff --git a/vllm/entrypoints/cli/benchmark/serve.py b/vllm/entrypoints/cli/benchmark/serve.py index 188afd6c703..41a65273ba8 100644 --- a/vllm/entrypoints/cli/benchmark/serve.py +++ b/vllm/entrypoints/cli/benchmark/serve.py @@ -1,11 +1,68 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import argparse +import os +import sys +from pathlib import Path -from vllm.benchmarks.serve import add_cli_args, main +from vllm.benchmarks.serve import add_cli_args +from vllm.benchmarks.serve import main as python_main from vllm.entrypoints.cli.benchmark.base import BenchmarkSubcommandBase +from vllm.logger import init_logger from vllm.utils.argparse_utils import FlexibleArgumentParser +logger = init_logger(__name__) +_RUST_CLI_PATH = Path(__file__).resolve().parents[3] / "vllm-rs" +_RUST_SUPPORTED_DATASETS = frozenset( + { + "custom", + "hf", + "prefix_repetition", + "random", + "random-mm", + "random-rerank", + "sharegpt", + "sonnet", + "speed_bench", + } +) +_RUST_SUPPORTED_BACKENDS = frozenset( + { + "openai", + "openai-chat", + "openai-embeddings", + "openai-embeddings-chat", + "vllm", + "vllm-pooling", + "vllm-rerank", + } +) + + +def _rust_unsupported_reason(args: argparse.Namespace) -> str | None: + if args.dataset_name not in _RUST_SUPPORTED_DATASETS: + return f"dataset {args.dataset_name!r} is not supported by the Rust benchmark" + if args.backend not in _RUST_SUPPORTED_BACKENDS: + return f"backend {args.backend!r} is not supported by the Rust benchmark" + return None + + +def _maybe_exec_rust_bench(args: argparse.Namespace) -> None: + if reason := _rust_unsupported_reason(args): + logger.info("Using Python benchmark: %s.", reason) + return + + if not _RUST_CLI_PATH.is_file(): + logger.warning( + "Rust benchmark binary not found at %s; falling back to Python.", + _RUST_CLI_PATH, + ) + return + + rust_cli = str(_RUST_CLI_PATH) + logger.info("Delegating `vllm bench serve` to Rust binary at %s.", rust_cli) + os.execv(rust_cli, [rust_cli, "bench", "serve", *sys.argv[3:]]) + class BenchmarkServingSubcommand(BenchmarkSubcommandBase): """The `serve` subcommand for `vllm bench`.""" @@ -19,4 +76,5 @@ class BenchmarkServingSubcommand(BenchmarkSubcommandBase): @staticmethod def cmd(args: argparse.Namespace) -> None: - main(args) + _maybe_exec_rust_bench(args) + python_main(args) From 818cf61e9149c3fdff302cddbb2d026090160aea Mon Sep 17 00:00:00 2001 From: Reid <61492567+reidliu41@users.noreply.github.com> Date: Mon, 20 Jul 2026 17:39:13 +0800 Subject: [PATCH 47/51] [Rust Frontend] Fix macro-based content format detection (#49042) Signed-off-by: reidliu41 --- rust/src/chat/src/renderer/hf/format.rs | 17 ++++++++++++++--- rust/src/chat/src/renderer/hf/mod.rs | 20 ++++++++++++++++++++ 2 files changed, 34 insertions(+), 3 deletions(-) diff --git a/rust/src/chat/src/renderer/hf/format.rs b/rust/src/chat/src/renderer/hf/format.rs index f419b9e3a0c..e8e9f6412bb 100644 --- a/rust/src/chat/src/renderer/hf/format.rs +++ b/rust/src/chat/src/renderer/hf/format.rs @@ -236,9 +236,10 @@ fn has_content_item_loop(root: &Stmt<'_>) -> bool { loops.into_iter().any(|loop_ast| { matches!(loop_ast.target, Expr::Var(_)) - && message_varnames - .iter() - .any(|varname| is_var_or_elems_access(&loop_ast.iter, varname, Some("content"))) + && (is_var_access(&loop_ast.iter, "content") + || message_varnames.iter().any(|varname| { + is_var_or_elems_access(&loop_ast.iter, varname, Some("content")) + })) }) } @@ -315,6 +316,16 @@ mod tests { ); } + #[test] + fn detects_openai_template_with_content_parameter_loop() { + assert_eq!( + detect( + "{% macro render(content) %}{% for item in content %}{{ item }}{% endfor %}{% endmacro %}{% for message in messages %}{{ render(message.content) }}{% endfor %}" + ), + ChatTemplateContentFormat::OpenAi + ); + } + #[test] fn detects_openai_template_with_messages_alias() { assert_eq!( diff --git a/rust/src/chat/src/renderer/hf/mod.rs b/rust/src/chat/src/renderer/hf/mod.rs index 7ba2a2ccc03..cfbc3f924a3 100644 --- a/rust/src/chat/src/renderer/hf/mod.rs +++ b/rust/src/chat/src/renderer/hf/mod.rs @@ -1309,6 +1309,26 @@ mod tests { .assert_eq(&rendered); } + #[test] + fn qwen35_template_auto_detects_openai_multimodal_content() { + let mut request = image_request(); + request.chat_options.generation_prompt_mode = GenerationPromptMode::NoGenerationPrompt; + + let rendered = render_mm( + QWEN3_5_0_8B_TEMPLATE, + &request, + ChatTemplateContentFormatOption::Auto, + ) + .unwrap(); + + expect![[r#" + Text( + "<|im_start|>user\na<|vision_start|><|image_pad|><|vision_end|>b<|im_end|>\n", + ) + "#]] + .assert_debug_eq(&rendered.prompt); + } + #[test] fn qwen35_template_renders_closed_empty_reasoning_span_when_thinking_disabled() { let mut request = sample_request(vec![ChatMessage::text(ChatRole::User, "hello")]); From 47d0597ca2259096f18c52029a47cf10f9d35e70 Mon Sep 17 00:00:00 2001 From: Lena Onyshchenko Date: Mon, 20 Jul 2026 12:44:25 +0300 Subject: [PATCH 48/51] [Misc][Docs] Fix broken csrc kernel links in fusions doc (#47211) Signed-off-by: oonyshch Co-authored-by: Cursor --- docs/design/fusions.md | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/design/fusions.md b/docs/design/fusions.md index c9991f75cdb..002984c113a 100644 --- a/docs/design/fusions.md +++ b/docs/design/fusions.md @@ -306,7 +306,7 @@ Supported quantization scheme/hardware combinations: - Pass: [`vllm/compilation/passes/fusion/rms_quant_fusion.py`](https://github.com/vllm-project/vllm/blob/main/vllm/compilation/passes/fusion/rms_quant_fusion.py) - ROCm AITER pass: [`vllm/compilation/passes/fusion/rocm_aiter_fusion.py`](https://github.com/vllm-project/vllm/blob/main/vllm/compilation/passes/fusion/rocm_aiter_fusion.py) -- CUDA/HIP kernels: [`csrc/layernorm_quant_kernels.cu`](https://github.com/vllm-project/vllm/blob/main/csrc/layernorm_quant_kernels.cu) +- CUDA/HIP kernels: [`csrc/libtorch_stable/layernorm_quant_kernels.cu`](https://github.com/vllm-project/vllm/blob/main/csrc/libtorch_stable/layernorm_quant_kernels.cu) ### SiLU+Mul + Quantization (`fuse_act_quant`) @@ -332,7 +332,7 @@ Supported quantization scheme/hardware combinations: - Pass: [`vllm/compilation/passes/fusion/act_quant_fusion.py`](https://github.com/vllm-project/vllm/blob/main/vllm/compilation/passes/fusion/act_quant_fusion.py) - ROCm AITER pass: [`vllm/compilation/passes/fusion/rocm_aiter_fusion.py`](https://github.com/vllm-project/vllm/blob/main/vllm/compilation/passes/fusion/rocm_aiter_fusion.py) - CUDA/HIP kernels: [`csrc/quantization/`](https://github.com/vllm-project/vllm/blob/main/csrc/quantization/) -- Fused SiLU+Mul+BlockQuant kernel: [`csrc/quantization/fused_kernels/fused_silu_mul_block_quant.cu`](https://github.com/vllm-project/vllm/blob/main/csrc/quantization/fused_kernels/fused_silu_mul_block_quant.cu) +- Fused SiLU+Mul+BlockQuant kernel: [`csrc/libtorch_stable/quantization/fused_kernels/fused_silu_mul_block_quant.cu`](https://github.com/vllm-project/vllm/blob/main/csrc/libtorch_stable/quantization/fused_kernels/fused_silu_mul_block_quant.cu) ### RMSNorm + Padding (`fuse_act_padding`) From d835ad572ca76d8ed6308e81750813c814fb3343 Mon Sep 17 00:00:00 2001 From: Salt Sato Date: Mon, 20 Jul 2026 11:00:06 +0100 Subject: [PATCH 49/51] [Bugfix][Rust Frontend] Map missing prompt logprobs for single-token prompts in chat and raw generate (#49111) Signed-off-by: Feathbow --- .../server/src/routes/inference/generate.rs | 63 ++++++++++++++++--- .../src/routes/openai/chat_completions.rs | 13 ++-- .../server/src/routes/openai/completions.rs | 26 +------- .../src/routes/openai/utils/logprobs.rs | 45 ++++++++----- 4 files changed, 94 insertions(+), 53 deletions(-) diff --git a/rust/src/server/src/routes/inference/generate.rs b/rust/src/server/src/routes/inference/generate.rs index 6fcdb427ccf..e176938db40 100644 --- a/rust/src/server/src/routes/inference/generate.rs +++ b/rust/src/server/src/routes/inference/generate.rs @@ -238,13 +238,18 @@ fn collect_generate( None }; let prompt_logprobs = if include_prompt_logprobs { - let prompt_logprobs = collected.prompt_logprobs.as_ref().ok_or_else(|| { - ApiError::server_error( - "raw generate response requested prompt_logprobs but generation returned none" - .to_string(), - ) - })?; - Some(raw_prompt_logprobs_to_maps(prompt_logprobs)) + match collected.prompt_logprobs.as_ref() { + Some(prompt_logprobs) => Some(raw_prompt_logprobs_to_maps(prompt_logprobs)), + // A single-token prompt has no scored positions; same mapping + // as /v1/completions. + None if collected.prompt_token_ids.len() == 1 => Some(vec![None]), + None => { + return Err(ApiError::server_error( + "raw generate response requested prompt_logprobs but generation returned none" + .to_string(), + )); + } + } } else { None }; @@ -472,4 +477,48 @@ mod tests { Some(2) ); } + + #[test] + fn collect_generate_maps_prompt_logprobs_for_single_token_prompt() { + let output_without_payload = |prompt_token_ids: Vec| CollectedGenerateOutput { + request_id: "raw-1".to_string(), + prompt_logprobs: None, + token_ids: vec![3], + logprobs: None, + finish_reason: FinishReason::stop_eos(), + usage: vllm_llm::TokenUsage { + prompt_token_count: prompt_token_ids.len(), + output_token_count: 1, + cached_token_count: 0, + }, + kv_transfer_params: None, + ec_transfer_params: None, + prompt_token_ids, + }; + + let response = collect_generate( + output_without_payload(vec![9707]), + "raw-1".to_string(), + ApiServerOptions::default(), + ResponseOptions { + include_prompt_logprobs: true, + ..Default::default() + }, + ) + .expect("single-token prompt without payload maps to [None]"); + let prompt_logprobs = response.prompt_logprobs.expect("prompt logprobs present"); + assert_eq!(prompt_logprobs.len(), 1); + assert!(prompt_logprobs[0].is_none()); + + collect_generate( + output_without_payload(vec![9707, 11]), + "raw-2".to_string(), + ApiServerOptions::default(), + ResponseOptions { + include_prompt_logprobs: true, + ..Default::default() + }, + ) + .expect_err("multi-token prompt without payload is an engine failure"); + } } diff --git a/rust/src/server/src/routes/openai/chat_completions.rs b/rust/src/server/src/routes/openai/chat_completions.rs index b08c533a51b..bd86a06136a 100644 --- a/rust/src/server/src/routes/openai/chat_completions.rs +++ b/rust/src/server/src/routes/openai/chat_completions.rs @@ -35,7 +35,7 @@ use crate::routes::openai::chat_completions::types::{ ChatMessageDelta, }; use crate::routes::openai::utils::logprobs::{ - decoded_logprobs_to_openai_chat, decoded_prompt_logprobs_to_maps, + decoded_logprobs_to_openai_chat, prompt_logprobs_to_maps, }; use crate::routes::openai::utils::types::{ ChatLogProbs, FunctionCallDelta, FunctionCallResponse, ToolCall, ToolCallDelta, Usage, @@ -181,14 +181,11 @@ async fn collect_chat_completion( None }; let prompt_logprobs = if include_prompt_logprobs { - Some(decoded_prompt_logprobs_to_maps( - prompt_logprobs.as_ref().ok_or_else(|| { - server_error!( - "chat response requested prompt_logprobs but generation returned none" - ) - })?, + Some(prompt_logprobs_to_maps( + prompt_logprobs.as_ref(), + &prompt_token_ids, return_tokens_as_token_ids, - )) + )?) } else { None }; diff --git a/rust/src/server/src/routes/openai/completions.rs b/rust/src/server/src/routes/openai/completions.rs index cf15e138bf3..b6f2138f127 100644 --- a/rust/src/server/src/routes/openai/completions.rs +++ b/rust/src/server/src/routes/openai/completions.rs @@ -5,7 +5,6 @@ mod convert; mod types; mod validate; -use std::collections::HashMap; use std::convert::Infallible; use std::result::Result; use std::sync::Arc; @@ -29,8 +28,8 @@ use vllm_text::{ use self::convert::{ResponseOptions, prepare_completion_request}; use super::utils::logprobs::{ - collected_logprobs_to_openai, decoded_logprobs_to_openai, decoded_prompt_logprobs_to_maps, - decoded_prompt_logprobs_to_openai, text_len, + collected_logprobs_to_openai, decoded_logprobs_to_openai, decoded_prompt_logprobs_to_openai, + prompt_logprobs_to_maps, text_len, }; use super::utils::types::Usage; use crate::config::ApiServerOptions; @@ -505,27 +504,6 @@ fn prompt_only_logprobs_to_openai( )) } -fn prompt_logprobs_to_maps( - prompt_logprobs: Option<&DecodedPromptLogprobs>, - prompt_token_ids: &[u32], - return_tokens_as_token_ids: bool, -) -> Result>>, ApiError> { - if let Some(prompt_logprobs) = prompt_logprobs { - return Ok(decoded_prompt_logprobs_to_maps( - prompt_logprobs, - return_tokens_as_token_ids, - )); - } - - if let [_token_id] = prompt_token_ids { - return Ok(vec![None]); - } - - Err(server_error!( - "completion response requested prompt_logprobs but generation returned none" - )) -} - fn usage_chunk( request_id: &str, response_model: &str, diff --git a/rust/src/server/src/routes/openai/utils/logprobs.rs b/rust/src/server/src/routes/openai/utils/logprobs.rs index 078c42b507f..e300e124b1a 100644 --- a/rust/src/server/src/routes/openai/utils/logprobs.rs +++ b/rust/src/server/src/routes/openai/utils/logprobs.rs @@ -100,20 +100,31 @@ pub fn decoded_prompt_logprobs_to_openai( }) } -/// Convert decoded prompt logprobs into the vLLM-style prompt-logprobs response -/// shape. -pub fn decoded_prompt_logprobs_to_maps( - prompt_logprobs: &DecodedPromptLogprobs, +/// Map decoded prompt logprobs into vLLM-style per-position maps, treating a +/// missing single-token payload as `[None]`. +pub fn prompt_logprobs_to_maps( + prompt_logprobs: Option<&DecodedPromptLogprobs>, + prompt_token_ids: &[u32], return_tokens_as_token_ids: bool, -) -> Vec>> { - std::iter::once(None) - .chain(prompt_logprobs.scored_positions.iter().map(|position| { - Some(position_top_logprobs_map( - position, - return_tokens_as_token_ids, - )) - })) - .collect() +) -> Result>>, ApiError> { + if let Some(prompt_logprobs) = prompt_logprobs { + return Ok(std::iter::once(None) + .chain(prompt_logprobs.scored_positions.iter().map(|position| { + Some(position_top_logprobs_map( + position, + return_tokens_as_token_ids, + )) + })) + .collect()); + } + + if let [_token_id] = prompt_token_ids { + return Ok(vec![None]); + } + + Err(server_error!( + "prompt_logprobs were requested but generation returned none" + )) } /// Convert decoded token-position logprobs into the OpenAI chat `logprobs` @@ -275,7 +286,13 @@ pub fn clamp_logprob(logprob: f32) -> f32 { mod tests { use vllm_text::{DecodedLogprobs, DecodedPositionLogprobs, DecodedTokenLogprob}; - use super::decoded_logprobs_to_openai_chat; + use super::{decoded_logprobs_to_openai_chat, prompt_logprobs_to_maps}; + + #[test] + fn prompt_logprobs_maps_reject_missing_multi_token_payload() { + prompt_logprobs_to_maps(None, &[9707, 11], false) + .expect_err("multi-token prompt without payload is an engine failure"); + } fn sample_logprobs() -> DecodedLogprobs { DecodedLogprobs { From 530ee36a0d99869804c0af68b41007072b609480 Mon Sep 17 00:00:00 2001 From: hcl Date: Mon, 20 Jul 2026 18:04:50 +0800 Subject: [PATCH 50/51] fix(openai): reject non-numeric logprobs with 400 instead of 500 (#49144) Signed-off-by: Chenglun Hu --- .../openai/chat_completion/test_chat_error.py | 12 ++++++++++++ .../openai/completion/test_completion_error.py | 13 +++++++++++++ vllm/entrypoints/openai/chat_completion/protocol.py | 12 ++++++++++++ vllm/entrypoints/openai/completion/protocol.py | 12 ++++++++++++ 4 files changed, 49 insertions(+) diff --git a/tests/entrypoints/openai/chat_completion/test_chat_error.py b/tests/entrypoints/openai/chat_completion/test_chat_error.py index 4b6be87ae5c..4b42f522e81 100644 --- a/tests/entrypoints/openai/chat_completion/test_chat_error.py +++ b/tests/entrypoints/openai/chat_completion/test_chat_error.py @@ -515,3 +515,15 @@ def test_structured_outputs_structural_tag_invalid(structural_tag): messages=[{"role": "user", "content": "hello"}], structured_outputs={"structural_tag": structural_tag}, ) + + +@pytest.mark.parametrize("field_name", ["prompt_logprobs", "top_logprobs"]) +def test_non_numeric_logprobs_rejected(field_name): + """A non-numeric logprobs value must be a clean 400 validation error, not a + TypeError from the mode='before' comparison (which surfaces as HTTP 500).""" + with pytest.raises(ValidationError, match=f"`{field_name}` must be an integer"): + ChatCompletionRequest( + model=MODEL_NAME, + messages=[{"role": "user", "content": "hello"}], + **{field_name: "2"}, + ) diff --git a/tests/entrypoints/openai/completion/test_completion_error.py b/tests/entrypoints/openai/completion/test_completion_error.py index aa9e9c1d72e..818cad738d9 100644 --- a/tests/entrypoints/openai/completion/test_completion_error.py +++ b/tests/entrypoints/openai/completion/test_completion_error.py @@ -610,3 +610,16 @@ class TestCompletionPromptListLimit: max_tokens=1, ) assert len(request.prompt_embeds) == 5 + + +@pytest.mark.parametrize("field_name", ["prompt_logprobs", "logprobs"]) +def test_non_numeric_logprobs_rejected(field_name): + """A non-numeric logprobs value must be a clean 400 validation error, not a + TypeError from the mode='before' comparison (which surfaces as HTTP 500).""" + with pytest.raises(ValidationError, match=f"`{field_name}` must be an integer"): + CompletionRequest( + model=MODEL_NAME, + prompt="Test prompt", + max_tokens=10, + **{field_name: "2"}, + ) diff --git a/vllm/entrypoints/openai/chat_completion/protocol.py b/vllm/entrypoints/openai/chat_completion/protocol.py index 8c1694bbf73..af8a19c6432 100644 --- a/vllm/entrypoints/openai/chat_completion/protocol.py +++ b/vllm/entrypoints/openai/chat_completion/protocol.py @@ -757,6 +757,18 @@ class ChatCompletionRequest(OpenAIBaseModel): parameter="logprob_token_ids", ) + # These fields are integers, but `mode="before"` runs on the raw + # request data, so a non-numeric value (e.g. a JSON string) would + # reach the comparisons below and raise TypeError -> HTTP 500. Reject + # it here so the client gets a clean 400 instead. + for field_name in ("prompt_logprobs", "top_logprobs"): + field_value = data.get(field_name) + if field_value is not None and not isinstance(field_value, (int, float)): + raise VLLMValidationError( + f"`{field_name}` must be an integer.", + parameter=field_name, + value=field_value, + ) if (prompt_logprobs := data.get("prompt_logprobs")) is not None: if data.get("stream") and (prompt_logprobs > 0 or prompt_logprobs == -1): raise VLLMValidationError( diff --git a/vllm/entrypoints/openai/completion/protocol.py b/vllm/entrypoints/openai/completion/protocol.py index 73677b16af9..6a2bf47bb15 100644 --- a/vllm/entrypoints/openai/completion/protocol.py +++ b/vllm/entrypoints/openai/completion/protocol.py @@ -468,6 +468,18 @@ class CompletionRequest(OpenAIBaseModel): parameter="logprob_token_ids", ) + # These fields are integers, but `mode="before"` runs on the raw + # request data, so a non-numeric value (e.g. a JSON string) would + # reach the comparisons below and raise TypeError -> HTTP 500. Reject + # it here so the client gets a clean 400 instead. + for field_name in ("prompt_logprobs", "logprobs"): + field_value = data.get(field_name) + if field_value is not None and not isinstance(field_value, (int, float)): + raise VLLMValidationError( + f"`{field_name}` must be an integer.", + parameter=field_name, + value=field_value, + ) if (prompt_logprobs := data.get("prompt_logprobs")) is not None: if data.get("stream") and (prompt_logprobs > 0 or prompt_logprobs == -1): raise VLLMValidationError( From ae10e855abf4ff5e24e2088aef16029ee1cb7de8 Mon Sep 17 00:00:00 2001 From: Lena Onyshchenko Date: Mon, 20 Jul 2026 13:05:36 +0300 Subject: [PATCH 51/51] [Misc][Docs] Remove duplicate CodeGeex4 row in XPU model table (#47210) Signed-off-by: oonyshch --- docs/models/hardware_supported_models/xpu.md | 2 -- 1 file changed, 2 deletions(-) diff --git a/docs/models/hardware_supported_models/xpu.md b/docs/models/hardware_supported_models/xpu.md index d065b4b6890..49636d0d5c9 100644 --- a/docs/models/hardware_supported_models/xpu.md +++ b/docs/models/hardware_supported_models/xpu.md @@ -31,10 +31,8 @@ | THUDM/CodeGeex4-All-9B | CodeGeexForCausalLM | ✅ | | | | chuhac/TeleChat2-35B | LlamaForCausalLM (TeleChat2 based on Llama arch) | ✅ | | | | 01-ai/Yi1.5-34B-Chat | YiForCausalLM | ✅ | | | -| THUDM/CodeGeex4-All-9B | CodeGeexForCausalLM | ✅ | | | | deepseek-ai/DeepSeek-Coder-33B-base | DeepSeekCoderForCausalLM | ✅ | | | | meta-llama/Llama-2-13b-chat-hf | LlamaForCausalLM | ✅ | | | -| THUDM/CodeGeex4-All-9B | CodeGeexForCausalLM | ✅ | | | | Qwen/Qwen1.5-14B-Chat | QwenForCausalLM | ✅ | | | | Qwen/Qwen1.5-32B-Chat | QwenForCausalLM | ✅ | | | | RedHatAI/Meta-Llama-3.1-8B-Instruct-FP8-dynamic | LlamaForCausalLM | | ✅ | |