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
486 changes: 439 additions & 47 deletions cpp/tensorrt_llm/kernels/minimaxM3SelectBlocks.cu

Large diffs are not rendered by default.

2 changes: 1 addition & 1 deletion cpp/tensorrt_llm/kernels/minimaxM3SelectBlocks.h
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@ namespace kernels

void invokeMinimaxM3SelectBlocks(float const* scores, int64_t headStride, int64_t blockStride, int64_t queryStride,
int32_t const* nValidBlocks, int32_t* output, int32_t numKvHeads, int32_t numBlocks, int32_t totalQueries,
int32_t initBlocks, int32_t localBlocks, cudaStream_t stream);
int32_t initBlocks, int32_t localBlocks, bool headMajorOutput, cudaStream_t stream);

} // namespace kernels

Expand Down
11 changes: 7 additions & 4 deletions cpp/tensorrt_llm/thop/minimaxM3SelectBlocksOp.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@ namespace torch_ext
{

torch::Tensor minimaxM3SelectBlocks(torch::Tensor const& scores, torch::Tensor const& nValidBlocks, int64_t topK,
int64_t initBlocks, int64_t localBlocks)
int64_t initBlocks, int64_t localBlocks, bool headMajorOutput)
{
constexpr int64_t kRequiredTopK = 16;
constexpr int64_t kMaxBlockIndex = 65'535;
Expand Down Expand Up @@ -61,13 +61,16 @@ torch::Tensor minimaxM3SelectBlocks(torch::Tensor const& scores, torch::Tensor c
TORCH_CHECK(initBlocks <= std::numeric_limits<int32_t>::max() && localBlocks <= std::numeric_limits<int32_t>::max(),
"minimax_m3_select_blocks forcing ranges exceed int32 range");

auto output = torch::empty({scores.size(2), scores.size(0), topK}, scores.options().dtype(torch::kInt32));
auto const outputOptions = scores.options().dtype(torch::kInt32);
auto output = headMajorOutput
? torch::empty({scores.size(0), scores.size(2), topK}, outputOptions).permute({1, 0, 2})
: torch::empty({scores.size(2), scores.size(0), topK}, outputOptions);
auto const stream = at::cuda::getCurrentCUDAStream(scores.get_device());
tensorrt_llm::kernels::invokeMinimaxM3SelectBlocks(scores.data_ptr<float>(), scores.stride(0), scores.stride(1),
scores.stride(2), nValidBlocks.data_ptr<int32_t>(), output.data_ptr<int32_t>(),
static_cast<int32_t>(scores.size(0)), static_cast<int32_t>(scores.size(1)),
static_cast<int32_t>(scores.size(2)), static_cast<int32_t>(initBlocks), static_cast<int32_t>(localBlocks),
stream);
headMajorOutput, stream);
return output;
}

Expand All @@ -79,7 +82,7 @@ TORCH_LIBRARY_FRAGMENT(trtllm, m)
{
m.def(
"minimax_m3_select_blocks(Tensor scores, Tensor n_valid_blocks, int topk, int init_blocks, int "
"local_blocks) -> Tensor");
"local_blocks, bool head_major_output=False) -> Tensor");
}

