Compare commits

...
Author SHA1 Message Date
Kevin H. Luu 568afb3a13 [CI/Build] Refresh tags before building macOS wheel (#49901)
Signed-off-by: khluu <khluu000@gmail.com>
Co-authored-by: OpenAI Codex <codex@openai.com>
(cherry picked from commit 0934b26790)
2026-07-26 17:57:50 -07:00
TJianandkhluu f2654939e6 [ROCm] [Release] [Bugfix] Fix the per commit wheel release pipeline. (#49245)
Signed-off-by: tjtanaa <tunjian.tan@embeddedllm.com>
2026-07-25 00:23:46 -07:00
djramicandkhluu ffd46bfab2 [Bugfix] Register axk1 config to fix A.X-K1 init (#49727)
Signed-off-by: Djordje Ramic <djoramic@amd.com>
(cherry picked from commit e222c33f2f)
2026-07-24 18:48:21 -07:00
Andrey Talmanandkhluu ffd6ee4bcc [CI] Bump PyTorch Compilation Unit Tests timeout to 150 min (#49606) 2026-07-23 21:20:25 -07:00
Nick Hillandkhluu bb26ce8e93 [CI] Increase timeout of pytorch-compilation-unit-tests (#49450)
Signed-off-by: Nick Hill <nickhill123@gmail.com>
2026-07-23 21:20:25 -07:00
Kevin H. LuuandOpenAI Codex 091db8b58f [CI] Increase timeouts for jobs exceeding current limits (#49374)
Signed-off-by: khluu <khluu000@gmail.com>
Co-authored-by: OpenAI Codex <codex@openai.com>
2026-07-23 21:20:25 -07:00
Nick Hillandkhluu ba694b86f2 [CI] Bump timeout of entrypoints-integration-api-server-openai-part-2 (#49359)
Signed-off-by: Nick Hill <nickhill123@gmail.com>
2026-07-23 21:20:25 -07:00
zhrrrandkhluu e5949f1000 [Bugfix] handle grammar compilation failures to avoid engine crash (#47312)
Signed-off-by: zhuhaoran <zhuhaoran.zhr@alibaba-inc.com>
Signed-off-by: Nick Hill <nickhill123@gmail.com>
Co-authored-by: Nick Hill <nickhill123@gmail.com>
(cherry picked from commit 12213c6795)
2026-07-23 11:13:25 -07:00
Michael Goinandkhluu 8b30569e83 [Bugfix] Fix DeepGEMM warmup when using FlashInferFp8DeepGEMMDynamicBlockScaledKernel (#49467)
Signed-off-by: mgoin <mgoin64@gmail.com>
(cherry picked from commit 917fdb5bf7)
2026-07-23 11:13:25 -07:00
Lucas Wilkinsonandkhluu 9d37a50c80 [Bugfix][Attention] Ignore empty MLA context chunks during merge (#49294)
Signed-off-by: Lucas Wilkinson <lwilkins@redhat.com>
Co-authored-by: OpenAI Codex <codex@openai.com>
(cherry picked from commit 060b5f61dc)
2026-07-23 11:13:25 -07:00
aoshen02andkhluu 2dd1e7cd3b Update BGE-M3 token expectations for leading spaces (#49269)
Signed-off-by: aoshen02 <aoshen02@users.noreply.github.com>
Co-authored-by: aoshen02 <aoshen02@users.noreply.github.com>
Co-authored-by: Codex <noreply@openai.com>
(cherry picked from commit d9aa35161d)
2026-07-23 11:13:25 -07:00
Alejandro Paredes La Torreandkhluu a54c93a146 [Bugfix] Fix WSL circular import from pin_memory warning_once (#48444)
Signed-off-by: AlejandroParedesLT <alejandroparedeslatorre@gmail.com>
Co-authored-by: Shengqi Chen <harry-chen@outlook.com>
(cherry picked from commit 0a684ab0c0)
2026-07-23 11:13:25 -07:00
30 changed files with 367 additions and 73 deletions
+2 -2
View File
@@ -813,8 +813,8 @@ steps:
# Download artifacts from current build # Download artifacts from current build
echo "Downloading 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-base-wheels/*.whl" .
# buildkite-agent artifact download "artifacts/rocm-vllm-wheel/*.whl" . buildkite-agent artifact download "artifacts/rocm-vllm-wheel/*.whl" .
# # Run upload script # # Run upload script
bash .buildkite/scripts/upload-rocm-wheels.sh bash .buildkite/scripts/upload-rocm-wheels.sh
+3
View File
@@ -7,6 +7,9 @@
set -euo pipefail set -euo pipefail
# The macmini queue uses persistent checkouts, so refresh tags for setuptools-scm.
git fetch --tags --force origin
# The Rust frontend build needs protoc. # The Rust frontend build needs protoc.
if ! command -v protoc >/dev/null 2>&1; then if ! command -v protoc >/dev/null 2>&1; then
brew install protobuf brew install protobuf
+1 -1
View File
@@ -4,7 +4,7 @@ depends_on:
steps: steps:
- label: Basic Correctness - label: Basic Correctness
key: basic-correctness key: basic-correctness
timeout_in_minutes: 45 timeout_in_minutes: 68
device: h200_18gb device: h200_18gb
source_file_dependencies: source_file_dependencies:
- vllm/ - vllm/
+1 -1
View File
@@ -4,7 +4,7 @@ depends_on:
steps: steps:
- label: Benchmarks CLI Test - label: Benchmarks CLI Test
key: benchmarks-cli-test key: benchmarks-cli-test
timeout_in_minutes: 30 timeout_in_minutes: 45
device: h200_18gb device: h200_18gb
source_file_dependencies: source_file_dependencies:
- vllm/ - vllm/
+1 -1
View File
@@ -51,7 +51,7 @@ steps:
- label: e2e Scheduling (1 GPU) - label: e2e Scheduling (1 GPU)
key: e2e-scheduling-1-gpu key: e2e-scheduling-1-gpu
timeout_in_minutes: 35 timeout_in_minutes: 53
device: h200_18gb device: h200_18gb
source_file_dependencies: source_file_dependencies:
- vllm/v1/ - vllm/v1/
+4 -4
View File
@@ -39,7 +39,7 @@ steps:
- label: Entrypoints Integration (API Server) - label: Entrypoints Integration (API Server)
key: entrypoints-integration-api-server key: entrypoints-integration-api-server
device: h200_35gb device: h200_35gb
timeout_in_minutes: 50 timeout_in_minutes: 75
working_dir: "/vllm-workspace/tests" working_dir: "/vllm-workspace/tests"
source_file_dependencies: source_file_dependencies:
- vllm/ - vllm/
@@ -59,7 +59,7 @@ steps:
- label: Entrypoints Integration (API Server OpenAI - Part 1) - label: Entrypoints Integration (API Server OpenAI - Part 1)
device: h200_35gb device: h200_35gb
key: entrypoints-integration-api-server-openai-part-1 key: entrypoints-integration-api-server-openai-part-1
timeout_in_minutes: 45 timeout_in_minutes: 68
working_dir: "/vllm-workspace/tests" working_dir: "/vllm-workspace/tests"
source_file_dependencies: source_file_dependencies:
- vllm/ - vllm/
@@ -78,7 +78,7 @@ steps:
- label: Entrypoints Integration (API Server OpenAI - Part 2) - label: Entrypoints Integration (API Server OpenAI - Part 2)
device: h200_35gb device: h200_35gb
key: entrypoints-integration-api-server-openai-part-2 key: entrypoints-integration-api-server-openai-part-2
timeout_in_minutes: 45 timeout_in_minutes: 83
working_dir: "/vllm-workspace/tests" working_dir: "/vllm-workspace/tests"
source_file_dependencies: source_file_dependencies:
- vllm/ - vllm/
@@ -156,7 +156,7 @@ steps:
- label: Entrypoints Integration (Pooling) - label: Entrypoints Integration (Pooling)
device: h200_35gb device: h200_35gb
key: entrypoints-integration-pooling key: entrypoints-integration-pooling
timeout_in_minutes: 50 timeout_in_minutes: 75
working_dir: "/vllm-workspace/tests" working_dir: "/vllm-workspace/tests"
source_file_dependencies: source_file_dependencies:
- vllm/ - vllm/
+1 -1
View File
@@ -31,7 +31,7 @@ steps:
- label: V1 Sample + Logits - label: V1 Sample + Logits
key: v1-sample-logits key: v1-sample-logits
timeout_in_minutes: 45 timeout_in_minutes: 83
device: h200_18gb device: h200_18gb
source_file_dependencies: source_file_dependencies:
- vllm/config/ - vllm/config/
+1 -1
View File
@@ -5,7 +5,7 @@ steps:
- label: Model Executor - label: Model Executor
device: h200_35gb device: h200_35gb
key: model-executor key: model-executor
timeout_in_minutes: 45 timeout_in_minutes: 60
source_file_dependencies: source_file_dependencies:
- vllm/engine/arg_utils.py - vllm/engine/arg_utils.py
- vllm/config/model.py - vllm/config/model.py
+1 -1
View File
@@ -137,7 +137,7 @@ steps:
- label: Language Models Test (MTEB) - label: Language Models Test (MTEB)
key: language-models-test-mteb key: language-models-test-mteb
timeout_in_minutes: 45 timeout_in_minutes: 68
device: h200_18gb device: h200_18gb
optional: true optional: true
source_file_dependencies: source_file_dependencies:
+4 -4
View File
@@ -4,7 +4,7 @@ depends_on:
steps: steps:
- label: "Multi-Modal Models (Standard) 1: qwen2" - label: "Multi-Modal Models (Standard) 1: qwen2"
key: multi-modal-models-standard-1-qwen2 key: multi-modal-models-standard-1-qwen2
timeout_in_minutes: 45 timeout_in_minutes: 68
device: h200_18gb device: h200_18gb
source_file_dependencies: source_file_dependencies:
- vllm/ - vllm/
@@ -20,7 +20,7 @@ steps:
- label: "Multi-Modal Models (Standard) 2: qwen3 + gemma" - label: "Multi-Modal Models (Standard) 2: qwen3 + gemma"
key: multi-modal-models-standard-2-qwen3-gemma key: multi-modal-models-standard-2-qwen3-gemma
timeout_in_minutes: 50 timeout_in_minutes: 75
device: h200_18gb device: h200_18gb
source_file_dependencies: source_file_dependencies:
- vllm/ - vllm/
@@ -54,7 +54,7 @@ steps:
- label: "Multi-Modal Models (Standard) 4: other + whisper" - label: "Multi-Modal Models (Standard) 4: other + whisper"
device: h200_35gb device: h200_35gb
key: multi-modal-models-standard-4-other-whisper key: multi-modal-models-standard-4-other-whisper
timeout_in_minutes: 50 timeout_in_minutes: 75
source_file_dependencies: source_file_dependencies:
- vllm/ - vllm/
- tests/models/multimodal - tests/models/multimodal
@@ -85,7 +85,7 @@ steps:
- label: Multi-Modal Processor # 44min - label: Multi-Modal Processor # 44min
key: multi-modal-processor key: multi-modal-processor
timeout_in_minutes: 65 timeout_in_minutes: 98
device: h200_18gb device: h200_18gb
source_file_dependencies: source_file_dependencies:
- vllm/ - vllm/
+1 -1
View File
@@ -5,7 +5,7 @@ steps:
- label: PyTorch Compilation Unit Tests - label: PyTorch Compilation Unit Tests
device: h200_35gb device: h200_35gb
key: pytorch-compilation-unit-tests key: pytorch-compilation-unit-tests
timeout_in_minutes: 90 timeout_in_minutes: 150
source_file_dependencies: source_file_dependencies:
- vllm/__init__.py - vllm/__init__.py
- vllm/_aiter_ops.py - vllm/_aiter_ops.py
@@ -6,6 +6,7 @@ max_concurrency: 100
server_args: >- server_args: >-
--enforce-eager --enforce-eager
--max-model-len 4096 --max-model-len 4096
--max-num-batched-tokens 32768
--safetensors-load-strategy prefetch --safetensors-load-strategy prefetch
--moe-backend flashinfer_cutlass --moe-backend flashinfer_cutlass
--prefill-context-parallel-size 4 --prefill-context-parallel-size 4
@@ -6,6 +6,7 @@ max_concurrency: 100
server_args: >- server_args: >-
--enforce-eager --enforce-eager
--max-model-len 4096 --max-model-len 4096
--max-num-batched-tokens 32768
--safetensors-load-strategy prefetch --safetensors-load-strategy prefetch
--moe-backend flashinfer_cutlass --moe-backend flashinfer_cutlass
--tensor-parallel-size 2 --tensor-parallel-size 2
@@ -11,6 +11,9 @@ from vllm._custom_ops import (
scaled_fp8_quant, scaled_fp8_quant,
) )
from vllm.platforms import current_platform from vllm.platforms import current_platform
from vllm.v1.attention.ops.triton_merge_attn_states import (
mask_empty_context,
)
from vllm.v1.attention.ops.triton_merge_attn_states import ( from vllm.v1.attention.ops.triton_merge_attn_states import (
merge_attn_states as merge_attn_states_triton, merge_attn_states as merge_attn_states_triton,
) )
@@ -73,6 +76,59 @@ DTYPES = [torch.float32, torch.half, torch.bfloat16]
all_case_info: list[tuple] = [] all_case_info: list[tuple] = []
def test_mask_empty_context() -> None:
query_lens = torch.tensor([2] + [1] * 31 + [131, 1], dtype=torch.int32)
query_start_loc = torch.cat(
(torch.zeros(1, dtype=torch.int32), query_lens.cumsum(0))
).cuda()
context_lens = torch.tensor([4] * 32 + [0, 3], dtype=torch.int32)
context_start_loc = torch.cat(
(torch.zeros(1, dtype=torch.int32), context_lens.cumsum(0))
).cuda()
num_heads, num_tokens, head_dim = 4, 165, 16
lse = torch.randn(num_heads, num_tokens, device="cuda")
output = torch.randn(num_tokens, num_heads, head_dim, device="cuda")
# Empty-context rows carry undefined (possibly non-finite) attention output.
output[33:164] = float("nan")
expected_lse = lse.clone()
expected_lse[:, 33:164] = float("-inf")
expected_output = output.clone()
expected_output[33:164] = 0.0
mask_empty_context(lse, output, query_start_loc, context_start_loc)
torch.testing.assert_close(lse, expected_lse)
torch.testing.assert_close(output, expected_output)
@pytest.mark.parametrize("merge_fn", [merge_attn_states_cuda, merge_attn_states_triton])
@pytest.mark.parametrize("output_dtype", [torch.float32, torch.half, torch.bfloat16])
def test_merge_attn_states_both_empty(merge_fn, output_dtype) -> None:
"""When a token is empty on both sides (both LSE -inf), the 0/0 softmax
scales must not surface as NaN in the merged output."""
num_tokens, num_heads, head_size = 6, 8, 128
prefix_output = torch.zeros(
num_tokens, num_heads, head_size, device="cuda", dtype=output_dtype
)
prefix_lse = torch.randn(num_heads, num_tokens, device="cuda")
suffix_output = torch.zeros(
num_tokens, num_heads, head_size, device="cuda", dtype=output_dtype
)
suffix_lse = torch.randn(num_heads, num_tokens, device="cuda")
# Tokens 2 and 3 are empty on both sides (mask_empty_context already zeroed
# their outputs and set both LSEs to -inf).
empty = slice(2, 4)
prefix_lse[:, empty] = float("-inf")
suffix_lse[:, empty] = float("-inf")
output = torch.empty_like(prefix_output)
merge_fn(output, prefix_output, prefix_lse, suffix_output, suffix_lse)
assert not output.isnan().any()
def generate_markdown_table(): def generate_markdown_table():
global all_case_info global all_case_info
table_header = ( table_header = (
@@ -48,12 +48,12 @@ def _check_dense_embedding(data, index=0):
def _check_sparse_embedding(data, check_tokens=False): def _check_sparse_embedding(data, check_tokens=False):
expected_weights = [ expected_weights = [
{"token_id": 32, "weight": 0.0552978515625, "token": "?"}, {"token_id": 32, "weight": 0.0552978515625, "token": "?"},
{"token_id": 70, "weight": 0.09808349609375, "token": "the"}, {"token_id": 70, "weight": 0.09808349609375, "token": " the"},
{"token_id": 83, "weight": 0.08154296875, "token": "is"}, {"token_id": 83, "weight": 0.08154296875, "token": " is"},
{"token_id": 111, "weight": 0.11810302734375, "token": "of"}, {"token_id": 111, "weight": 0.11810302734375, "token": " of"},
{"token_id": 4865, "weight": 0.1171875, "token": "What"}, {"token_id": 4865, "weight": 0.1171875, "token": " What"},
{"token_id": 9942, "weight": 0.292236328125, "token": "France"}, {"token_id": 9942, "weight": 0.292236328125, "token": " France"},
{"token_id": 10323, "weight": 0.2802734375, "token": "capital"}, {"token_id": 10323, "weight": 0.2802734375, "token": " capital"},
] ]
expected_embed = {x["token_id"]: x for x in expected_weights} expected_embed = {x["token_id"]: x for x in expected_weights}
+3 -1
View File
@@ -9,6 +9,7 @@ from vllm.v1.core.sched.async_scheduler import AsyncScheduler
from vllm.v1.core.sched.output import CachedRequestData, SchedulerOutput from vllm.v1.core.sched.output import CachedRequestData, SchedulerOutput
from vllm.v1.outputs import ModelRunnerOutput from vllm.v1.outputs import ModelRunnerOutput
from vllm.v1.request import RequestStatus from vllm.v1.request import RequestStatus
from vllm.v1.structured_output import StructuredOutputGrammar
from vllm.v1.utils import ConstantList from vllm.v1.utils import ConstantList
from .utils import create_requests, create_scheduler from .utils import create_requests, create_scheduler
@@ -262,7 +263,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance():
scheduler = object.__new__(AsyncScheduler) scheduler = object.__new__(AsyncScheduler)
request = create_requests(num_requests=1, num_tokens=1)[0] request = create_requests(num_requests=1, num_tokens=1)[0]
request.structured_output_request = Mock() request.structured_output_request = Mock()
request.structured_output_request.grammar = Mock() request.structured_output_request.grammar = Mock(spec=StructuredOutputGrammar)
request.structured_output_request.grammar.accept_tokens.return_value = False request.structured_output_request.grammar.accept_tokens.return_value = False
request.status = RequestStatus.RUNNING request.status = RequestStatus.RUNNING
request.num_computed_tokens = request.num_tokens request.num_computed_tokens = request.num_tokens
@@ -284,6 +285,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance():
scheduler.kv_event_publisher = Mock() scheduler.kv_event_publisher = Mock()
scheduler.finished_req_ids = set() scheduler.finished_req_ids = set()
scheduler.finished_req_ids_dict = None scheduler.finished_req_ids_dict = None
scheduler.grammar_compile_error_reqs = set()
scheduler.vllm_config = Mock() scheduler.vllm_config = Mock()
scheduler.vllm_config.model_config.enable_return_routed_experts = False scheduler.vllm_config.model_config.enable_return_routed_experts = False
scheduler.enable_return_routed_experts = False scheduler.enable_return_routed_experts = False
+56 -2
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project # SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import dataclasses import dataclasses
from concurrent.futures import Future
from unittest.mock import Mock from unittest.mock import Mock
import pytest import pytest
@@ -36,7 +37,7 @@ from vllm.v1.kv_cache_interface import (
) )
from vllm.v1.outputs import DraftTokenIds, KVConnectorOutput, ModelRunnerOutput from vllm.v1.outputs import DraftTokenIds, KVConnectorOutput, ModelRunnerOutput
from vllm.v1.request import Request, RequestStatus from vllm.v1.request import Request, RequestStatus
from vllm.v1.structured_output import StructuredOutputManager from vllm.v1.structured_output import StructuredOutputGrammar, StructuredOutputManager
from .utils import EOS_TOKEN_ID, create_requests, create_scheduler, mock_kv from .utils import EOS_TOKEN_ID, create_requests, create_scheduler, mock_kv
@@ -3006,6 +3007,58 @@ def test_schedule_skip_tokenizer_init_structured_output_request():
assert len(scheduler.skipped_waiting) == 1 assert len(scheduler.skipped_waiting) == 1
@pytest.mark.parametrize("async_grammar", [True, False])
def test_grammar_compile_error_finishes_only_request(async_grammar: bool):
scheduler = create_scheduler()
manager = scheduler.structured_output_manager
manager.backend = Mock()
manager.backend.compile_grammar.side_effect = RuntimeError(
"forced FSM compilation error"
)
manager._use_async_grammar_compilation = async_grammar
sampling_params = SamplingParams(
max_tokens=16,
structured_outputs=StructuredOutputsParams(json='{"type": "object"}'),
)
sampling_params.update_from_generation_config({}, EOS_TOKEN_ID)
request = Request(
request_id="grammar-error",
prompt_token_ids=[0, 1],
sampling_params=sampling_params,
pooling_params=None,
)
manager.grammar_init(request)
assert request.structured_output_request is not None
grammar_future = request.structured_output_request._grammar
assert isinstance(grammar_future, Future)
assert isinstance(grammar_future.exception(timeout=5), RuntimeError)
scheduler.add_request(request)
scheduler_output = scheduler.schedule()
assert not scheduler_output.num_scheduled_tokens
engine_core_outputs = scheduler.update_from_output(
scheduler_output,
ModelRunnerOutput(req_ids=[], req_id_to_index={}),
)
assert request.status == RequestStatus.FINISHED_ERROR
assert request.request_id not in scheduler.requests
output = engine_core_outputs[0].outputs[0]
assert output.request_id == request.request_id
assert output.finish_reason == FinishReason.ERROR
assert output.stop_reason is None
healthy_request = create_requests(num_requests=1, req_ids=["healthy-request"])[0]
scheduler.add_request(healthy_request)
next_output = scheduler.schedule()
assert [req.req_id for req in next_output.scheduled_new_reqs] == [
healthy_request.request_id
]
def test_abort_request_when_structured_output_fsm_cannot_advance(): def test_abort_request_when_structured_output_fsm_cannot_advance():
scheduler = object.__new__(Scheduler) scheduler = object.__new__(Scheduler)
sampling_params = SamplingParams(ignore_eos=True, max_tokens=4) sampling_params = SamplingParams(ignore_eos=True, max_tokens=4)
@@ -3019,7 +3072,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance():
pooling_params=None, pooling_params=None,
) )
request.structured_output_request = Mock() request.structured_output_request = Mock()
request.structured_output_request.grammar = Mock() request.structured_output_request.grammar = Mock(spec=StructuredOutputGrammar)
request.structured_output_request.grammar.accept_tokens.return_value = False request.structured_output_request.grammar.accept_tokens.return_value = False
request.status = RequestStatus.RUNNING request.status = RequestStatus.RUNNING
request.num_computed_tokens = request.num_tokens request.num_computed_tokens = request.num_tokens
@@ -3040,6 +3093,7 @@ def test_abort_request_when_structured_output_fsm_cannot_advance():
scheduler.kv_event_publisher = Mock() scheduler.kv_event_publisher = Mock()
scheduler.finished_req_ids = set() scheduler.finished_req_ids = set()
scheduler.finished_req_ids_dict = None scheduler.finished_req_ids_dict = None
scheduler.grammar_compile_error_reqs = set()
scheduler.vllm_config = Mock() scheduler.vllm_config = Mock()
scheduler.vllm_config.model_config.enable_return_routed_experts = False scheduler.vllm_config.model_config.enable_return_routed_experts = False
scheduler.enable_return_routed_experts = False scheduler.enable_return_routed_experts = False
@@ -277,6 +277,7 @@ from vllm.v1.attention.backends.utils import (
from vllm.v1.attention.ops.common import cp_lse_ag_out_ar, 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.dcp_alltoall import dcp_a2a_lse_reduce
from vllm.v1.attention.ops.merge_attn_states import merge_attn_states from vllm.v1.attention.ops.merge_attn_states import merge_attn_states
from vllm.v1.attention.ops.triton_merge_attn_states import mask_empty_context
from vllm.v1.attention.selector import get_attn_backend from vllm.v1.attention.selector import get_attn_backend
from vllm.v1.kv_cache_interface import ( from vllm.v1.kv_cache_interface import (
AttentionSpec, AttentionSpec,
@@ -1342,6 +1343,7 @@ class MLACommonPrefillMetadata:
workspace: torch.Tensor workspace: torch.Tensor
token_to_seq: torch.Tensor token_to_seq: torch.Tensor
chunk_total_token: list[int] chunk_total_token: list[int]
has_empty_context: list[bool]
# for mla DCP # for mla DCP
padded_local_chunk_seq_lens: list[list[int]] | None = None padded_local_chunk_seq_lens: list[list[int]] | None = None
@@ -1551,6 +1553,7 @@ def build_mla_chunked_context_metadata(
) )
chunk_seq_lens = chunk_ends - chunk_starts chunk_seq_lens = chunk_ends - chunk_starts
chunk_seq_lens.clamp_(min=0) chunk_seq_lens.clamp_(min=0)
has_empty_context = torch.any(chunk_seq_lens == 0, dim=1).tolist()
cu_seq_lens_cpu = torch.zeros( cu_seq_lens_cpu = torch.zeros(
num_chunks, num_prefills + 1, dtype=torch.int32, pin_memory=True num_chunks, num_prefills + 1, dtype=torch.int32, pin_memory=True
@@ -1629,6 +1632,7 @@ def build_mla_chunked_context_metadata(
token_to_seq=token_to_seq_cpu.to(device, non_blocking=True), token_to_seq=token_to_seq_cpu.to(device, non_blocking=True),
chunk_total_token=chunk_total_token.tolist(), chunk_total_token=chunk_total_token.tolist(),
workspace=chunked_prefill_workspace, workspace=chunked_prefill_workspace,
has_empty_context=has_empty_context,
prefill_tokens_with_context=prefill_tokens_with_context, prefill_tokens_with_context=prefill_tokens_with_context,
padded_local_chunk_seq_lens=padded_local_chunk_seq_lens.tolist(), padded_local_chunk_seq_lens=padded_local_chunk_seq_lens.tolist(),
local_context_lens_allranks=local_context_lens_allranks.tolist(), local_context_lens_allranks=local_context_lens_allranks.tolist(),
@@ -1651,6 +1655,7 @@ def build_mla_chunked_context_metadata(
token_to_seq=token_to_seq_cpu.to(device, non_blocking=True), token_to_seq=token_to_seq_cpu.to(device, non_blocking=True),
chunk_total_token=chunk_total_token, chunk_total_token=chunk_total_token,
workspace=chunked_prefill_workspace, workspace=chunked_prefill_workspace,
has_empty_context=has_empty_context,
prefill_tokens_with_context=prefill_tokens_with_context, prefill_tokens_with_context=prefill_tokens_with_context,
) )
@@ -2238,6 +2243,13 @@ class MLACommonBaseImpl(MLAAttentionImpl[A], Generic[A]):
v=v, v=v,
) )
) )
if prefill_metadata.chunked_context.has_empty_context[i]:
mask_empty_context(
attn_softmax_lse,
attn_output,
prefill_metadata.query_start_loc,
prefill_metadata.chunked_context.cu_seq_lens[i],
)
if output is None: if output is None:
output = attn_output output = attn_output
@@ -2388,6 +2400,13 @@ class MLACommonBaseImpl(MLAAttentionImpl[A], Generic[A]):
v=v, v=v,
) )
) )
if prefill_metadata.chunked_context.has_empty_context[i]:
mask_empty_context(
attn_softmax_lse,
attn_output,
prefill_metadata.query_start_loc,
prefill_metadata.chunked_context.cu_seq_lens[i],
)
if output is None: if output is None:
output = attn_output output = attn_output
+14 -4
View File
@@ -133,6 +133,18 @@ def _extract_data_from_fused_moe_module(
return w13, w13_s, w2, w2_s, num_topk return w13, w13_s, w2, w2_s, num_topk
def _is_deep_gemm_backed_kernel(fp8_linear: object) -> bool:
"""
Return True if the selected linear kernel dispatches to DeepGEMM, either
directly or as the fallback branch of a dynamic wrapper.
"""
if isinstance(fp8_linear, DeepGemmFp8BlockScaledMMKernel):
return True
return isinstance(
getattr(fp8_linear, "fallback", None), DeepGemmFp8BlockScaledMMKernel
)
def _fp8_linear_may_use_deep_gemm(module: torch.nn.Module) -> bool: def _fp8_linear_may_use_deep_gemm(module: torch.nn.Module) -> bool:
""" """
Return True if the input module/layer could be processed with DeepGEMM. Return True if the input module/layer could be processed with DeepGEMM.
@@ -147,10 +159,8 @@ def _fp8_linear_may_use_deep_gemm(module: torch.nn.Module) -> bool:
): ):
return False return False
if not isinstance( fp8_linear = getattr(module.quant_method, "fp8_linear", None)
getattr(module.quant_method, "fp8_linear", None), if not _is_deep_gemm_backed_kernel(fp8_linear):
DeepGemmFp8BlockScaledMMKernel,
):
return False return False
block_size = get_mk_alignment_for_contiguous_layout()[0] block_size = get_mk_alignment_for_contiguous_layout()[0]
+2 -1
View File
@@ -290,7 +290,8 @@ class CudaPlatformBase(Platform):
# kernel with limited pinned memory support for CUDA. # kernel with limited pinned memory support for CUDA.
version = _get_wsl_kernel_version() version = _get_wsl_kernel_version()
if version is None or version < (4, 19, 121): if version is None or version < (4, 19, 121):
logger.warning_once( # warning_once() causes a circular import on WSL, see #48397.
logger.warning(
"Using 'pin_memory=False' as WSL is detected and the " "Using 'pin_memory=False' as WSL is detected and the "
"WSL2 kernel version is below 4.19.121. This may slow " "WSL2 kernel version is below 4.19.121. This may slow "
"down performance. Please run `wsl --update`." "down performance. Please run `wsl --update`."
+2 -1
View File
@@ -991,7 +991,8 @@ class Platform:
# Pinned memory support under WSL depends on the vendor and driver # Pinned memory support under WSL depends on the vendor and driver
# version. Conservative default: return False. Platform subclasses # version. Conservative default: return False. Platform subclasses
# that can verify support (e.g. CudaPlatformBase) override this. # that can verify support (e.g. CudaPlatformBase) override this.
logger.warning_once( # warning_once() causes a circular import on WSL, see #48397.
logger.warning(
"Using 'pin_memory=False' as WSL is detected. " "Using 'pin_memory=False' as WSL is detected. "
"This may slow down performance." "This may slow down performance."
) )
+1
View File
@@ -72,6 +72,7 @@ class LazyConfigDict(dict):
_CONFIG_REGISTRY: dict[str, type[PretrainedConfig]] = LazyConfigDict( _CONFIG_REGISTRY: dict[str, type[PretrainedConfig]] = LazyConfigDict(
afmoe="AfmoeConfig", afmoe="AfmoeConfig",
arctic="ArcticConfig", arctic="ArcticConfig",
axk1="AXK1Config",
bagel="BagelConfig", bagel="BagelConfig",
umm="CheersConfig", umm="CheersConfig",
chatglm="ChatGLMConfig", chatglm="ChatGLMConfig",
+1 -1
View File
@@ -114,7 +114,7 @@ class AXK1Config(PretrainedConfig):
The dropout ratio for the attention probabilities. The dropout ratio for the attention probabilities.
""" """
model_type = "AXK1" model_type = "axk1"
keys_to_ignore_at_inference = ["past_key_values"] keys_to_ignore_at_inference = ["past_key_values"]
def __init__( def __init__(
@@ -261,7 +261,7 @@ class ModelArchConfigConvertorBase:
if not hasattr(self.hf_text_config, "model_type"): if not hasattr(self.hf_text_config, "model_type"):
return False return False
elif self.hf_text_config.model_type in ( elif self.hf_text_config.model_type in (
"AXK1", "axk1",
"deepseek_v2", "deepseek_v2",
"deepseek_v3", "deepseek_v3",
"deepseek_v32", "deepseek_v32",
@@ -290,7 +290,7 @@ class ModelArchConfigConvertorBase:
return ( return (
self.hf_text_config.model.model_type self.hf_text_config.model.model_type
in ( in (
"AXK1", "axk1",
"deepseek_v2", "deepseek_v2",
"deepseek_v3", "deepseek_v3",
"deepseek_v32", "deepseek_v32",
@@ -9,6 +9,118 @@ from vllm.triton_utils import tl, triton
float8_info = torch.finfo(current_platform.fp8_dtype()) float8_info = torch.finfo(current_platform.fp8_dtype())
def mask_empty_context(
lse: torch.Tensor,
output: torch.Tensor,
query_start_loc: torch.Tensor,
context_start_loc: torch.Tensor,
) -> None:
"""Neutralize context chunks that cover no keys before merging.
A prefill query whose context chunk is empty attended to no keys, so its
partial attention is undefined: the backend leaves the output rows as
uninitialized scratch (which may hold NaN/Inf) even when it reports an LSE
of -inf. Sanitize both here so ``merge_attn_states`` can stay generic:
force the LSE to -inf (zero softmax weight) and zero the undefined output
rows (so a zero weight cannot combine with NaN/Inf). Emptiness is derived
from the context offsets, not from the -inf LSE, so no merge kernel has to
reason about undefined partials.
Args:
lse: Chunk log-sum-exp, shape [num_heads, num_tokens].
output: Chunk attention output, shape [num_tokens, num_heads, ...].
query_start_loc: Prefill query cumulative offsets, shape [num_reqs + 1].
context_start_loc: Chunk context cumulative offsets,
shape [num_reqs + 1]; an empty chunk has a zero-length span.
"""
num_heads, num_tokens = lse.shape
num_reqs = query_start_loc.shape[0] - 1
block_size = 128
# Reserve the worst-case number of request-local blocks.
num_query_blocks = num_tokens // block_size + num_reqs
is_empty = torch.zeros(num_tokens, dtype=torch.bool, device=lse.device)
mask_empty_context_kernel[(num_query_blocks,)](
lse,
is_empty,
query_start_loc,
context_start_loc,
lse.stride(0),
lse.stride(1),
num_reqs,
NUM_HEADS=num_heads,
BLOCK_SIZE=block_size,
BLOCK_HEADS=8,
num_warps=8,
)
output.masked_fill_(is_empty[:, None, None], 0.0)
@triton.jit
def mask_empty_context_kernel(
lse,
is_empty,
query_start_loc,
context_start_loc,
lse_head_stride,
lse_token_stride,
num_reqs,
NUM_HEADS: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
BLOCK_HEADS: tl.constexpr,
):
query_block_idx = tl.program_id(0)
lanes = tl.arange(0, 32)
chunk_start = 0
req_idx = 0
req_idx_found = False
while (chunk_start < num_reqs) & (not req_idx_found):
req_offsets = chunk_start + lanes
req_mask = req_offsets < num_reqs
query_starts = tl.load(query_start_loc + req_offsets, mask=req_mask)
# Assume the worst-case number of blocks for each request.
req_block_starts = query_starts // BLOCK_SIZE + req_offsets
matched_idx = tl.sum(
(req_mask & (req_block_starts <= query_block_idx)).to(tl.int32)
)
# matched_idx == 32 means the match is past this warp chunk.
req_idx = chunk_start + matched_idx - 1
req_idx_found = matched_idx < 32
chunk_start += 32
query_start = tl.load(query_start_loc + req_idx)
query_end = tl.load(query_start_loc + req_idx + 1)
query_len = query_end - query_start
req_first_block = query_start // BLOCK_SIZE + req_idx
block_in_req = query_block_idx - req_first_block
token_offset = block_in_req * BLOCK_SIZE
if token_offset >= query_len:
return
context_start = tl.load(context_start_loc + req_idx)
context_end = tl.load(context_start_loc + req_idx + 1)
if context_start != context_end:
return
token_offsets = token_offset + tl.arange(0, BLOCK_SIZE)
token_indices = query_start + token_offsets
token_lse_offsets = token_indices * lse_token_stride
valid_tokens = token_offsets < query_len
tl.store(is_empty + token_indices, True, mask=valid_tokens)
head_offsets = tl.arange(0, BLOCK_HEADS)
for head_start in range(0, NUM_HEADS, BLOCK_HEADS):
head_indices = head_start + head_offsets
lse_ptrs = (
lse + head_indices[:, None] * lse_head_stride + token_lse_offsets[None, :]
)
valid_heads = head_indices < NUM_HEADS
tl.store(
lse_ptrs,
float("-inf"),
mask=valid_heads[:, None] & valid_tokens[None, :],
)
# Implements section 2.2 of https://www.arxiv.org/pdf/2501.01005 # Implements section 2.2 of https://www.arxiv.org/pdf/2501.01005
# can be used to combine partial attention results (in the split-KV case) # can be used to combine partial attention results (in the split-KV case)
def merge_attn_states( def merge_attn_states(
@@ -136,6 +248,9 @@ def merge_attn_states_kernel(
if OUTPUT_LSE: if OUTPUT_LSE:
out_lse = tl.log(out_se) + max_lse out_lse = tl.log(out_se) + max_lse
# Both sides empty (max_lse == -inf) => undefined merge; keep -inf so
# downstream merges continue to treat the token as empty.
out_lse = tl.where(max_lse == float("-inf"), float("-inf"), out_lse)
tl.store(output_lse + head_idx * num_tokens + token_idx, out_lse) tl.store(output_lse + head_idx * num_tokens + token_idx, out_lse)
p_out = tl.load( p_out = tl.load(
@@ -159,6 +274,10 @@ def merge_attn_states_kernel(
p_scale = p_se / out_se p_scale = p_se / out_se
s_scale = s_se / out_se s_scale = s_se / out_se
out = p_out * p_scale + s_out * s_scale out = p_out * p_scale + s_out * s_scale
# If both sides are empty (max_lse == -inf) the scales are 0/0 = NaN; emit
# zeros rather than NaN. Callers with empty chunks (see mask_empty_context)
# zero those inputs, so this only guards the fully-undefined corner.
out = tl.where(max_lse == float("-inf"), 0.0, out)
if USE_FP8: if USE_FP8:
out = out * (1.0 / tl.load(output_scale)) out = out * (1.0 / tl.load(output_scale))
+3 -3
View File
@@ -145,7 +145,7 @@ class SchedulerInterface(ABC):
self, self,
request_ids: str | Iterable[str] | None, request_ids: str | Iterable[str] | None,
finished_status: "RequestStatus", finished_status: "RequestStatus",
) -> list[tuple[str, int]]: ) -> "list[Request]":
"""Finish the requests in the scheduler's internal queue. If the request """Finish the requests in the scheduler's internal queue. If the request
is not in the queue, this method will do nothing for that request. is not in the queue, this method will do nothing for that request.
@@ -159,8 +159,8 @@ class SchedulerInterface(ABC):
finished_status: The finished status of the given requests. finished_status: The finished status of the given requests.
Returns: Returns:
Tuple of (req_id, client_index) for requests that were aborted. Will not List of requests that were aborted. Will not include any that were
include any that were already finished. already finished.
""" """
raise NotImplementedError raise NotImplementedError
+24 -12
View File
@@ -61,7 +61,7 @@ from vllm.v1.outputs import DraftTokenIds, KVConnectorOutput, ModelRunnerOutput
from vllm.v1.request import Request, RequestStatus, StreamingUpdate from vllm.v1.request import Request, RequestStatus, StreamingUpdate
from vllm.v1.spec_decode.dynamic.utils import build_dynamic_sd_schedule_lookup from vllm.v1.spec_decode.dynamic.utils import build_dynamic_sd_schedule_lookup
from vllm.v1.spec_decode.metrics import SpecDecodingStats from vllm.v1.spec_decode.metrics import SpecDecodingStats
from vllm.v1.structured_output import StructuredOutputManager from vllm.v1.structured_output import StructuredOutputGrammar, StructuredOutputManager
from vllm.v1.utils import record_function_or_nullcontext from vllm.v1.utils import record_function_or_nullcontext
logger = init_logger(__name__) logger = init_logger(__name__)
@@ -202,6 +202,10 @@ class Scheduler(SchedulerInterface):
self.finished_recving_kv_req_ids: set[str] = set() self.finished_recving_kv_req_ids: set[str] = set()
self.failed_recving_kv_req_ids: set[str] = set() self.failed_recving_kv_req_ids: set[str] = set()
# Grammar compilation failures to finish as per-request errors in
# update_from_output.
self.grammar_compile_error_reqs: set[str] = set()
# Encoder-related. # Encoder-related.
# Calculate encoder cache size if applicable # Calculate encoder cache size if applicable
supports_mm_inputs = mm_registry.supports_multimodal_inputs( supports_mm_inputs = mm_registry.supports_multimodal_inputs(
@@ -1729,7 +1733,7 @@ class Scheduler(SchedulerInterface):
struct_output_request = request.structured_output_request struct_output_request = request.structured_output_request
assert struct_output_request is not None assert struct_output_request is not None
grammar = struct_output_request.grammar grammar = struct_output_request.grammar
assert grammar is not None assert isinstance(grammar, StructuredOutputGrammar)
# new_token_ids can be a mixed block of reasoning content, then # new_token_ids can be a mixed block of reasoning content, then
# the reasoning end marker, then the start of the grammar content. # the reasoning end marker, then the start of the grammar content.
# Trim the reasoning content so the grammar only sees grammar content. # Trim the reasoning content so the grammar only sees grammar content.
@@ -1859,10 +1863,16 @@ class Scheduler(SchedulerInterface):
# This is a rare case and unlikely to impact performance. # This is a rare case and unlikely to impact performance.
self.waiting.remove_requests(stopped_preempted_reqs) self.waiting.remove_requests(stopped_preempted_reqs)
error_req_ids = set(self.grammar_compile_error_reqs)
self.grammar_compile_error_reqs.clear()
if failed_kv_load_req_ids and not self.recompute_kv_load_failures: if failed_kv_load_req_ids and not self.recompute_kv_load_failures:
requests = [self.requests[req_id] for req_id in failed_kv_load_req_ids] error_req_ids.update(failed_kv_load_req_ids)
self.finish_requests(failed_kv_load_req_ids, RequestStatus.FINISHED_ERROR)
for request in requests: if error_req_ids:
error_reqs = self.finish_requests(
error_req_ids, RequestStatus.FINISHED_ERROR
)
for request in error_reqs:
outputs[request.client_index].append( outputs[request.client_index].append(
EngineCoreOutput( EngineCoreOutput(
request_id=request.request_id, request_id=request.request_id,
@@ -2092,8 +2102,7 @@ class Scheduler(SchedulerInterface):
# Filter out spec tokens which do not adhere to the grammar. # Filter out spec tokens which do not adhere to the grammar.
if self.structured_output_manager.should_advance(request): if self.structured_output_manager.should_advance(request):
metadata = request.structured_output_request metadata = request.structured_output_request
assert metadata is not None and metadata.grammar is not None spec_token_ids = metadata.grammar.validate_tokens(spec_token_ids) # type: ignore[union-attr]
spec_token_ids = metadata.grammar.validate_tokens(spec_token_ids)
# Pad to original number of spec tokens. # Pad to original number of spec tokens.
num_invalid_tokens = orig_num_spec_tokens - len(spec_token_ids) num_invalid_tokens = orig_num_spec_tokens - len(spec_token_ids)
if num_invalid_tokens: if num_invalid_tokens:
@@ -2134,7 +2143,7 @@ class Scheduler(SchedulerInterface):
def finish_requests( def finish_requests(
self, request_ids: str | Iterable[str] | None, finished_status: RequestStatus self, request_ids: str | Iterable[str] | None, finished_status: RequestStatus
) -> list[tuple[str, int]]: ) -> list[Request]:
"""Handles the finish signal from outside the scheduler. """Handles the finish signal from outside the scheduler.
For example, the API server can abort a request when the client For example, the API server can abort a request when the client
@@ -2143,8 +2152,8 @@ class Scheduler(SchedulerInterface):
If request_ids is None, all requests will be finished. If request_ids is None, all requests will be finished.
Returns: Returns:
Tuple of (req_id, client_index) for requests that were aborted. Will not List of requests that were aborted. Will not include any that were
include any that were already finished. already finished.
""" """
assert RequestStatus.is_finished(finished_status) assert RequestStatus.is_finished(finished_status)
if isinstance(request_ids, str): if isinstance(request_ids, str):
@@ -2193,7 +2202,7 @@ class Scheduler(SchedulerInterface):
request.status = finished_status request.status = finished_status
self._free_request(request, delay_free_blocks=delay_free_blocks) self._free_request(request, delay_free_blocks=delay_free_blocks)
return [(r.request_id, r.client_index) for r in valid_requests] return valid_requests
def _free_request( def _free_request(
self, request: Request, delay_free_blocks: bool = False self, request: Request, delay_free_blocks: bool = False
@@ -2593,7 +2602,10 @@ class Scheduler(SchedulerInterface):
if request.status == RequestStatus.WAITING_FOR_STRUCTURED_OUTPUT_GRAMMAR: if request.status == RequestStatus.WAITING_FOR_STRUCTURED_OUTPUT_GRAMMAR:
structured_output_req = request.structured_output_request structured_output_req = request.structured_output_request
if not (structured_output_req and structured_output_req.grammar): if not structured_output_req or structured_output_req.grammar is None:
return False
if isinstance(structured_output_req.grammar, Exception):
self.grammar_compile_error_reqs.add(request.request_id)
return False return False
request.status = RequestStatus.WAITING request.status = RequestStatus.WAITING
return True return True
+3 -3
View File
@@ -1830,13 +1830,13 @@ class EngineCoreProc(EngineCore):
) -> None: ) -> None:
self._send_finish_outputs_to_client(req_ids, client_index, FinishReason.ERROR) self._send_finish_outputs_to_client(req_ids, client_index, FinishReason.ERROR)
def _send_abort_outputs(self, aborted_reqs: list[tuple[str, int]]) -> None: def _send_abort_outputs(self, aborted_reqs: list[Request]) -> None:
# TODO(nick) this will be moved inside the scheduler # TODO(nick) this will be moved inside the scheduler
if aborted_reqs: if aborted_reqs:
# Map client_index to list of request_ids that belong to that client. # Map client_index to list of request_ids that belong to that client.
by_client = defaultdict[int, set[str]](set) by_client = defaultdict[int, set[str]](set)
for req_id, client_index in aborted_reqs: for request in aborted_reqs:
by_client[client_index].add(req_id) by_client[request.client_index].add(request.request_id)
for client_index, req_ids in by_client.items(): for client_index, req_ids in by_client.items():
self._send_abort_outputs_to_client(list(req_ids), client_index) self._send_abort_outputs_to_client(list(req_ids), client_index)
+25 -14
View File
@@ -164,24 +164,33 @@ class StructuredOutputManager:
else: else:
raise ValueError(f"Unsupported structured output backend: {backend}") raise ValueError(f"Unsupported structured output backend: {backend}")
grammar: Future[StructuredOutputGrammar] | StructuredOutputGrammar
if self._use_async_grammar_compilation: if self._use_async_grammar_compilation:
grammar = self.executor.submit(self._create_grammar, request) grammar = self.executor.submit(self._create_grammar, request)
else: else:
grammar = self._create_grammar(request) # type: ignore[assignment] try:
request.structured_output_request.grammar = grammar # type: ignore[assignment] grammar = self._create_grammar(request)
except Exception as e:
grammar = Future()
grammar.set_exception(e)
request.structured_output_request.grammar = grammar
def _create_grammar(self, request: "Request") -> StructuredOutputGrammar: def _create_grammar(self, request: "Request") -> StructuredOutputGrammar:
key = request.structured_output_request.structured_output_key # type: ignore[union-attr] struct_request = request.structured_output_request
assert struct_request is not None
# Note that the request was validated in the engine core client, # Note that the request was validated in the engine core client,
# so at this point we know it is a supported type of request. # so at this point we know it is a supported type of request. Grammar
# # compilation may still fail; the Future carries that error to the
# TODO: we still need to handle xgrammar compilation failures, # scheduler so it can fail only this request.
# though it should be unlikely as we test that up front as well. try:
request_type, grammar_spec = key request_type, grammar_spec = struct_request.structured_output_key
assert self.backend is not None
assert self.backend is not None return self.backend.compile_grammar(request_type, grammar_spec)
return self.backend.compile_grammar(request_type, grammar_spec) except Exception:
logger.exception(
"Failed to compile grammar for request %s", request.request_id
)
raise
def _fill_bitmasks( def _fill_bitmasks(
self, batch: Iterable[tuple[StructuredOutputGrammar, int, bool]] self, batch: Iterable[tuple[StructuredOutputGrammar, int, bool]]
@@ -244,8 +253,9 @@ class StructuredOutputManager:
structured_output_request = request.structured_output_request structured_output_request = request.structured_output_request
if TYPE_CHECKING: if TYPE_CHECKING:
assert structured_output_request is not None assert structured_output_request is not None
assert structured_output_request.grammar is not None
grammar = structured_output_request.grammar grammar = structured_output_request.grammar
if TYPE_CHECKING:
assert isinstance(grammar, StructuredOutputGrammar)
apply_bitmask = self.should_fill_bitmask(request) apply_bitmask = self.should_fill_bitmask(request)
batch.append((grammar, cumulative_index, apply_bitmask)) batch.append((grammar, cumulative_index, apply_bitmask))
@@ -268,8 +278,9 @@ class StructuredOutputManager:
if TYPE_CHECKING: if TYPE_CHECKING:
assert structured_output_request is not None assert structured_output_request is not None
assert structured_output_request.grammar is not None
grammar = structured_output_request.grammar grammar = structured_output_request.grammar
if TYPE_CHECKING:
assert isinstance(grammar, StructuredOutputGrammar)
apply_bitmask = self.should_fill_bitmask(request) apply_bitmask = self.should_fill_bitmask(request)
reasoner = self._get_reasoner(request) reasoner = self._get_reasoner(request)
+9 -6
View File
@@ -21,7 +21,9 @@ if TYPE_CHECKING:
@dataclasses.dataclass @dataclasses.dataclass
class StructuredOutputRequest: class StructuredOutputRequest:
params: StructuredOutputsParams params: StructuredOutputsParams
_grammar: Future[StructuredOutputGrammar] | StructuredOutputGrammar | None = None _grammar: (
Future[StructuredOutputGrammar] | StructuredOutputGrammar | Exception | None
) = None
reasoning_ended: bool | None = None reasoning_ended: bool | None = None
# Absolute index into the request's all_token_ids of the last reasoning # Absolute index into the request's all_token_ids of the last reasoning
# token (the reasoning-end marker). Tokens at or before this index are # token (the reasoning-end marker). Tokens at or before this index are
@@ -56,6 +58,8 @@ class StructuredOutputRequest:
self.status = RequestStatus.WAITING self.status = RequestStatus.WAITING
except TimeoutError: except TimeoutError:
return False return False
except Exception as e:
self._grammar = e
return True return True
@property @property
@@ -63,11 +67,10 @@ class StructuredOutputRequest:
return self._check_grammar_completion() return self._check_grammar_completion()
@property @property
def grammar(self) -> StructuredOutputGrammar | None: def grammar(self) -> StructuredOutputGrammar | Exception | None:
completed = self._check_grammar_completion() if not self._check_grammar_completion():
return ( return None
cast(StructuredOutputGrammar | None, self._grammar) if completed else None return cast(StructuredOutputGrammar | Exception | None, self._grammar)
)
@grammar.setter @grammar.setter
def grammar( def grammar(