[Bugfix][V1][TurboQuant] Reserve workspace before CUDA graph capture (#44053)

Signed-off-by: Guipeng Zhang <zhangguipeng23z@ict.ac.cn>
Co-authored-by: Codex <codex@openai.com>
Co-authored-by: Michael Goin <mgoin64@gmail.com>
This commit is contained in:
Guipeng Zhang
2026-06-22 15:47:48 -07:00
committed by GitHub
co-authored by Codex Michael Goin
parent ca5b24695b
commit 183b5f27ea
2 changed files with 163 additions and 0 deletions
+124
View File
@@ -6,6 +6,7 @@ Run: .venv/bin/python -m pytest tests/quantization/test_turboquant.py -v
"""
import math
from types import SimpleNamespace
import pytest
import torch
@@ -276,6 +277,129 @@ class TestHybridAttentionIndices:
assert _get_full_attention_layer_indices(mc) == []
class TestTurboQuantWorkspaceReservation:
@staticmethod
def _fake_vllm_config(
*,
max_num_seqs: int = 16,
max_num_batched_tokens: int = 4096,
enable_chunked_prefill: bool = True,
max_model_len: int = 8192,
dtype: torch.dtype = torch.float16,
max_num_kv_splits: int = 4,
):
return SimpleNamespace(
scheduler_config=SimpleNamespace(
max_num_seqs=max_num_seqs,
max_num_batched_tokens=max_num_batched_tokens,
enable_chunked_prefill=enable_chunked_prefill,
),
model_config=SimpleNamespace(
max_model_len=max_model_len,
dtype=dtype,
get_num_attention_heads=lambda parallel_config: 8,
),
parallel_config=SimpleNamespace(
tensor_parallel_size=2,
decode_context_parallel_size=1,
),
attention_config=SimpleNamespace(
tq_max_kv_splits_for_cuda_graph=max_num_kv_splits
),
)
@staticmethod
def _fake_kv_cache_spec():
from vllm.v1.kv_cache_interface import TQFullAttentionSpec
return TQFullAttentionSpec(
block_size=32,
num_kv_heads=4,
head_size=128,
head_size_v=128,
dtype=torch.uint8,
tq_slot_size=102,
)
def test_metadata_builder_reserves_decode_and_continuation_prefill_workspace(
self, monkeypatch
):
from vllm.v1.attention.backends import turboquant_attn
calls = []
class FakeWorkspaceManager:
def get_simultaneous(self, *shapes_and_dtypes):
calls.append(shapes_and_dtypes)
monkeypatch.setattr(
turboquant_attn,
"current_workspace_manager",
lambda: FakeWorkspaceManager(),
)
monkeypatch.setattr(
turboquant_attn,
"is_workspace_manager_initialized",
lambda: True,
)
turboquant_attn.TurboQuantMetadataBuilder(
kv_cache_spec=self._fake_kv_cache_spec(),
layer_names=["layers.0.self_attn.attn"],
vllm_config=self._fake_vllm_config(),
device=torch.device("cuda"),
)
assert calls == [
(
((16, 8, 4, 129), torch.float32),
((16, 8, 128), torch.float16),
((16, 8), torch.float32),
),
(
((1, 4, 8192, 128), torch.float16),
((1, 4, 8192, 128), torch.float16),
),
]
def test_metadata_builder_skips_continuation_prefill_when_disabled(
self, monkeypatch
):
from vllm.v1.attention.backends import turboquant_attn
calls = []
class FakeWorkspaceManager:
def get_simultaneous(self, *shapes_and_dtypes):
calls.append(shapes_and_dtypes)
monkeypatch.setattr(
turboquant_attn,
"current_workspace_manager",
lambda: FakeWorkspaceManager(),
)
monkeypatch.setattr(
turboquant_attn,
"is_workspace_manager_initialized",
lambda: True,
)
turboquant_attn.TurboQuantMetadataBuilder(
kv_cache_spec=self._fake_kv_cache_spec(),
layer_names=["layers.0.self_attn.attn"],
vllm_config=self._fake_vllm_config(enable_chunked_prefill=False),
device=torch.device("cuda"),
)
assert calls == [
(
((16, 8, 4, 129), torch.float32),
((16, 8, 128), torch.float16),
((16, 8), torch.float32),
)
]
# ============================================================================
# Centroids tests (CPU-only)
# ============================================================================
@@ -30,6 +30,7 @@ from vllm.model_executor.layers.quantization.turboquant.centroids import (
get_centroids,
)
from vllm.triton_utils import triton
from vllm.utils.math_utils import round_up
from vllm.v1.attention.backend import (
AttentionBackend,
AttentionCGSupport,
@@ -201,6 +202,44 @@ class TurboQuantMetadataBuilder(AttentionMetadataBuilder[TurboQuantMetadata]):
def __init__(self, kv_cache_spec, layer_names, vllm_config, device):
super().__init__(kv_cache_spec, layer_names, vllm_config, device)
self._init_reorder_batch_threshold(1, supports_spec_as_decode=False)
self._reserve_workspace()
def _reserve_workspace(self) -> None:
if not is_workspace_manager_initialized():
return
scheduler_config = self.vllm_config.scheduler_config
model_config = self.vllm_config.model_config
parallel_config = self.vllm_config.parallel_config
max_num_reqs = scheduler_config.max_num_seqs
num_heads = model_config.get_num_attention_heads(parallel_config)
num_kv_heads = self.kv_cache_spec.num_kv_heads
head_size = self.kv_cache_spec.head_size
max_num_splits = (
self.vllm_config.attention_config.tq_max_kv_splits_for_cuda_graph
)
current_workspace_manager().get_simultaneous(
((max_num_reqs, num_heads, max_num_splits, head_size + 1), torch.float32),
((max_num_reqs, num_heads, head_size), model_config.dtype),
((max_num_reqs, num_heads), torch.float32),
)
reserve_continuation_prefill = (
scheduler_config.enable_chunked_prefill
and scheduler_config.max_num_batched_tokens > _CONTINUATION_DECODE_THRESHOLD
)
if not reserve_continuation_prefill:
return
max_cached_len = max(0, model_config.max_model_len - 1)
alloc_len = round_up(max_cached_len, self.kv_cache_spec.block_size)
cache_buf_shape = (1, num_kv_heads, alloc_len, head_size)
current_workspace_manager().get_simultaneous(
(cache_buf_shape, torch.float16),
(cache_buf_shape, torch.float16),
)
def build_for_cudagraph_capture(
self, common_attn_metadata: CommonAttentionMetadata