diff --git a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/msa_utils.py b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/msa_utils.py index b3aa264d893e..774c21d88126 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/msa_utils.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/msa_utils.py @@ -23,6 +23,21 @@ MSA_REQUIRED_HEAD_DIM = 128 +def check_decode_span_shape(kernel: str, total_q: int, batch: int, query_len: int) -> None: + """Reject a q that does not cover exactly the batch it was handed. + + The decode kernels derive the request id as token // query_len, so a longer + q reads page table rows and lengths past the batch's last one. The caller + names itself so the error does too, rather than surfacing as an assert + several frames inside a kernel. + """ + if total_q != batch * query_len: + raise ValueError( + f"{kernel}: total_q ({total_q}) must be batch ({batch}) * " + f"decode_query_len ({query_len})." + ) + + def is_msa_layer(attn) -> bool: """Whether this layer's attention is served by the MiniMax-M3 MSA kernels.""" sparse_params = attn.sparse_params @@ -103,6 +118,7 @@ def write_msa_main_kv( out_cache_loc: torch.Tensor, k: torch.Tensor, v: torch.Tensor, + num_live_tokens: int, ) -> None: """Write new-token K and V into the paged main cache at out_cache_loc. @@ -116,10 +132,18 @@ def write_msa_main_kv( head_dim = int(k_view.shape[3]) num_tokens = int(k.shape[0]) write_kv_slots( - k_view, out_cache_loc, k.reshape(num_tokens, num_kv_heads, head_dim), layout="HND" + k_view, + out_cache_loc, + k.reshape(num_tokens, num_kv_heads, head_dim), + num_live_tokens, + layout="HND", ) write_kv_slots( - v_view, out_cache_loc, v.reshape(num_tokens, num_kv_heads, head_dim), layout="HND" + v_view, + out_cache_loc, + v.reshape(num_tokens, num_kv_heads, head_dim), + num_live_tokens, + layout="HND", ) @@ -163,6 +187,8 @@ def write_msa_phase_kv( metadata.msa_out_cache_loc[token_offset : token_offset + num_tokens], k, v, + # k and v are this phase's own token slice, so every row owns a slot. + num_tokens, ) @@ -268,6 +294,7 @@ def select_blocks_from_maxscore( "MSA_REQUIRED_HEAD_DIM", "MSA_REQUIRED_TOPK", "build_kv_page_indices", + "check_decode_span_shape", "msa_package_available", "msa_paged_kv", "per_token_valid_blocks", diff --git a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/paged_cache.py b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/paged_cache.py index 3a20070d56c3..b3cbb2716240 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/paged_cache.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/paged_cache.py @@ -13,6 +13,7 @@ def write_kv_slots( cache: torch.Tensor, out_cache_loc: torch.Tensor, values: torch.Tensor, + num_live_tokens: int, *, layout: Literal["NHD", "HND"] = "NHD", ) -> None: @@ -28,7 +29,27 @@ def write_kv_slots( mapping satisfies this contract: ``get_block_ids_per_seq`` canonicalizes padded ``BAD_PAGE_INDEX`` entries before ``build_paged_kv_slot_mapping`` selects only the allocated live-token positions. + + `num_live_tokens` is how many leading rows own a real cache slot; the rest + are dropped. It carries no default because a caller passing the padded + token extent of a piecewise CUDA graph corrupts the cache silently: torch + wraps a negative index, so the -1 sentinel past the live count lands in the + last page instead of raising. """ + # Trimming by count keeps this sync-free, the sentinel tail being contiguous + # by construction, where masking on the slot values would not. + if num_live_tokens < 0: + raise ValueError(f"num_live_tokens must be non-negative, got {num_live_tokens}") + if out_cache_loc.shape[0] < num_live_tokens or values.shape[0] < num_live_tokens: + raise ValueError( + f"num_live_tokens={num_live_tokens} exceeds the rows supplied " + f"(out_cache_loc={out_cache_loc.shape[0]}, values={values.shape[0]})" + ) + if num_live_tokens == 0: + return + out_cache_loc = out_cache_loc[:num_live_tokens] + values = values[:num_live_tokens] + with torch.no_grad(): if cache.ndim >= 4: token_axis = 2 if layout == "HND" else 1 diff --git a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/triton_sparse_decode.py b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/triton_sparse_decode.py index a23234636d5c..805cccc28a41 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/triton_sparse_decode.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/triton_sparse_decode.py @@ -45,6 +45,8 @@ from tensorrt_llm._torch.memory_buffer_utils import get_memory_buffers +from .msa_utils import check_decode_span_shape + # One sparse block is exactly one KV page. SPARSE_BLOCK_SIZE = 128 @@ -1048,11 +1050,9 @@ def minimax_m3_sparse_attn_decode( """ total_q, num_heads, head_dim = q.shape num_kv_heads = int(k_paged.shape[1]) - if total_q != int(seq_lens.shape[0]) * decode_query_len: - raise ValueError( - f"total_q ({total_q}) must be batch ({int(seq_lens.shape[0])}) * " - f"decode_query_len ({decode_query_len})." - ) + check_decode_span_shape( + "MiniMax-M3 Triton sparse decode", total_q, int(seq_lens.shape[0]), decode_query_len + ) if int(k_paged.shape[2]) != SPARSE_BLOCK_SIZE: raise ValueError( f"MiniMax-M3 sparse decode requires page_size={SPARSE_BLOCK_SIZE}; " diff --git a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/trtllm_gen_dense_decode.py b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/trtllm_gen_dense_decode.py index afaad2ba8430..452863586356 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/trtllm_gen_dense_decode.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/kernels/trtllm_gen_dense_decode.py @@ -27,6 +27,8 @@ from tensorrt_llm._torch.memory_buffer_utils import get_memory_buffers +from .msa_utils import check_decode_span_shape + @functools.lru_cache(maxsize=None) def _counter_size(num_heads: int, max_num_requests: int, device_index: int) -> int: @@ -258,6 +260,15 @@ def minimax_m3_trtllm_gen_dense_decode( """ import flashinfer + # The multi-CTA KV counters are sized against max_num_requests, so a batch + # read out of a longer q would undersize them. + check_decode_span_shape( + "MiniMax-M3 trtllm-gen dense decode", + int(q.shape[0]), + int(seq_lens.shape[0]), + decode_query_len, + ) + kv_pool, subpages_per_slot = kv_cache_manager.get_kv_subpage_pool(layer_idx, "HND") num_heads = int(q.shape[1]) diff --git a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/msa_backend.py b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/msa_backend.py index 250a77ee50ae..e8af33cbbc68 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/msa_backend.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/msa_backend.py @@ -1044,6 +1044,15 @@ def _build_msa_fields(self) -> None: f"MSA out_cache_loc buffer ({self.msa_out_cache_loc.shape[0]}) is " f"smaller than the step's new-token count ({total_new_tokens})." ) + # The cache writers trim to num_tokens, so that count and the mapping + # have to describe the same rows. The mapping emits one slot per new + # token, so they agree unless a caller staged lengths this metadata's + # seq_lens does not match. + if total_new_tokens != int(self.num_tokens): + raise ValueError( + f"MSA slot mapping covers {total_new_tokens} new tokens, but the " + f"step's token count is {int(self.num_tokens)}." + ) if kv_indices is not None and int(kv_indices.shape[0]) > self.msa_kv_indices.shape[0]: raise ValueError( f"MSA kv_indices buffer ({self.msa_kv_indices.shape[0]}) is " @@ -1058,8 +1067,11 @@ def _build_msa_fields(self) -> None: ) self.msa_out_cache_loc[:total_new_tokens].copy_(out_cache_loc, non_blocking=True) - # Captured producers also execute padded rows. Invalidate only the - # unwritten tail so they cannot reuse the previous step's live slots. + # Captured producers also execute padded rows: the fused index producer + # sits inside the captured region, so trimming it to a host-side count + # would make its shape dynamic, and a negative slot is what makes those + # rows cache-write no-ops instead. Invalidate the unwritten tail so they + # cannot reuse the previous step's live slots, which address real pages. if total_new_tokens < self.msa_out_cache_loc.shape[0]: self.msa_out_cache_loc[total_new_tokens:].fill_(-1) if kv_indices is not None: @@ -1168,6 +1180,9 @@ def msa_write_idx_k(self, layer_idx: int, idx_k: torch.Tensor) -> None: cache, self.msa_out_cache_loc[:num_tokens], idx_k.reshape(num_tokens, 1, sparse_index_dim), + # idx_k arrives over the padded token extent; the live prefix is + # where msa_out_cache_loc stops holding real slots. + int(self.num_tokens), layout="HND", ) @@ -1349,16 +1364,21 @@ def write_layer_caches( return num_kv_heads = int(k_view.shape[1]) head_dim = int(k_view.shape[3]) + # The dispatch clips k/v/idx_k to the step's live tokens before this + # runs, so every supplied row owns a real slot: num_tokens is the live + # count write_kv_slots requires. write_kv_slots( k_view, out_cache_loc, k.reshape(num_tokens, num_kv_heads, head_dim), + num_tokens, layout="HND", ) write_kv_slots( v_view, out_cache_loc, v.reshape(num_tokens, num_kv_heads, head_dim), + num_tokens, layout="HND", ) if idx_k is not None: @@ -1366,6 +1386,7 @@ def write_layer_caches( idx_cache, out_cache_loc, idx_k.reshape(num_tokens, 1, int(idx_cache.shape[-1])), + num_tokens, layout="HND", ) diff --git a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/triton_backend.py b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/triton_backend.py index 844fd979469c..fe53b4de0fa0 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/triton_backend.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/triton_backend.py @@ -158,8 +158,11 @@ def _write_main_kv_slots_to_pool( ``out_cache_loc`` is the 1-D ``[num_new_tokens]`` int tensor of flat slot ids to update. ``pool[:, kv_index]`` is a storage-sharing view, so the shared :func:`common.write_kv_slots` propagates the write to the pool. + + Every row owns a slot, so the row count is the live count: the slot mapping + emits one real slot per new token and no sentinel. """ - write_kv_slots(pool[:, kv_index], out_cache_loc, values) + write_kv_slots(pool[:, kv_index], out_cache_loc, values, int(values.shape[0])) def _write_main_kv_slots( @@ -173,7 +176,7 @@ def _write_main_kv_slots( flat-slot layout used by focused unit tests and the 4-D paged view of ``kv_pool[:, 0]`` / ``kv_pool[:, 1]``. """ - write_kv_slots(cache, out_cache_loc, values) + write_kv_slots(cache, out_cache_loc, values, int(values.shape[0])) def _scatter_topk_to_block_mask( diff --git a/tensorrt_llm/_torch/models/modeling_minimaxm3.py b/tensorrt_llm/_torch/models/modeling_minimaxm3.py index 02a314cb3117..c422c9d6313c 100644 --- a/tensorrt_llm/_torch/models/modeling_minimaxm3.py +++ b/tensorrt_llm/_torch/models/modeling_minimaxm3.py @@ -952,6 +952,40 @@ def _minimax_m3_fused_sparse_qkv_producer_fake( ) +def _dispatch_attention_over_live_tokens( + attn_layer: "MiniMaxM3Attention", + q: torch.Tensor, + k: Optional[torch.Tensor], + v: Optional[torch.Tensor], + idx_q: Optional[torch.Tensor], + idx_k: Optional[torch.Tensor], + attn_metadata: AttentionMetadata, + output: torch.Tensor, +) -> None: + """Run the attention core over the step's live tokens alone. + + A piecewise CUDA graph pads token-shaped inputs up to its capture bucket + without adding requests to go with them (see _get_padding_params in + model_engine), so q can outrun the rows the batch has. The kernels below + read a request out of a token index, so the pad comes off here, once for + both dispatch paths. + + The pad rows of output are left as the buffer supplied them. Nothing reads + them for its own result, and clearing them would put a device launch behind + a host-side count on every layer of every step. + """ + num_tokens = int(attn_metadata.num_tokens) + attn_layer._dispatch_attention_backend( + q[:num_tokens], + k[:num_tokens] if k is not None else None, + v[:num_tokens] if v is not None else None, + idx_q[:num_tokens] if idx_q is not None else None, + idx_k[:num_tokens] if idx_k is not None else None, + attn_metadata, + output[:num_tokens], + ) + + @torch.library.custom_op("trtllm::minimax_m3_attn_custom_op_inplace", mutates_args=("output",)) def minimax_m3_attn_custom_op_inplace( q: Optional[torch.Tensor], @@ -989,15 +1023,9 @@ def minimax_m3_attn_custom_op_inplace( k = v = idx_k = None if q is None: raise RuntimeError(f"MiniMax-M3 attention layer {layer_idx} received no query tensor.") - attn_layer._dispatch_attention_backend( - q[:num_tokens], - k[:num_tokens] if k is not None else None, - v[:num_tokens] if v is not None else None, - idx_q[:num_tokens] if idx_q is not None else None, - idx_k[:num_tokens] if idx_k is not None else None, - attn_metadata, - output[:num_tokens], - ) + # The live token count is a host value, so the compiled graph above must + # not see it: it would guard on it and recapture per count. + _dispatch_attention_over_live_tokens(attn_layer, q, k, v, idx_q, idx_k, attn_metadata, output) maybe_bcg_minimax_m3_attn_custom_op_inplace = eager_on_graph(minimax_m3_attn_custom_op_inplace) @@ -1959,7 +1987,9 @@ def _forward_attention_core( output, ) else: - self._dispatch_attention_backend(q, k, v, idx_q, idx_k, attn_metadata, output) + # A step that runs here rather than through compile is padded all + # the same, since the bucket is agreed across ranks. + _dispatch_attention_over_live_tokens(self, q, k, v, idx_q, idx_k, attn_metadata, output) return output def _dispatch_attention_backend( diff --git a/tests/unittest/_torch/attention/sparse/msa/test_msa_backend.py b/tests/unittest/_torch/attention/sparse/msa/test_msa_backend.py index 4fe0aef5e53b..a006e6e87fcc 100644 --- a/tests/unittest/_torch/attention/sparse/msa/test_msa_backend.py +++ b/tests/unittest/_torch/attention/sparse/msa/test_msa_backend.py @@ -76,6 +76,9 @@ def test_msa_metadata_clears_padded_cache_slot_tail( metadata._msa_qo_lens_cpu = torch.tensor([count], dtype=torch.int32) metadata._msa_kv_lens_cpu = metadata._msa_qo_lens_cpu.clone() metadata._msa_qo_offset_cpu = torch.zeros(1, dtype=torch.int32) + # What the base seq_lens setter would derive; the cache writers read the + # live count off it. + metadata._num_tokens = count metadata._build_msa_fields() assert metadata.msa_out_cache_loc.tolist() == list(range(12, 12 + count)) + [-1] * ( 4 - count @@ -685,6 +688,8 @@ def get_index_k_buffer(self, layer_idx): metadata.kv_cache_manager = manager metadata.msa_out_cache_loc = torch.tensor([2, page_size + 5], dtype=torch.int32) values = torch.arange(2 * head_dim, dtype=torch.float32).reshape(2, 1, head_dim) + metadata._msa_fields_ready = True + metadata._num_tokens = 2 returned = metadata.msa_idx_k_cache(3) metadata.msa_write_idx_k(3, values) @@ -809,6 +814,7 @@ def msa_write_idx_k(self, layer_idx: int, idx_k: torch.Tensor) -> None: self.cache, self.msa_out_cache_loc, idx_k, + int(idx_k.shape[0]), layout="HND", ) @@ -1226,6 +1232,97 @@ def test_the_decode_span_of_a_mixed_step_is_its_generation_suffix(): assert metadata.msa_max_kv_len == 40 +def test_the_triton_sparse_decode_rejects_a_q_that_outruns_the_batch(): + """The guard has to be reached, not merely available. + + Eleven speculative decode requests of 4 query tokens each, padded by a + piecewise CUDA graph out to 512: the kernel would read 512 page table rows + out of a batch that has 11. It reads its shapes before touching a device, + so the refusal happens at the call and needs no GPU. + """ + batch, query_len, total_q, num_heads, head_dim = 11, 4, 512, 4, 128 + q = torch.empty(total_q, num_heads, head_dim) + paged = torch.empty(1, 1, 128, head_dim) + with pytest.raises( + ValueError, match=r"Triton sparse decode: total_q \(512\) must be batch \(11\)" + ): + minimax_m3_sparse_attn_decode( + q, + paged, + paged, + torch.zeros(1, total_q, 64, dtype=torch.int64), + torch.zeros(batch, 4, dtype=torch.int32), + torch.zeros(batch, dtype=torch.int32), + sm_scale=head_dim**-0.5, + output=torch.empty_like(q), + decode_query_len=query_len, + ) + + +def test_the_trtllm_gen_dense_decode_rejects_a_q_that_outruns_the_batch(): + """The same padded q against the dense kernel, which had no such guard. + + It is handed no cache manager: its multi-CTA counters are sized off the + batch, so it has to decline before consulting one. + """ + pytest.importorskip("flashinfer") + from tensorrt_llm._torch.attention.backends.sparse.minimax_m3.kernels.trtllm_gen_dense_decode import ( + minimax_m3_trtllm_gen_dense_decode, + ) + + batch, query_len, total_q, num_heads, head_dim = 11, 4, 512, 4, 128 + q = torch.empty(total_q, num_heads, head_dim) + with pytest.raises( + ValueError, match=r"trtllm-gen dense decode: total_q \(512\) must be batch \(11\)" + ): + minimax_m3_trtllm_gen_dense_decode( + q, + None, + 0, + torch.zeros(batch, 4, dtype=torch.int32), + torch.zeros(batch, dtype=torch.int32), + sm_scale=head_dim**-0.5, + output=torch.empty_like(q), + decode_query_len=query_len, + max_seq_len=1024, + max_num_requests=batch, + ) + + +def test_the_eager_writer_drops_the_sentinel_tail_it_is_handed(): + """The padded rows of a step must reach no page at all, the wrap target of a + surviving -1 included.""" + num_pages, page_size, head_dim = 4, 8, 16 + cache = torch.zeros(num_pages, 1, page_size, head_dim, dtype=torch.bfloat16) + # Two live tokens, then the -1 tail a capture bucket pads the step out to. + out_cache_loc = torch.tensor([2, page_size + 5, -1, -1], dtype=torch.int32) + values = torch.arange(4 * head_dim, dtype=torch.float32).reshape(4, 1, head_dim) + + write_kv_slots(cache, out_cache_loc, values, 2, layout="HND") + + torch.testing.assert_close(cache[0, 0, 2], values[0, 0].to(torch.bfloat16)) + torch.testing.assert_close(cache[1, 0, 5], values[1, 0].to(torch.bfloat16)) + # The page a wrapped -1 would have hit. + assert not cache[num_pages - 1].any() + + +def test_the_eager_writer_refuses_a_live_count_it_has_no_rows_for(): + """A count past the rows supplied is a caller bug, not a short write.""" + cache = torch.zeros(4, 1, 8, 16, dtype=torch.bfloat16) + out_cache_loc = torch.tensor([2, 5], dtype=torch.int32) + values = torch.zeros(2, 1, 16) + + with pytest.raises(ValueError, match=r"num_live_tokens=3 exceeds the rows supplied"): + write_kv_slots(cache, out_cache_loc, values, 3, layout="HND") + + with pytest.raises(ValueError, match="must be non-negative"): + write_kv_slots(cache, out_cache_loc, values, -1, layout="HND") + + # A step that scheduled nothing writes nothing rather than erroring. + write_kv_slots(cache, out_cache_loc, values, 0, layout="HND") + assert not cache.any() + + def test_a_pure_prefill_step_has_no_decode_span(): """A step with no generation row has nothing for the decode kernels, and fmha_sm100 keeps every plan and the page table they read.""" @@ -1864,16 +1961,21 @@ def test_fused_scatter_matches_reference(src_dtype, cache_dtype, with_idx): ref_pool = pool.clone() ref_idx_pool = idx_pool.clone() + # The sentinel sits at the head here rather than in a tail, so the + # reference masks it out and states the surviving row count. + num_live_slots = int(valid.sum()) write_kv_slots( ref_pool[1:-1, 0], slots[valid], k.reshape(num_tokens, num_kv_heads, head_dim)[valid], + num_live_slots, layout="HND", ) write_kv_slots( ref_pool[1:-1, 1], slots[valid], v.reshape(num_tokens, num_kv_heads, head_dim)[valid], + num_live_slots, layout="HND", ) if with_idx: @@ -1881,6 +1983,7 @@ def test_fused_scatter_matches_reference(src_dtype, cache_dtype, with_idx): ref_idx_pool[1:-1, 0], slots[valid], idx_k.reshape(num_tokens, 1, head_dim)[valid], + num_live_slots, layout="HND", ) diff --git a/tests/unittest/_torch/models/test_minimax_m3.py b/tests/unittest/_torch/models/test_minimax_m3.py index 91c94ff0d800..0638b5ba3685 100644 --- a/tests/unittest/_torch/models/test_minimax_m3.py +++ b/tests/unittest/_torch/models/test_minimax_m3.py @@ -70,6 +70,7 @@ MiniMaxM3MoE, MiniMaxM3QKVIndexerLinear, _build_swiglu_oai_dense_mlp, + _dispatch_attention_over_live_tokens, _load_qkv_index_proj_weights, _minimax_m3_swiglu_oai, _moe_routed_output_is_global, @@ -598,6 +599,7 @@ def _dispatch_attention_backend(self, q, k, v, idx_q, idx_k, attn_metadata, outp assert layer.producer_shapes == ((2, 5), (1, 2), 2) torch.testing.assert_close(output[:2], packed[:2, :3]) + # The dispatch clips to the live tokens and leaves the pad rows as they came. torch.testing.assert_close(output[2:], torch.full((2, 3), -1.0)) @@ -1141,6 +1143,64 @@ def _has_cuda() -> bool: # --------------------------------------------------------------------------- +def test_attention_dispatch_clips_the_piecewise_token_pad(): + """Only the live tokens reach the attention core, and the output's pad rows + are left as the buffer supplied them.""" + seen = {} + + def capture(q, k, v, idx_q, idx_k, attn_metadata, output): + del attn_metadata + seen["rows"] = [None if t is None else int(t.shape[0]) for t in (q, k, v, idx_q, idx_k)] + seen["out_rows"] = int(output.shape[0]) + output.fill_(7.0) + + attn_layer = SimpleNamespace(_dispatch_attention_backend=capture) + # Eleven speculative decode requests of 4 query tokens, padded to 64. + padded, live, hidden = 64, 44, 8 + q = torch.ones((padded, hidden)) + output = torch.full((padded, hidden), float("nan")) + + _dispatch_attention_over_live_tokens( + attn_layer, + q, + q, + q, + None, + None, + SimpleNamespace(num_tokens=live), + output, + ) + + assert seen["rows"] == [live, live, live, None, None] + assert seen["out_rows"] == live + assert torch.equal(output[:live], torch.full((live, hidden), 7.0)) + assert output[live:].isnan().all() + + +def test_attention_dispatch_leaves_an_unpadded_step_alone(): + """No pad, so nothing to clip.""" + seen = {} + + def capture(q, k, v, idx_q, idx_k, attn_metadata, output): + del q, k, v, idx_q, idx_k, attn_metadata + seen["out"] = output + + output = torch.full((5, 8), float("nan")) + _dispatch_attention_over_live_tokens( + SimpleNamespace(_dispatch_attention_backend=capture), + torch.ones((5, 8)), + None, + None, + None, + None, + SimpleNamespace(num_tokens=5), + output, + ) + + assert seen["out"].shape == (5, 8) + assert output.isnan().all() + + def test_is_minimax_m3_vl_config_detects_vl(): assert is_minimax_m3_vl_config(_make_vl_config()) is True diff --git a/tests/unittest/_torch/thop/parallel_hw_agnostic/test_minimax_m3_fp8_indexer.py b/tests/unittest/_torch/thop/parallel_hw_agnostic/test_minimax_m3_fp8_indexer.py index 042ccd858e6a..bbe8cad52171 100644 --- a/tests/unittest/_torch/thop/parallel_hw_agnostic/test_minimax_m3_fp8_indexer.py +++ b/tests/unittest/_torch/thop/parallel_hw_agnostic/test_minimax_m3_fp8_indexer.py @@ -78,6 +78,37 @@ def _strided_cache(num_pages: int, page_size: int = 128, stride_scale: int = 7) return backing[::stride_scale] +def _guarded_cache( + num_pages: int, page_size: int = 128, stride_scale: int = 7, guard_pages: int = 2 +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Allocate a strided cache flanked by zeroed guard regions. + + Returns the cache view plus the backing below and above it. A store from a + negative slot lands below the base and one from an out-of-range slot lands + above the last page, so both stay inspectable rather than depending on a + fault that an out-of-bounds FP8 store may never raise. + """ + total_pages = num_pages + 2 * guard_pages + backing = torch.zeros( + total_pages * stride_scale, + 1, + page_size, + 128, + dtype=torch.float8_e4m3fn, + device="cuda", + ) + first = guard_pages * stride_scale + last = (guard_pages + num_pages) * stride_scale + view = backing[first:last:stride_scale] + assert view.shape[0] == num_pages + return view, backing[:first], backing[last:] + + +def _assert_all_zero(*regions: torch.Tensor) -> None: + for region in regions: + assert torch.count_nonzero(region.reshape(-1).view(torch.uint8)).item() == 0 + + def _run( qk: torch.Tensor, cache: torch.Tensor, @@ -163,6 +194,55 @@ def test_minimax_m3_fp8_indexer_defensively_skips_invalid_direct_op_slots() -> N assert torch.count_nonzero(backing[2].view(torch.uint8)).item() == 0 +@pytest.mark.parametrize("num_live_tokens", [0, 4]) +@pytest.mark.parametrize("tail_slot", [-1, 4 * 128, 4 * 128 + 61, 5 * 128 + 127]) +def test_minimax_m3_fp8_indexer_ignores_a_padded_slot_tail( + tail_slot: int, num_live_tokens: int +) -> None: + """Rows past the live prefix must leave every cache byte untouched. + + A direct caller can hand the kernel a padded token height whose trailing + slots are the -1 sentinel or a stale out-of-range id, which the two guards + have to drop. The failure modes differ: -1 truncates to page 0 at offset + -1, one token below the cache base, while an out-of-range page scatters + above the pool. The tail slots start at the first page past the end and + stay inside the guard region, so a regression trips an assert rather than + faulting. + """ + torch.manual_seed(2468) + num_heads_q = 4 + page_size = 128 + num_pages = 4 + padded_tokens = 17 + + qk = torch.randn(padded_tokens, (num_heads_q + 1) * 128, dtype=torch.bfloat16, device="cuda") + q_weight = torch.randn(128, dtype=torch.bfloat16, device="cuda") + k_weight = torch.randn(128, dtype=torch.bfloat16, device="cuda") + position_ids = torch.arange(padded_tokens, dtype=torch.int32, device="cuda") + 1024 + + cache, below, above = _guarded_cache(num_pages, page_size) + slots = torch.full((padded_tokens,), tail_slot, dtype=torch.int32, device="cuda") + # One page per live row, so a stray tail store cannot be mistaken for one. + pages = torch.arange(num_live_tokens, dtype=torch.int32, device="cuda") + within = (pages * 37) % page_size + slots[:num_live_tokens] = pages * page_size + within + + q_out = _run(qk, cache, slots, q_weight, k_weight, position_ids, num_heads_q) + q_ref, k_ref = _reference(qk, num_heads_q, q_weight, k_weight, position_ids) + + # Only the cache store is slot-gated; index-Q is produced for every row. + _assert_fp8_close(q_out, q_ref) + _assert_all_zero(below, above) + if num_live_tokens: + _assert_fp8_close(cache[pages.long(), 0, within.long()], k_ref[:num_live_tokens]) + + # A tail store landing on a valid page at the wrong offset would clear both + # guard regions, so require every unwritten slot to stay zero as well. + written = torch.zeros(num_pages, page_size, dtype=torch.bool, device="cuda") + written[pages.long(), within.long()] = True + _assert_all_zero(cache[:, 0][~written]) + + def test_minimax_m3_fp8_indexer_cuda_graph_replay_updates_outputs() -> None: torch.manual_seed(5678) num_tokens = 16