TORCH_LIBRARY_IMPL(trtllm, CUDA, m)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1082,6 +1082,9 @@ def run_indexer(
config = self.m3_config
idx_sm_scale = idx_sm_scale if idx_sm_scale is not None else config.sparse_index_dim**-0.5
num_tokens = int(idx_q.shape[0])
head_major_output = (
int(metadata.num_contexts or 0) > 0 and int(metadata.num_generations or 0) == 0
)
# idx_q and idx_k may be strided column-views of a fused buffer, so
# reshape to keep them zero-copy. The proxy fmha_sm100 and the index-K
# scatter below both honor the source strides.
Expand Down Expand Up @@ -1131,12 +1134,18 @@ def run_indexer(
max_score = None
# The empty case was resolved on the host, so short-circuit here.
if metadata.msa_eager_all_blocks_empty:
return torch.full(
(num_tokens, config.num_kv_heads, MSA_REQUIRED_TOPK),
output_shape = (
(config.num_kv_heads, num_tokens, MSA_REQUIRED_TOPK)
if head_major_output
else (num_tokens, config.num_kv_heads, MSA_REQUIRED_TOPK)
)
output = torch.full(
output_shape,
-1,
dtype=torch.int32,
device=idx_q.device,
)
return output.permute(1, 0, 2) if head_major_output else output
n_valid_blocks = metadata.msa_eager_n_valid_blocks
if n_valid_blocks is not None:
n_valid_blocks = n_valid_blocks[:num_tokens]
Expand All @@ -1151,6 +1160,7 @@ def run_indexer(
proxy_plan=proxy_plan,
max_score=max_score,
n_valid_blocks=n_valid_blocks,
head_major_output=head_major_output,
)

def sparse_attn_predict(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -131,6 +131,7 @@ def select_blocks(
proxy_plan: Optional[tuple] = None,
max_score: Optional[torch.Tensor] = None,
n_valid_blocks: Optional[torch.Tensor] = None,
head_major_output: bool = False,
) -> torch.Tensor:
"""Return [total_q, num_kv_heads, topk] selected block indices.

Expand Down Expand Up @@ -180,18 +181,25 @@ def select_blocks(
# Empty-selection guard. n_valid_blocks is a host tensor on
# this path, so the .item() read does not sync the device.
if n_valid_blocks.numel() == 0 or int(n_valid_blocks.max().item()) <= 0:
return torch.full(
(idx_q.shape[0], config.num_kv_heads, MSA_REQUIRED_TOPK),
output_shape = (
(config.num_kv_heads, idx_q.shape[0], MSA_REQUIRED_TOPK)
if head_major_output
else (idx_q.shape[0], config.num_kv_heads, MSA_REQUIRED_TOPK)
)
output = torch.full(
output_shape,
-1,
dtype=torch.int32,
device=idx_q.device,
)
return output.permute(1, 0, 2) if head_major_output else output
return select_blocks_from_maxscore(
max_score_kv,
topk=MSA_REQUIRED_TOPK,
n_valid_blocks=n_valid_blocks,
init_blocks=config.init_blocks,
local_blocks=config.local_blocks,
head_major_output=head_major_output,
)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -215,17 +215,25 @@ def select_blocks_from_maxscore(
n_valid_blocks: torch.Tensor,
init_blocks: int,
local_blocks: int,
head_major_output: bool = False,
) -> torch.Tensor:
"""Select per-query top-k blocks from per-KV-head block scores.

Applies init and local forced blocks and per-query valid-block masking
on the amax-reduced scores [num_kv_heads, n_blocks, total_q]. Returns
[total_q, num_kv_heads, topk] int32 ascending block ids with -1 tail
padding.
padding. When ``head_major_output`` is set, the logical result uses a
head-major backing so ``result.permute(1, 0, 2)`` is contiguous without a
copy.
"""
nvb = n_valid_blocks.to(device=max_score_kv.device, dtype=torch.int32).contiguous()
return torch.ops.trtllm.minimax_m3_select_blocks(
max_score_kv, nvb, topk, init_blocks, local_blocks
max_score_kv,
nvb,
topk,
init_blocks,
local_blocks,
head_major_output,
)


Expand Down
12 changes: 11 additions & 1 deletion tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -277,8 +277,18 @@ def _(logits,
pass

@torch.library.register_fake("trtllm::minimax_m3_select_blocks")
def _(scores, n_valid_blocks, topk, init_blocks, local_blocks):
def _(
scores,
n_valid_blocks,
topk,
init_blocks,
local_blocks,
head_major_output=False,
):
del n_valid_blocks, init_blocks, local_blocks
if head_major_output:
return scores.new_empty((scores.shape[0], scores.shape[2], topk),
dtype=torch.int32).permute(1, 0, 2)
return scores.new_empty((scores.shape[2], scores.shape[0], topk),
dtype=torch.int32)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,8 @@
reference is covered by the SM100 integration accuracy test.
"""

from types import SimpleNamespace

import pytest
import torch

Expand Down Expand Up @@ -252,6 +254,67 @@ def msa_idx_k_cache(self, layer_idx):
assert "write" not in captured


@pytest.mark.parametrize(
("num_contexts", "num_generations", "expected_head_major"),
[(2, 0, True), (1, 1, False), (0, 2, False)],
)
def test_run_indexer_routes_head_major_output_by_batch_mode(
num_contexts, num_generations, expected_head_major
):
num_tokens, num_index_heads, sparse_index_dim = 3, 4, 128
captured = {}

class FakeIndexer:
def select_blocks(self, *args, **kwargs):
del args
captured["head_major_output"] = kwargs["head_major_output"]
return torch.zeros(num_tokens, 1, 16, dtype=torch.int32)

class FakeMetadata:
msa_decode_proxy_plan = None
msa_eager_proxy_plan = ("eager",)
msa_eager_all_blocks_empty = False
msa_eager_n_valid_blocks = torch.ones(num_tokens, dtype=torch.int32)
msa_kv_indices = torch.arange(num_tokens, dtype=torch.int32)
msa_qo_lens_cpu = torch.tensor([num_tokens], dtype=torch.int32)
msa_kv_lens_cpu = torch.tensor([num_tokens], dtype=torch.int32)
msa_qo_offset_cpu = torch.tensor([0], dtype=torch.int32)

def __init__(self):
self.num_contexts = num_contexts
self.num_generations = num_generations
self.idx_k_cache = None

def msa_write_idx_k(self, layer_idx, idx_k):
del layer_idx
self.idx_k_cache = idx_k

def msa_idx_k_cache(self, layer_idx):
del layer_idx
return self.idx_k_cache

attention = SimpleNamespace(
layer_idx=0,
m3_config=SimpleNamespace(
sparse_index_dim=sparse_index_dim,
num_index_heads=num_index_heads,
num_kv_heads=1,
),
indexer=FakeIndexer(),
)
metadata = FakeMetadata()

result = MiniMaxM3MsaSparseAttention.run_indexer(
attention,
torch.zeros(num_tokens, num_index_heads * sparse_index_dim),
torch.zeros(num_tokens, sparse_index_dim),
metadata,
)

assert result.shape == (num_tokens, 1, 16)
assert captured["head_major_output"] is expected_head_major


def test_msa_proxy_max_score_strided_index_k_matches_packed():
if not torch.cuda.is_available():
pytest.skip("CUDA required")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,9 @@ def _reference_select_blocks(


@pytest.mark.parametrize("num_kv_heads", [1, 4])
@pytest.mark.parametrize("num_blocks", [1, 8, 16, 17, 127, 1024, 1537])
@pytest.mark.parametrize(
"num_blocks", [1, 8, 16, 17, 32, 33, 64, 65, 96, 127, 128, 129, 1024, 1537]
)
def test_fused_selector_matches_reference_random(num_kv_heads, num_blocks):
total_q = 19
generator = torch.Generator(device="cuda").manual_seed(num_blocks)
Expand Down Expand Up @@ -95,7 +97,82 @@ def test_fused_selector_matches_reference_random(num_kv_heads, num_blocks):

assert actual.dtype == torch.int32
assert actual.shape == (total_q, num_kv_heads, 16)
assert actual.stride() == (num_kv_heads * 16, 16, 1)
assert actual.is_contiguous()
assert torch.equal(actual, expected)


@pytest.mark.parametrize("num_blocks", [65, 96, 128])
def test_fused_selector_128_path_matches_reference_ties_and_forcing(num_blocks):
scores = torch.zeros((2, num_blocks, 5), device="cuda", dtype=torch.float32)
n_valid_blocks = torch.tensor(
[0, 8, 16, num_blocks - 1, num_blocks], device="cuda", dtype=torch.int32
)

expected = _reference_select_blocks(
scores,
topk=16,
n_valid_blocks=n_valid_blocks,
init_blocks=8,
local_blocks=12,
)
actual = select_blocks_from_maxscore(
scores,
topk=16,
n_valid_blocks=n_valid_blocks,
init_blocks=8,
local_blocks=12,
)

assert torch.equal(actual, expected)


@pytest.mark.parametrize("num_blocks", [16, 48, 96, 129])
def test_fused_selector_head_major_output_is_zero_copy_q2k(num_blocks):
total_q, num_kv_heads = 7, 4
generator = torch.Generator(device="cuda").manual_seed(num_blocks)
scores = torch.randn(
num_kv_heads,
num_blocks,
total_q,
generator=generator,
device="cuda",
dtype=torch.float32,
)
n_valid_blocks = torch.randint(
0,
num_blocks + 1,
(total_q,),
generator=generator,
device="cuda",
dtype=torch.int32,
)

expected = _reference_select_blocks(
scores,
topk=16,
n_valid_blocks=n_valid_blocks,
init_blocks=8,
local_blocks=12,
)
actual = select_blocks_from_maxscore(
scores,
topk=16,
n_valid_blocks=n_valid_blocks,
init_blocks=8,
local_blocks=12,
head_major_output=True,
)
q2k = actual.permute(1, 0, 2).contiguous().to(torch.int32)

assert torch.equal(actual, expected)
assert actual.shape == (total_q, num_kv_heads, 16)
assert actual.stride() == (16, total_q * 16, 1)
assert not actual.is_contiguous()
assert q2k.shape == (num_kv_heads, total_q, 16)
assert q2k.is_contiguous()
assert q2k.data_ptr() == actual.data_ptr()
assert q2k.untyped_storage().data_ptr() == actual.untyped_storage().data_ptr()


@pytest.mark.parametrize(
Expand Down
Loading