Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
187 changes: 96 additions & 91 deletions tests/kernels/attention/test_deepgemm_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
fp8_fp4_paged_mqa_logits,
get_num_sms,
get_paged_mqa_logits_metadata,
native_next_n_supported,
)
from vllm.utils.import_utils import has_deep_gemm
from vllm.utils.math_utils import cdiv
Expand Down Expand Up @@ -205,106 +206,110 @@ def _ref_fp8_fp4_paged_mqa_logits(
@pytest.mark.skipif(
not current_platform.has_device_capability(90), reason="SM90 and SM100 only"
)
def test_deepgemm_fp8_fp4_paged_mqa_logits():
# next_n = 1 + num_speculative_tokens, so next_n=4 is MTP=3 (issue #35878).
@pytest.mark.parametrize("batch_size,next_n", [(4, 1), (2, 2), (2, 4)])
def test_deepgemm_fp8_fp4_paged_mqa_logits(batch_size: int, next_n: int):
if not native_next_n_supported(next_n):
pytest.skip(f"next_n={next_n} has no native kernel on this architecture")

# NOTE: clean_logits=True is incompatible with the 2D context_lens
# required by csrc/apis/attention.hpp; only the False path is exercised.
clean_logits = False
torch.manual_seed(0)
random.seed(0)

max_model_len = 4096
for batch_size, next_n in [(4, 1), (2, 2)]:
for heads, index_dim in [(32, 128)]:
for avg_kv in (2048,):
num_blocks, blocksize = max_model_len * 2, 64

q = torch.randn(
(batch_size, next_n, heads, index_dim),
device="cuda",
dtype=torch.bfloat16,
)
kv_cache = torch.randn(
(num_blocks, blocksize, 1, index_dim),
device="cuda",
dtype=torch.bfloat16,
)
weights = torch.randn(
(batch_size * next_n, heads),
device="cuda",
dtype=torch.float32,
)

context_lens = (
torch.randint(int(0.8 * avg_kv), int(1.2 * avg_kv), (batch_size,))
.cuda()
.to(torch.int32)
)
max_block_len = (
(context_lens.max().item() + blocksize - 1) // blocksize * blocksize
)
block_tables = torch.zeros(
(batch_size, max_block_len),
device="cuda",
dtype=torch.int32,
)
for heads, index_dim in [(32, 128)]:
for avg_kv in (2048,):
num_blocks, blocksize = max_model_len * 2, 64

counter = 0
block_idx_pool = list(range(num_blocks))
random.shuffle(block_idx_pool)
for i in range(batch_size):
ctx_len = int(context_lens[i].item())
for j in range((ctx_len + blocksize - 1) // blocksize):
block_tables[i][j] = block_idx_pool[counter]
counter += 1
q = torch.randn(
(batch_size, next_n, heads, index_dim),
device="cuda",
dtype=torch.bfloat16,
)
kv_cache = torch.randn(
(num_blocks, blocksize, 1, index_dim),
device="cuda",
dtype=torch.bfloat16,
)
weights = torch.randn(
(batch_size * next_n, heads),
device="cuda",
dtype=torch.float32,
)

q_fp8 = q.to(torch.float8_e4m3fn)
kv_cache_fp8 = kv_cache_cast_to_fp8(kv_cache)

# deep_gemm paged MQA logits requires 2D context_lens of
# shape (B, next_n) (csrc/apis/attention.hpp:332-335);
# see indexer.py:607-608. For each batch/next_n token, the
# effective context length is context_lens[b] - next_n + j + 1.
next_n_arange = torch.arange(next_n, device="cuda", dtype=torch.int32)
context_lens_2d = (
context_lens.unsqueeze(-1) - next_n + 1 + next_n_arange
).contiguous()
schedule_metadata = get_paged_mqa_logits_metadata(
context_lens_2d, blocksize, get_num_sms()
)
logits = fp8_fp4_paged_mqa_logits(
(q_fp8, None),
kv_cache_fp8,
weights,
context_lens_2d,
block_tables,
schedule_metadata,
max_model_len,
clean_logits=clean_logits,
)
context_lens = (
torch.randint(int(0.8 * avg_kv), int(1.2 * avg_kv), (batch_size,))
.cuda()
.to(torch.int32)
)
max_block_len = (
(context_lens.max().item() + blocksize - 1) // blocksize * blocksize
)
block_tables = torch.zeros(
(batch_size, max_block_len),
device="cuda",
dtype=torch.int32,
)

ref_logits = _ref_fp8_fp4_paged_mqa_logits(
q,
kv_cache,
weights,
context_lens,
block_tables,
max_model_len,
)
counter = 0
block_idx_pool = list(range(num_blocks))
random.shuffle(block_idx_pool)
for i in range(batch_size):
ctx_len = int(context_lens[i].item())
for j in range((ctx_len + blocksize - 1) // blocksize):
block_tables[i][j] = block_idx_pool[counter]
counter += 1

q_fp8 = q.to(torch.float8_e4m3fn)
kv_cache_fp8 = kv_cache_cast_to_fp8(kv_cache)

# deep_gemm paged MQA logits requires 2D context_lens of
# shape (B, next_n) (csrc/apis/attention.hpp:332-335);
# see indexer.py:607-608. For each batch/next_n token, the
# effective context length is context_lens[b] - next_n + j + 1.
next_n_arange = torch.arange(next_n, device="cuda", dtype=torch.int32)
context_lens_2d = (
context_lens.unsqueeze(-1) - next_n + 1 + next_n_arange
).contiguous()
schedule_metadata = get_paged_mqa_logits_metadata(
context_lens_2d,
blocksize,
get_num_sms(),
)
logits = fp8_fp4_paged_mqa_logits(
(q_fp8, None),
kv_cache_fp8,
weights,
context_lens_2d,
block_tables,
schedule_metadata,
max_model_len,
clean_logits=clean_logits,
)

positions = (
torch.arange(max_model_len, device="cuda")
.unsqueeze(0)
.expand(batch_size * next_n, -1)
)
row_indices = torch.arange(batch_size * next_n, device="cuda") // next_n
next_n_offset = (
torch.arange(batch_size * next_n, device="cuda") % next_n
)
mask = positions <= (
context_lens[row_indices] - next_n + next_n_offset
).unsqueeze(1)
ref_logits = _ref_fp8_fp4_paged_mqa_logits(
q,
kv_cache,
weights,
context_lens,
block_tables,
max_model_len,
)

logits = logits.masked_fill(~mask, 0)
ref_logits = ref_logits.masked_fill(~mask, 0)
diff = calc_diff(logits, ref_logits)
assert diff < 1e-3, f"{diff=}"
positions = (
torch.arange(max_model_len, device="cuda")
.unsqueeze(0)
.expand(batch_size * next_n, -1)
)
row_indices = torch.arange(batch_size * next_n, device="cuda") // next_n
next_n_offset = torch.arange(batch_size * next_n, device="cuda") % next_n
mask = positions <= (
context_lens[row_indices] - next_n + next_n_offset
).unsqueeze(1)

logits = logits.masked_fill(~mask, 0)
ref_logits = ref_logits.masked_fill(~mask, 0)
diff = calc_diff(logits, ref_logits)
assert diff < 1e-3, f"{diff=}"
75 changes: 75 additions & 0 deletions tests/v1/attention/test_indexer_native_next_n.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Which next_n the DSA indexer decode path may hand to DeepGEMM unflattened.

Getting this wrong is not a slow path but a crash: `fp8_fp4_paged_mqa_logits`
asserts both that the architecture implements the requested `next_n` and that
the schedule metadata was sized for the matching slot count.
"""

import pytest

from vllm.platforms import current_platform
from vllm.utils.deep_gemm import _paged_mqa_logits_schedule_slots
from vllm.v1.attention.backends.mla import indexer

NUM_SMS = 114 # H100 PCIe


def _set_arch(monkeypatch, family: int, *, cuda: bool = True, deep_gemm: bool = True):
monkeypatch.setattr(current_platform, "is_cuda", lambda: cuda)
monkeypatch.setattr(
current_platform,
"is_device_capability_family",
lambda capability, device_id=0: capability // 10 == family,
)
monkeypatch.setattr(indexer, "has_deep_gemm", lambda: deep_gemm)


@pytest.mark.parametrize(
"family,expected_native",
[
# SM90 gained next_n=4 (MTP=3) via 2-CTA multicast, but never 3.
(9, {1, 2, 4}),
# SM100 schedules any next_n with multi-atom tiles.
(10, {1, 2, 3, 4, 5, 8}),
# SM120 advertises multi-atom too but is unvalidated on hardware, so
# it stays on the conservative gate. Loosen it only with measurements.
(12, {1, 2}),
],
)
def test_native_decode_gate_per_architecture(monkeypatch, family, expected_native):
_set_arch(monkeypatch, family)
for next_n in (1, 2, 3, 4, 5, 8):
assert indexer._supports_native_decode(next_n) == (next_n in expected_native), (
f"family={family} next_n={next_n}"
)


@pytest.mark.parametrize(
"cuda,deep_gemm", [(False, True), (True, False), (False, False)]
)
def test_native_decode_gate_without_deepgemm(monkeypatch, cuda, deep_gemm):
"""Without the DeepGEMM kernels only the shapes every backend handles."""
_set_arch(monkeypatch, 9, cuda=cuda, deep_gemm=deep_gemm)
assert [indexer._supports_native_decode(n) for n in (1, 2, 3, 4)] == [
True,
True,
False,
False,
]


def test_sm90_next_n_4_halves_the_schedule_slots(monkeypatch):
"""SM90 next_n=4 runs one scheduler task per 2-CTA cluster, not per SM."""
_set_arch(monkeypatch, 9)
assert _paged_mqa_logits_schedule_slots(NUM_SMS, 4) == NUM_SMS // 2
for next_n in (1, 2, 3):
assert _paged_mqa_logits_schedule_slots(NUM_SMS, next_n) == NUM_SMS


@pytest.mark.parametrize("family", [10, 12])
def test_multicast_is_sm90_only(monkeypatch, family):
_set_arch(monkeypatch, family)
for next_n in (1, 2, 3, 4):
assert _paged_mqa_logits_schedule_slots(NUM_SMS, next_n) == NUM_SMS
36 changes: 31 additions & 5 deletions vllm/utils/deep_gemm.py
Original file line number Diff line number Diff line change
Expand Up @@ -552,6 +552,29 @@ def fp8_fp4_mqa_logits(
)


def native_next_n_supported(next_n: int) -> bool:
"""Whether the paged MQA logits kernel takes `next_n` Q rows per request.

SM90 implements only {1, 2, 4}; SM100 and SM120 schedule any `next_n` via
multi-atom tiles. Unsupported values must be flattened to one row per query.
"""
if current_platform.is_device_capability_family(90):
return next_n in (1, 2, 4)
return True


def _paged_mqa_logits_schedule_slots(num_sms: int, next_n: int) -> int:
"""Scheduler tasks the paged MQA logits kernel launches.

SM90 `next_n=4` runs one task per 2-CTA multicast cluster rather than per
SM, and `fp8_fp4_paged_mqa_logits` asserts its metadata is sized to match.
"""
num_kv_multicast = (
2 if next_n == 4 and current_platform.is_device_capability_family(90) else 1
)
return num_sms // num_kv_multicast


def get_paged_mqa_logits_metadata(
context_lens: torch.Tensor,
block_size: int,
Expand All @@ -561,22 +584,24 @@ def get_paged_mqa_logits_metadata(
"""Build scheduling metadata for paged MQA logits.

Args:
context_lens: Tensor of shape [B], dtype int32; effective context length
per batch element.
context_lens: Tensor of shape [B, next_n], dtype int32; effective
context length per Q row.
block_size: KV-cache block size in tokens (e.g., 64).
num_sms: Number of SMs available. 132 for Hopper
indices: Optional request index for each varlen row.

Returns:
Backend-specific tensor consumed by `fp8_fp4_paged_mqa_logits` to
schedule work across SMs.
Tensor of shape [slots + 1, 2] consumed by `fp8_fp4_paged_mqa_logits`
to schedule work across SMs.
"""
_lazy_init()
if _get_paged_mqa_logits_metadata_impl is None:
return _missing()
next_n = context_lens.shape[1] if context_lens.dim() == 2 else 1
num_slots = _paged_mqa_logits_schedule_slots(num_sms, next_n)
kwargs = {} if indices is None else {"indices": indices}
return _get_paged_mqa_logits_metadata_impl(
context_lens, block_size, num_sms, **kwargs
context_lens, block_size, num_slots, **kwargs
)


Expand Down Expand Up @@ -752,6 +777,7 @@ def should_use_deepgemm_for_fp8_linear(
"fp8_fp4_mqa_logits",
"fp8_fp4_paged_mqa_logits",
"get_paged_mqa_logits_metadata",
"native_next_n_supported",
"per_block_cast_to_fp8",
"is_deep_gemm_e8m0_used",
"is_deep_gemm_supported",
Expand Down
Loading
Loading