diff --git a/tensorrt_llm/_torch/attention_backend/sparse/dsa.py b/tensorrt_llm/_torch/attention_backend/sparse/dsa.py index 4dbaccc687da..bca0b367cc47 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/dsa.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/dsa.py @@ -122,26 +122,59 @@ def warmup_heuristic_topk_decode(top_k: int = 2048, _DG_SCHEDULE_BLOCK_KV = 64 -def _pick_fp4_dsl_expand(next_n: int) -> Tuple[int, int]: - """Pick (expand_factor, effective_next_n) for the DSL FP4 paged kernel. - - The CuTe DSL FP4 paged MQA logits kernel only supports - ``effective_next_n ∈ {1, 2, 3}``. For larger ``next_n`` we reshape - ``[B, next_n, ...]`` to ``[B * expand_factor, effective_next_n, ...]`` - caller-side. To minimize HBM bandwidth we pick the smallest - ``expand_factor`` (== largest ``effective_next_n``) that divides - ``next_n`` cleanly. - - Examples: - next_n=4 -> (2, 2): two atoms of next_n=2 (DG's kNextNAtom=2 style) - next_n=5 -> (5, 1): only divisor; 5x HBM - next_n=6 -> (2, 3): two atoms of next_n=3 (prefer eff=3 over eff=2) - next_n=9 -> (3, 3): three atoms of next_n=3 +def _pick_dsl_expand( + next_n: int, + num_sms: int, + batch_size: int = 0, + max_ctx: int = 0, + kernel_atoms: Tuple[int, ...] = (1, 2, 3)) -> Tuple[int, int]: + """Pick (expand_factor, effective_next_n) for the DSL paged kernel + using a wave-aware strategy. Used by both FP4 and FP8 DSL paths. + + The DSL kernel natively supports ``effective_next_n ∈ kernel_atoms`` + (FP4: ``(1, 2, 3)``; FP8: ``(1, 2, 3, 4)``). For ``next_n`` not natively + supported or when SM utilization can be improved, reshape + ``[B, next_n, ...]`` -> ``[B * expand_factor, effective_next_n, ...]`` + caller-side. + + Strategy: enumerate ``(expand_factor, effective_next_n)`` pairs with + ``expand_factor * effective_next_n == next_n`` and ``effective_next_n + in kernel_atoms``. Score each by ``(waves, -expand_factor)`` where + ``waves = ceil(B * expand_factor * ceil(max_ctx/256) / num_sms)``. + Pick min waves; on tie, prefer LARGER expand_factor (more SMs busy per + wave; pays HBM cost of expand_factor x KV re-reads). + + When ``batch_size == 0`` or ``max_ctx == 0`` (workload unknown), fall + back to the legacy HBM-minimizing heuristic: largest effective_next_n + that divides next_n cleanly (still constrained to ``kernel_atoms``). + + Examples (wave-aware, num_sms=148 [B200], SPLIT_KV=256 tokens): + FP4, next_n=4, B=1, ctx=4096 -> (4, 1): ntask=64<148, 1 wave, max factor + FP4, next_n=4, B=32, ctx=4096 -> (2, 2): ntask=1024>148, multi-wave, min factor + FP4, next_n=2, B=1, ctx=4096 -> (2, 1): wave-tie, larger factor + FP8, next_n=4, B=1, ctx=4096 -> (4, 1): kernel_atoms incl. 4 doesn't change small-B pick """ - for eff in (3, 2): + # Legacy fallback when workload is unknown. + if batch_size <= 0 or max_ctx <= 0: + for eff in sorted(kernel_atoms, reverse=True): + if next_n % eff == 0: + return next_n // eff, eff + return next_n, 1 + + SPLIT_KV_TOKENS = 256 + cands = [] + for eff in kernel_atoms: if next_n % eff == 0: - return next_n // eff, eff - return next_n, 1 + factor = next_n // eff + ntask = batch_size * factor * ( + (max_ctx + SPLIT_KV_TOKENS - 1) // SPLIT_KV_TOKENS) + waves = (ntask + num_sms - 1) // num_sms + cands.append((waves, factor, eff)) + if not cands: + return next_n, 1 + cands.sort(key=lambda x: (x[0], -x[1])) # min waves, max factor + _, factor, eff = cands[0] + return factor, eff def _compute_slot_mappings( @@ -425,13 +458,19 @@ class DSAtrtllmAttentionMetadata(TrtllmAttentionMetadata): skip_indexer_for_gen_reqs: bool = False # Whether to use the expanded buffers for MTP support use_expanded_buffers_for_mtp: bool = False - # Whether to reshape the DSL FP4 paged MQA logits Q tensor into - # supported next_n ∈ {1, 2, 3} via caller-side atom-split (see - # `_pick_fp4_dsl_expand`). Reuses `kv_lens_expanded_cuda` / - # `block_table_expanded` / `scheduler_metadata_buffer_expanded`; runtime - # mutually exclusive with `use_expanded_buffers_for_mtp` (the latter - # requires `not _use_dsl`). - expand_for_dsl_fp4: bool = False + # Whether to reshape the DSL paged MQA logits Q tensor into a kernel- + # supported `effective_next_n` via caller-side atom-split (FP4: {1,2,3}; + # FP8: {1,2,3,4}; see `_pick_dsl_expand`). Reuses + # `kv_lens_expanded_cuda` / `block_table_expanded` / + # `scheduler_metadata_buffer_expanded`; runtime mutually exclusive with + # `use_expanded_buffers_for_mtp` (the latter requires `not _use_dsl`). + expand_for_dsl: bool = False + # Cached (expand_factor, atom) decision from the wave-aware picker. Set at + # `prepare()` time and read by forward call sites — avoids re-running the + # picker per call and guarantees prepare/forward use the SAME decision + # (otherwise the populated buffers would mismatch the kernel reshape). + dsl_expand_factor: int = 1 + dsl_atom: int = 1 def __init__(self, *args, **kwargs): """Initialize DSA metadata with SM count and indexer chunk size.""" @@ -1076,39 +1115,68 @@ def prepare(self): self.block_table_expanded.clamp_(min=0) # CuTe DSL FP4 paged MQA logits kernel natively supports - # next_n ∈ {1, 2, 3} only. For next_n ≥ 4 we caller-side reshape - # [B, next_n, ...] -> [B * expand_factor, eff_next_n, ...] (see - # `_pick_fp4_dsl_expand`). The expanded buffers (`kv_lens_expanded_cuda` - # / `block_table_expanded` / `scheduler_metadata_buffer_expanded`, - # all sized for the worst-case `1+max_draft_tokens` factor) are - # reused: under `_use_dsl=True` the existing FP8 DG expand path - # never writes to them (it's gated on `not _use_dsl`), so there's - # no conflict. - self.expand_for_dsl_fp4 = (_use_dsl - and self.kv_cache_manager is not None - and self.kv_cache_manager.use_fp4 - and self.max_draft_tokens >= 3) - if self.expand_for_dsl_fp4 and self.num_generations > 0: + # next_n ∈ {1, 2, 3} only. For next_n ≥ 4 atom-split is mandatory. + # For next_n ∈ {2, 3} atom-split is also beneficial when the wave-aware + # picker decides more SM utilization outweighs the (expand_factor)x + # HBM cost (e.g., low batch with idle SMs). The expanded buffers + # (`kv_lens_expanded_cuda` / `block_table_expanded` / + # `scheduler_metadata_buffer_expanded`, all sized for the worst-case + # `1+max_draft_tokens` factor) are reused: under `_use_dsl=True` the + # existing FP8 DG expand path never writes to them (gated on + # `not _use_dsl`), so there's no conflict. + # Trigger relaxed to `max_draft_tokens >= 1` (i.e., next_n >= 2) so the + # picker can choose to expand when waves vs HBM trade-off favors it. + # Trigger atom-split for both FP4 and FP8 DSL paths. FP4 kernel + # supports atom ∈ {1, 2, 3}; FP8 supports {1, 2, 3, 4}. Picker is + # given the appropriate kernel_atoms set so it only enumerates + # decompositions the kernel can handle. + self.expand_for_dsl = (_use_dsl and self.kv_cache_manager is not None + and self.max_draft_tokens >= 1) + if self.expand_for_dsl and self.num_generations > 0: next_n = 1 + self.max_draft_tokens - expand_factor, _ = _pick_fp4_dsl_expand(next_n) - num_tokens = self.num_generations * expand_factor + kernel_atoms = (1, 2, + 3) if self.kv_cache_manager.use_fp4 else (1, 2, 3, + 4) + # Wave-aware picker. max_ctx ≈ longest gen kv_len (decode iter + # upper-bound observed at this prepare). num_sms is hardware. gen_kv_lens = kv_lens[self.num_contexts:self.num_seqs] - gen_kv_lens_expanded = gen_kv_lens.repeat_interleave(expand_factor) - self.kv_lens_expanded_host[:num_tokens].copy_(gen_kv_lens_expanded) - self.kv_lens_expanded_cuda[:num_tokens].copy_( - self.kv_lens_expanded_host[:num_tokens], non_blocking=True) - if self.kv_cache_manager is not None: - max_len = self.host_indexer_k_cache_block_offsets.shape[1] - gen_block_tensor = self.host_indexer_k_cache_block_offsets[ - self.num_contexts:self.num_seqs, :max_len] - expanded_blocks = gen_block_tensor.repeat_interleave( - expand_factor, dim=0) - self.host_block_table_expanded[:num_tokens, :max_len].copy_( - expanded_blocks, non_blocking=True) - self.block_table_expanded[:num_tokens].copy_( - self.host_block_table_expanded[:num_tokens], - non_blocking=True) - self.block_table_expanded.clamp_(min=0) + max_ctx = int( + gen_kv_lens.max().item()) if gen_kv_lens.numel() else 0 + expand_factor, atom = _pick_dsl_expand( + next_n, + batch_size=self.num_generations, + max_ctx=max_ctx, + num_sms=self.num_sms, + kernel_atoms=kernel_atoms, + ) + self.dsl_expand_factor = expand_factor + self.dsl_atom = atom + # Only populate when picker chose to actually split (factor > 1); + # factor=1 means kernel-native, no expansion needed. + if expand_factor > 1: + num_tokens = self.num_generations * expand_factor + gen_kv_lens_expanded = gen_kv_lens.repeat_interleave( + expand_factor) + self.kv_lens_expanded_host[:num_tokens].copy_( + gen_kv_lens_expanded) + self.kv_lens_expanded_cuda[:num_tokens].copy_( + self.kv_lens_expanded_host[:num_tokens], non_blocking=True) + if self.kv_cache_manager is not None: + max_len = self.host_indexer_k_cache_block_offsets.shape[1] + gen_block_tensor = self.host_indexer_k_cache_block_offsets[ + self.num_contexts:self.num_seqs, :max_len] + expanded_blocks = gen_block_tensor.repeat_interleave( + expand_factor, dim=0) + self.host_block_table_expanded[:num_tokens, :max_len].copy_( + expanded_blocks, non_blocking=True) + self.block_table_expanded[:num_tokens].copy_( + self.host_block_table_expanded[:num_tokens], + non_blocking=True) + self.block_table_expanded.clamp_(min=0) + else: + # Reset cache; forward path uses kernel-native next_n. + self.dsl_expand_factor = 1 + self.dsl_atom = 1 + self.max_draft_tokens # Prepare metadata for indexer Indexer.prepare(metadata=self) @@ -1227,6 +1295,26 @@ def on_update_kv_lens(self): kv_lens_expanded_2d, _DG_SCHEDULE_BLOCK_KV, self.num_sms) self.scheduler_metadata_buffer_expanded.copy_( scheduler_metadata_buffer_expanded, non_blocking=True) + # DSL atom-split path: mirror the prepare()-time build so that + # overlap-scheduler / spec-dec runtime corrections to kv_lens_cuda + # propagate into kv_lens_expanded_cuda and the matching schedule. + # Reuse the cached (dsl_expand_factor, dsl_atom) — re-running the + # picker here would let the split decision drift between prepare + # and forward, breaking CUDA graph capture. + if self.expand_for_dsl and self.dsl_expand_factor > 1: + expand_factor = self.dsl_expand_factor + num_tokens = self.num_generations * expand_factor + gen_kv_lens_expanded = gen_kv_lens.repeat_interleave( + expand_factor) + self.kv_lens_expanded_cuda[:num_tokens].copy_( + gen_kv_lens_expanded) + kv_lens_expanded_2d = self.kv_lens_expanded_cuda[: + num_tokens].view( + -1, 1) + scheduler_metadata_buffer_expanded = get_paged_mqa_logits_metadata( + kv_lens_expanded_2d, _DG_SCHEDULE_BLOCK_KV, self.num_sms) + self.scheduler_metadata_buffer_expanded.copy_( + scheduler_metadata_buffer_expanded, non_blocking=True) self.prepare_dense_topk_indices(self.kv_lens_cuda, device=True) def update_for_spec_dec(self): @@ -1661,16 +1749,15 @@ def prepare(metadata: DSAtrtllmAttentionMetadata): metadata.scheduler_metadata_buffer_expanded.copy_( scheduler_metadata_buffer_expanded, non_blocking=True) - # DSL FP4 atom-split schedule. The DSL kernel reshapes - # [B, next_n, ...] -> [B * factor, eff_next_n, ...] caller-side - # for next_n > 3 (kernel only supports {1, 2, 3} natively); the - # matching schedule is built from the doubled-batch shape so - # `num_next_n_atoms` keeps encoding the effective next_n. Runtime - # mutually exclusive with the `else` branch above (the latter - # requires `use_expanded_buffers_for_mtp` which is False under DSL). - if metadata.expand_for_dsl_fp4 and metadata.num_generations > 0: - expand_factor, _ = _pick_fp4_dsl_expand( - 1 + metadata.max_draft_tokens) + # DSL atom-split schedule. Picker decision was cached on + # `metadata.dsl_{expand_factor, atom}` at metadata prepare + # time; only build the expanded schedule when picker chose to + # split (factor > 1). Runtime mutually exclusive with the `else` + # branch above (latter requires `use_expanded_buffers_for_mtp` + # which is False under DSL). + if metadata.expand_for_dsl and metadata.num_generations > 0 \ + and metadata.dsl_expand_factor > 1: + expand_factor = metadata.dsl_expand_factor num_tokens = metadata.num_generations * expand_factor kv_lens_expanded_2d = metadata.kv_lens_expanded_cuda[: num_tokens].view( @@ -2058,16 +2145,17 @@ def sparse_attn_indexer( dsl_block_table = block_table dsl_schedule_meta = metadata.scheduler_metadata_buffer - # DSL FP4 kernel natively supports next_n ∈ {1, 2, 3}; for - # larger next_n, reshape [B, next_n, ...] -> [B*factor, - # eff_next_n, ...] (HBM-optimal — see - # `_pick_fp4_dsl_expand`). Row layout of logits is - # preserved by the reshape so no reassembly is needed at - # the output side. The expanded metadata is populated - # in DSAtrtllmAttentionMetadata.prepare / - # Indexer.prepare under `expand_for_dsl_fp4`. - if next_n > 3: - factor, eff_next_n = _pick_fp4_dsl_expand(next_n) + # DSL FP4 kernel natively supports next_n ∈ {1, 2, 3}. + # The wave-aware picker in `_pick_dsl_expand` is run + # once per metadata prepare and the result cached on + # `metadata.dsl_{expand_factor, atom}`. Trigger expand + # whenever the picker decided to split (factor > 1), + # regardless of next_n — this lets next_n ∈ {2, 3} also + # benefit from atom-split when low-batch leaves SMs idle, + # in addition to the mandatory next_n=4 case. + if metadata.dsl_expand_factor > 1: + factor = metadata.dsl_expand_factor + eff_next_n = metadata.dsl_atom exp_B = num_generations * factor dsl_q = dsl_q.reshape(exp_B, eff_next_n, self.n_heads, self.head_dim // 2) @@ -2084,10 +2172,29 @@ def sparse_attn_indexer( dsl_context_lens, dsl_block_table, dsl_schedule_meta, max_seq_len) else: + # FP8 DSL kernel natively supports next_n ∈ {1, 2, 3, 4}. + # Apply wave-aware atom-split when the picker decided to + # split (factor > 1) — typically benefits small-batch / + # low-ntask configs by raising SM utilization at the cost + # of factor× KV HBM re-reads. Picker decision was cached + # on metadata.{dsl_expand_factor, dsl_atom} during prepare. + dsl_q = q_decode + fp8_ctx_lens = dsl_context_lens + fp8_block_table = block_table + fp8_schedule_meta = metadata.scheduler_metadata_buffer + if metadata.dsl_expand_factor > 1: + factor = metadata.dsl_expand_factor + atom = metadata.dsl_atom + exp_B = num_generations * factor + dsl_q = q_decode.reshape(exp_B, atom, self.n_heads, + self.head_dim) + fp8_ctx_lens = metadata.kv_lens_expanded_cuda[:exp_B] + fp8_block_table = metadata.block_table_expanded[:exp_B] + fp8_schedule_meta = ( + metadata.scheduler_metadata_buffer_expanded) logits_decode = torch.ops.trtllm.cute_dsl_fp8_paged_mqa_logits( - q_decode, k_cache, weights_decode, dsl_context_lens, - block_table, metadata.scheduler_metadata_buffer, - max_seq_len) + dsl_q, k_cache, weights_decode, fp8_ctx_lens, + fp8_block_table, fp8_schedule_meta, max_seq_len) else: decode_q_scale = q_scale[num_ctx_tokens:num_ctx_tokens + num_gen_tokens, diff --git a/tensorrt_llm/_torch/custom_ops/cute_dsl_custom_ops.py b/tensorrt_llm/_torch/custom_ops/cute_dsl_custom_ops.py index b2cd3b137640..030fd6f57845 100644 --- a/tensorrt_llm/_torch/custom_ops/cute_dsl_custom_ops.py +++ b/tensorrt_llm/_torch/custom_ops/cute_dsl_custom_ops.py @@ -5986,9 +5986,15 @@ def _compile(cls, compute_block_kv, phys_block_kv, num_heads, head_dim, max_blocks_per_seq = cute.sym_int() num_ctas = cute.sym_int() - kv_fake = cute.runtime.make_fake_compact_tensor( + # KV may come from the indexer K-cache pool view, which is + # strided in dim 0 (pool layout interleaves layers: + # [num_blocks, num_layers, kvFactor, blockSize]). Declare outer + # stride as sym so the actual per-block stride is read at + # runtime; innermost stride is fixed to 1 (byte-contig within a + # logical block view). + kv_fake = cute.runtime.make_fake_tensor( cutlass.Uint8, (sym_num_phys_blocks, block_bytes), - stride_order=(1, 0)) + stride=(cute.sym_int64(), 1)) q_fake = cute.runtime.make_fake_compact_tensor(cutlass.Uint8, (N, head_dim, sym_B), @@ -6770,12 +6776,21 @@ class CuteDSLFP4PagedMQALogitsRunner: kernel_cache = dict() @classmethod - def _compile(cls, compute_block_kv, phys_block_kv, num_heads, head_dim, - next_n, num_sms, num_epi_subtiles, epi_dtype, - output_dtype): + def _compile(cls, + compute_block_kv, + phys_block_kv, + num_heads, + head_dim, + next_n, + num_sms, + num_epi_subtiles, + epi_dtype, + output_dtype, + remove_online_sf_transpose=False): """Compile kernel using fake tensors + TVM FFI.""" key = (compute_block_kv, phys_block_kv, num_heads, head_dim, next_n, - num_sms, num_epi_subtiles, epi_dtype, output_dtype) + num_sms, num_epi_subtiles, epi_dtype, output_dtype, + remove_online_sf_transpose) if key in cls.kernel_cache: return @@ -6791,9 +6806,15 @@ def _compile(cls, compute_block_kv, phys_block_kv, num_heads, head_dim, max_blocks_per_seq = cute.sym_int() num_ctas = cute.sym_int() - kv_fake = cute.runtime.make_fake_compact_tensor( + # KV may come from the indexer K-cache pool view, which is + # strided in dim 0 (pool layout interleaves layers: + # [num_blocks, num_layers, kvFactor, blockSize]). Declare outer + # stride as sym so the actual per-block stride is read at + # runtime; innermost stride is fixed to 1 (byte-contig within a + # logical block view). + kv_fake = cute.runtime.make_fake_tensor( cutlass.Uint8, (sym_num_phys_blocks, block_bytes), - stride_order=(1, 0)) + stride=(cute.sym_int64(), 1)) # Q is FP4 packed bytes: head_dim/2 bytes per row q_fake = cute.runtime.make_fake_compact_tensor( @@ -6845,6 +6866,7 @@ def _compile(cls, compute_block_kv, phys_block_kv, num_heads, head_dim, num_epi_subtiles=num_epi_subtiles, epi_dtype=to_cutlass[epi_dtype], output_dtype=to_cutlass[output_dtype], + remove_online_sf_transpose=remove_online_sf_transpose, ) compiled = cute.compile( @@ -6879,6 +6901,7 @@ def forward( num_epi_subtiles: int = 1, epi_dtype: torch.dtype = torch.float32, output_dtype: torch.dtype = torch.float32, + remove_online_sf_transpose: bool = False, ) -> torch.Tensor: """Execute FP4 paged MQA logits kernel. @@ -6946,10 +6969,20 @@ def forward( # Compile if needed (fake tensors, no real data required) key = (compute_block_kv, phys_block_kv, H, D, next_n, num_sms, - num_epi_subtiles, epi_dtype, output_dtype) + num_epi_subtiles, epi_dtype, output_dtype, + remove_online_sf_transpose) if key not in cls.kernel_cache: - cls._compile(compute_block_kv, phys_block_kv, H, D, next_n, - num_sms, num_epi_subtiles, epi_dtype, output_dtype) + cls._compile( + compute_block_kv, + phys_block_kv, + H, + D, + next_n, + num_sms, + num_epi_subtiles, + epi_dtype, + output_dtype, + remove_online_sf_transpose=remove_online_sf_transpose) compiled = cls.kernel_cache[key] # TVM FFI: pass raw tensors, no dlpack/stream needed @@ -6972,6 +7005,7 @@ def cute_dsl_fp4_paged_mqa_logits( num_epi_subtiles: int = 1, epi_dtype: torch.dtype = torch.float32, output_dtype: torch.dtype = torch.float32, + remove_online_sf_transpose: bool = False, ) -> torch.Tensor: if not is_sm_100f(): raise ValueError( @@ -7008,7 +7042,8 @@ def cute_dsl_fp4_paged_mqa_logits( max_context_len, num_epi_subtiles=num_epi_subtiles, epi_dtype=epi_dtype, - output_dtype=output_dtype) + output_dtype=output_dtype, + remove_online_sf_transpose=remove_online_sf_transpose) @torch.library.register_fake("trtllm::cute_dsl_fp4_paged_mqa_logits") def _( @@ -7023,6 +7058,7 @@ def _( num_epi_subtiles: int = 1, epi_dtype: torch.dtype = torch.float32, output_dtype: torch.dtype = torch.float32, + remove_online_sf_transpose: bool = False, ) -> torch.Tensor: B = q.shape[0] next_n = q.shape[1] diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/blackwell/paged_mqa_logits/fp4_paged_mqa_logits.py b/tensorrt_llm/_torch/cute_dsl_kernels/blackwell/paged_mqa_logits/fp4_paged_mqa_logits.py index f53dc069da1a..1e4d5abf1bfc 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/blackwell/paged_mqa_logits/fp4_paged_mqa_logits.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/blackwell/paged_mqa_logits/fp4_paged_mqa_logits.py @@ -364,6 +364,8 @@ def __init__( num_epi_subtiles: int = 1, epi_dtype=cutlass.Float32, output_dtype=cutlass.Float32, + remove_online_sf_transpose: bool = False, + use_batched_store: bool = True, ): # Static FP4 invariants — see plan Sanity checklist. assert num_heads == 64, "FP4 kernel hardcodes num_heads=64 for TMEM/SMEM budget" @@ -398,6 +400,16 @@ def __init__( self.num_sms = num_sms self.num_epi_subtiles = num_epi_subtiles self.epi_dtype = epi_dtype + # When True, skip the in-kernel SMEM warp_transpose for KV SF; assume + # the host has pre-arranged GMEM SF into UTCCP chunk layout. Only valid + # for phys_block_kv=128 (1 phys block = 1 UTCCP atom). Q SF transpose + # is NOT affected by this flag (deferred to a separate phase). + if remove_online_sf_transpose and phys_block_kv != 128: + remove_online_sf_transpose = False + self.remove_online_sf_transpose = remove_online_sf_transpose + # When True, defer per-t STG to register array and emit all STGs in + # one contiguous LSU phase after the for-t loop (epilogue micro-opt). + self.use_batched_store = use_batched_store # epi_bytes covers fp16 and bf16 (FP8 only handled fp16). self.epi_bytes = 2 if epi_dtype in (cutlass.Float16, cutlass.BFloat16) else 4 # sW stage stride padded to 128-byte SMEM alignment for TMA bulk copy. @@ -625,7 +637,6 @@ def __call__( # [SF phys_block_kv*4 bytes (= phys_block_kv int32)] phys_block_kv = self.phys_block_kv half_head_dim = self.head_dim // 2 # FP4 packed bytes per row - phys_block_bytes = phys_block_kv * (half_head_dim + 4) scale_offset_bytes = phys_block_kv * half_head_dim # to SF region of each phys block # Recast the fused buffer to FP4. Each uint8 byte becomes 2 FP4 elements, @@ -636,23 +647,31 @@ def __call__( # type inference and TMA descriptors are correct. b = cute.recast_tensor(b, Float4E2M1FN) + # Read the real per-block stride (bytes) from the input tensor. + # When KV is the indexer K-cache pool view, the pool is laid out as + # [num_blocks, num_layers, kvFactor, blockSize], so dim-0 stride = + # num_layers * kvFactor * phys_block_bytes (not phys_block_bytes). + # Using the input stride keeps both the contiguous test path and + # the strided prod path correct. + kv_block_stride_bytes = kv_fused.layout.stride[0] + # KV data view: [phys_block_kv, head_dim, num_phys_blocks] FP4 elements. # Innermost stride 1 = consecutive FP4 elem = packed pair share a byte. # Per-row stride = head_dim FP4 elem = head_dim/2 bytes. - # Per-block stride = phys_block_bytes * 2 FP4 elem (data + SF region). + # Per-block stride (FP4 elem) = kv_block_stride_bytes * 2 (uint8→FP4 doubles). kv_layout = cute.make_layout( (phys_block_kv, self.head_dim, num_phys_blocks), - stride=(self.head_dim, 1, phys_block_bytes * 2), + stride=(self.head_dim, 1, kv_block_stride_bytes * 2), ) a = cute.make_tensor(kv_fp4.iterator, kv_layout) # SF KV view: int32 (4 UE8M0 packed). Build a uint8 view at the SF # offset, then recast to int32. - # Layout in bytes: (phys_block_kv * 4, num_phys_blocks) stride (1, phys_block_bytes) - # After recast int32: (phys_block_kv, num_phys_blocks) stride (1, phys_block_bytes/4) + # Layout in bytes: (phys_block_kv * 4, num_phys_blocks) stride (1, kv_block_stride_bytes) + # After recast int32: (phys_block_kv, num_phys_blocks) stride (1, kv_block_stride_bytes/4) sf_kv_uint8_layout = cute.make_layout( (phys_block_kv * 4, num_phys_blocks), - stride=(1, phys_block_bytes), + stride=(1, kv_block_stride_bytes), ) sf_kv_uint8 = cute.make_tensor(kv_fused.iterator + scale_offset_bytes, sf_kv_uint8_layout) sf_kv = cute.recast_tensor(sf_kv_uint8, cutlass.Int32) @@ -1568,14 +1587,18 @@ def kernel( # Step 5.6: SF KV transpose + UTCCP. block_kv = 128 = 1 # UTCCP atom; loop is constexpr-1 but kept for clarity. - sf_kv_atoms = self.block_kv // 128 - for atom_idx in cutlass.range_constexpr(sf_kv_atoms): - atom_offset = atom_idx * 128 - stage_offset = kv_stage * sSF_KV_0.layout.stride[1] - utccp_required_smem_warp_transpose( - sSF_KV_0.iterator + stage_offset + atom_offset - ) - cute.arch.fence_view_async_shared() + # When remove_online_sf_transpose=True, the host has already + # pre-arranged GMEM SF into UTCCP chunk layout, so the + # in-kernel SMEM transpose (and its fence) can be skipped. + if cutlass.const_expr(not self.remove_online_sf_transpose): + sf_kv_atoms = self.block_kv // 128 + for atom_idx in cutlass.range_constexpr(sf_kv_atoms): + atom_offset = atom_idx * 128 + stage_offset = kv_stage * sSF_KV_0.layout.stride[1] + utccp_required_smem_warp_transpose( + sSF_KV_0.iterator + stage_offset + atom_offset + ) + cute.arch.fence_view_async_shared() # int32 SMEM → UE8M0 view for UTCCP atom + chunk layout. sSF_KV_0_ue8m0 = cute.recast_tensor(sSF_KV_0, Float8E8M0FNU) stage_off_kv0_ue8m0 = kv_stage * sSF_KV_0_ue8m0.layout.stride[1] @@ -1705,14 +1728,18 @@ def kernel( kv_stage_1 = kv_cons_state_umma_1.index # Step 5.6: SF KV (group 1) transpose + UTCCP. - sf_kv_atoms_1 = self.block_kv // 128 - for atom_idx in cutlass.range_constexpr(sf_kv_atoms_1): - atom_offset = atom_idx * 128 - stage_offset = kv_stage_1 * sSF_KV_1.layout.stride[1] - utccp_required_smem_warp_transpose( - sSF_KV_1.iterator + stage_offset + atom_offset - ) - cute.arch.fence_view_async_shared() + # When remove_online_sf_transpose=True, the host has already + # pre-arranged GMEM SF into UTCCP chunk layout, so the + # in-kernel SMEM transpose (and its fence) can be skipped. + if cutlass.const_expr(not self.remove_online_sf_transpose): + sf_kv_atoms_1 = self.block_kv // 128 + for atom_idx in cutlass.range_constexpr(sf_kv_atoms_1): + atom_offset = atom_idx * 128 + stage_offset = kv_stage_1 * sSF_KV_1.layout.stride[1] + utccp_required_smem_warp_transpose( + sSF_KV_1.iterator + stage_offset + atom_offset + ) + cute.arch.fence_view_async_shared() sSF_KV_1_ue8m0 = cute.recast_tensor(sSF_KV_1, Float8E8M0FNU) stage_off_kv1_ue8m0 = kv_stage_1 * sSF_KV_1_ue8m0.layout.stride[1] sSF_KV_1_chunk = cute.make_tensor( @@ -1816,6 +1843,13 @@ def kernel( MAX_NUM_W_IN_REG = 56 if next_n == 3 else 64 NUM_W_IN_REG = min(MAX_NUM_W_IN_REG, num_heads) w_cache = cute.make_fragment(NUM_W_IN_REG * next_n, self.epi_dtype) + # Batched STG: hold reduced result per t in register; the + # actual STG happens once after the for-t loop to land all + # STGs in one contiguous LSU phase. + if cutlass.const_expr(self.use_batched_store): + result_arr = cute.make_fragment(next_n, self.output_dtype) + else: + result_arr = None q_stage_local = cutlass.Int32(0) while has_work: @@ -2022,9 +2056,18 @@ def kernel( result_t = sum_lo + sum_hi else: result_t = s0x + s0y + s1x + s1y - out_row = q_idx * next_n + t # Step 5.7: drop * scale_val (FP4 SF baked into acc). - mLogits[(out_row, kv_pos)] = self.output_dtype(result_t) + if cutlass.const_expr(self.use_batched_store): + result_arr[t] = self.output_dtype(result_t) + else: + out_row = q_idx * next_n + t + mLogits[(out_row, kv_pos)] = self.output_dtype(result_t) + + if cutlass.const_expr(self.use_batched_store): + # Batched STG: all result_arr[t] → mLogits in one pass. + for t in cutlass.range_constexpr(next_n): + out_row = q_idx * next_n + t + mLogits[(out_row, kv_pos)] = result_arr[t] # Advance: inline fetch_next_task next_kv_idx = kv_idx + NUM_MATH_WG @@ -2063,6 +2106,13 @@ def kernel( MAX_NUM_W_IN_REG = 56 if next_n == 3 else 64 NUM_W_IN_REG = min(MAX_NUM_W_IN_REG, num_heads) w_cache = cute.make_fragment(NUM_W_IN_REG * next_n, self.epi_dtype) + # Batched STG: hold reduced result per t in register; the + # actual STG happens once after the for-t loop to land all + # STGs in one contiguous LSU phase. + if cutlass.const_expr(self.use_batched_store): + result_arr = cute.make_fragment(next_n, self.output_dtype) + else: + result_arr = None q_stage_local = cutlass.Int32(0) while has_work: @@ -2259,9 +2309,18 @@ def kernel( result_t = sum_lo + sum_hi else: result_t = s0x + s0y + s1x + s1y - out_row = q_idx * next_n + t # Step 5.7: drop * scale_val (FP4 SF baked into acc). - mLogits[(out_row, kv_pos)] = self.output_dtype(result_t) + if cutlass.const_expr(self.use_batched_store): + result_arr[t] = self.output_dtype(result_t) + else: + out_row = q_idx * next_n + t + mLogits[(out_row, kv_pos)] = self.output_dtype(result_t) + + if cutlass.const_expr(self.use_batched_store): + # Batched STG: all result_arr[t] → mLogits in one pass. + for t in cutlass.range_constexpr(next_n): + out_row = q_idx * next_n + t + mLogits[(out_row, kv_pos)] = result_arr[t] # Advance: inline fetch_next_task next_kv_idx = kv_idx + NUM_MATH_WG diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/blackwell/paged_mqa_logits/fp8_paged_mqa_logits.py b/tensorrt_llm/_torch/cute_dsl_kernels/blackwell/paged_mqa_logits/fp8_paged_mqa_logits.py index dc2282256f39..217abc474bb1 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/blackwell/paged_mqa_logits/fp8_paged_mqa_logits.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/blackwell/paged_mqa_logits/fp8_paged_mqa_logits.py @@ -372,7 +372,6 @@ def __call__( # Derive KV and Scale views from fused buffer using CuTE ops. # Fused layout per physical block: [KV data (phys_block_kv*head_dim)] [Scales (phys_block_kv*4)] phys_block_kv = self.phys_block_kv - phys_block_bytes = phys_block_kv * (self.head_dim + 4) scale_offset_elems = phys_block_kv * self.head_dim # in FP8 elements # Recast fused buffer to FP8 (same 1-byte elements, needed for MMA type inference) @@ -382,11 +381,20 @@ def __call__( # recast back to FP8 so MMA type inference and TMA descriptors are correct. b = cute.recast_tensor(b, cutlass.Float8E4M3FN) + # Read the real per-block stride (bytes = FP8 elements) from the input. + # When KV is the indexer K-cache pool view, the pool is laid out as + # [num_blocks, num_layers, kvFactor, blockSize], so dim-0 stride = + # num_layers * kvFactor * phys_block_bytes (not phys_block_bytes). + # Using the input stride keeps both contiguous test path and the + # strided prod path correct. FP8 = 1 byte per element, so the byte + # stride is the same as the element stride after recast. + kv_block_stride = kv_fused.layout.stride[0] + # KV view: [phys_block_kv, head_dim, num_phys_blocks] FP8 # Each TMA loads one physical block; multiple TMAs fill a compute tile. kv_layout = cute.make_layout( (phys_block_kv, self.head_dim, num_phys_blocks), - stride=(self.head_dim, 1, phys_block_bytes), + stride=(self.head_dim, 1, kv_block_stride), ) a = cute.make_tensor(kv_fp8.iterator, kv_layout) @@ -394,7 +402,7 @@ def __call__( # [phys_block_kv, num_phys_blocks] float32 (after recast) scale_fp8_layout = cute.make_layout( (phys_block_kv * 4, num_phys_blocks), - stride=(1, phys_block_bytes), + stride=(1, kv_block_stride), ) scale_fp8 = cute.make_tensor(kv_fp8.iterator + scale_offset_elems, scale_fp8_layout) scales = cute.recast_tensor(scale_fp8, cutlass.Float32) diff --git a/tests/scripts/cute_dsl_kernels/paged_mqa_logits/run_fp4.py b/tests/scripts/cute_dsl_kernels/paged_mqa_logits/run_fp4.py index 1d91602f3a6d..d9b781a57340 100644 --- a/tests/scripts/cute_dsl_kernels/paged_mqa_logits/run_fp4.py +++ b/tests/scripts/cute_dsl_kernels/paged_mqa_logits/run_fp4.py @@ -112,22 +112,65 @@ def _cast_back_from_fp4(packed: torch.Tensor, sf: torch.Tensor, gran_k: int = 32 return x * sf[:, torch.arange(n, device=packed.device) // gran_k] -def _kv_cache_cast_to_fp4(x: torch.Tensor): - num_blocks, block_size, num_heads, head_dim = x.shape +def _apply_utccp_chunk_layout(sf_per_block: torch.Tensor) -> torch.Tensor: + """Reorder SF tokens within each 128-token atom to UTCCP chunk layout. + + Equivalent to the SMEM-side ``utccp_required_smem_warp_transpose``: maps + linear token order [t0, t1, ..., t127] to chunk order + [t0, t32, t64, t96, t1, t33, ...] per 128-int32 atom. + + Math: ``sf_out[m*4 + k] = sf_in[k*32 + m]`` per atom. + + Used when remove_online_sf_transpose=True so the kernel can skip the + runtime warp_transpose. Only valid when phys_block_kv == 128 (1 phys + block = 1 UTCCP atom of 128 tokens). + """ + assert sf_per_block.size(-1) == 128, ( + f"chunk layout requires atom_size=128 (phys_block_kv=128); got {sf_per_block.size(-1)}" + ) + new_shape = sf_per_block.shape[:-1] + (128,) + return ( + sf_per_block.reshape(*sf_per_block.shape[:-1], 4, 32) + .transpose(-1, -2) + .contiguous() + .reshape(*new_shape) + ) + + +def _kv_cache_cast_to_fp4(x: torch.Tensor, remove_online_sf_transpose: bool = False): + # page_size here = phys_block_kv (tokens per physical KV page). + num_blocks, page_size, num_heads, head_dim = x.shape assert num_heads == 1 and head_dim == 128 + # Mirror the kernel's graceful fallback: chunk layout only valid for + # page_size == 128 (1 phys page = 1 UTCCP atom of 128 tokens). + if remove_online_sf_transpose and page_size != 128: + print( + f"[_kv_cache_cast_to_fp4] remove_online_sf_transpose=True ignored: " + f"requires page_size (phys_block_kv) == 128, got {page_size}. " + f"Falling back to False." + ) + remove_online_sf_transpose = False + x_scaled, sf = _per_token_cast_to_fp4(x.view(-1, head_dim), gran_k=32) - x_back = _cast_back_from_fp4(x_scaled, sf, gran_k=32).view(num_blocks, block_size, 1, head_dim) + x_back = _cast_back_from_fp4(x_scaled, sf, gran_k=32).view(num_blocks, page_size, 1, head_dim) x_fp4 = torch.empty( - (num_blocks, block_size * (head_dim // 2 + 4)), + (num_blocks, page_size * (head_dim // 2 + 4)), device=x.device, dtype=torch.uint8, ) - x_fp4[:, : block_size * head_dim // 2] = x_scaled.view( - num_blocks, block_size * head_dim // 2 + x_fp4[:, : page_size * head_dim // 2] = x_scaled.view( + num_blocks, page_size * head_dim // 2 ).view(torch.uint8) - x_fp4[:, block_size * head_dim // 2 :] = sf.view(num_blocks, block_size).view(torch.uint8) + + # Reorder SF tokens to UTCCP chunk layout (mirrors what the in-kernel + # warp_transpose would do at runtime), so the kernel can skip the + # runtime SMEM transpose when remove_online_sf_transpose=True. + sf_per_block = sf.view(num_blocks, page_size) # (num_blocks, 128) int32 + if remove_online_sf_transpose: + sf_per_block = _apply_utccp_chunk_layout(sf_per_block) + x_fp4[:, page_size * head_dim // 2 :] = sf_per_block.view(torch.uint8) return ( - x_fp4.view(num_blocks, block_size, num_heads, head_dim // 2 + 4), + x_fp4.view(num_blocks, page_size, num_heads, head_dim // 2 + 4), x_back.to(x.dtype), ) @@ -370,6 +413,7 @@ def _compile_fp4_kernel( num_epi_subtiles: int, epi_dtype, output_dtype, + remove_online_sf_transpose: bool = False, ): """Compile FP4 kernel with fake tensors + TVM FFI; cached by static config.""" key = ( @@ -382,6 +426,7 @@ def _compile_fp4_kernel( num_epi_subtiles, epi_dtype, output_dtype, + remove_online_sf_transpose, ) if key in _compiled_cache: return _compiled_cache[key] @@ -428,6 +473,7 @@ def _compile_fp4_kernel( num_epi_subtiles=num_epi_subtiles, epi_dtype=epi_dtype, output_dtype=output_dtype, + remove_online_sf_transpose=remove_online_sf_transpose, ) compiled = cute.compile( kernel, @@ -466,6 +512,7 @@ def fp4_paged_mqa_logits( epi_dtype=cutlass.Float32, output_dtype=cutlass.Float32, num_sms: int = 148, + remove_online_sf_transpose: bool = False, ) -> torch.Tensor: """Standalone wrapper around FP4MQALogitsKernel; no trtllm dependency. @@ -513,6 +560,7 @@ def fp4_paged_mqa_logits( num_epi_subtiles, epi_dtype, output_dtype, + remove_online_sf_transpose=remove_online_sf_transpose, ) compiled( kv_flat, @@ -609,6 +657,7 @@ def run( seed: int = 42, num_sms: int = 148, verify_meta: bool = False, + remove_online_sf_transpose: bool = False, ) -> float: """Generate random inputs, run kernel, compare to reference, print result. @@ -660,7 +709,9 @@ def run( .to(torch.bfloat16) ) - kv_fused, kv_sim = _kv_cache_cast_to_fp4(kv_cache) + kv_fused, kv_sim = _kv_cache_cast_to_fp4( + kv_cache, remove_online_sf_transpose=remove_online_sf_transpose + ) # Reference uses original (B, next_n) layout regardless of how the kernel # is invoked below. @@ -710,6 +761,7 @@ def run( epi_dtype=epi_dtype, output_dtype=output_dtype, num_sms=num_sms, + remove_online_sf_transpose=remove_online_sf_transpose, ) positions = ( @@ -794,6 +846,15 @@ def run( action="store_true", help="verify CuTe DSL schedule_meta kernel against the pure-Python reference", ) + parser.add_argument( + "--remove_online_sf_transpose", + action="store_true", + help=( + "skip in-kernel SMEM warp_transpose for KV SF; host pre-arranges " + "GMEM SF in UTCCP chunk layout (A/B perf test). " + "Requires phys_block_kv=128; auto-falls-back to False otherwise." + ), + ) args = parser.parse_args() print("=== FP4 paged MQA logits standalone test ===") @@ -811,4 +872,5 @@ def run( tol=args.tol, num_sms=args.num_sms, verify_meta=args.verify_meta, + remove_online_sf_transpose=args.remove_online_sf_transpose, ) diff --git a/tests/unittest/_torch/attention/sparse/test_cute_dsl_fp4_paged_mqa_logits.py b/tests/unittest/_torch/attention/sparse/test_cute_dsl_fp4_paged_mqa_logits.py index cecb11080b48..8dae8201ba89 100644 --- a/tests/unittest/_torch/attention/sparse/test_cute_dsl_fp4_paged_mqa_logits.py +++ b/tests/unittest/_torch/attention/sparse/test_cute_dsl_fp4_paged_mqa_logits.py @@ -139,9 +139,45 @@ def cast_back_from_fp4( # --------------------------------------------------------------------------- -def kv_cache_cast_to_fp4(x: torch.Tensor): - num_blocks, block_size, num_heads, head_dim = x.shape +def _apply_utccp_chunk_layout(sf_per_block: torch.Tensor) -> torch.Tensor: + """Reorder SF tokens within each 128-token atom to UTCCP chunk layout. + + Equivalent to the SMEM-side ``utccp_required_smem_warp_transpose``: maps + linear token order [t0, t1, ..., t127] to chunk order + [t0, t32, t64, t96, t1, t33, ...] per 128-int32 atom. + + Math: ``sf_out[m*4 + k] = sf_in[k*32 + m]`` per atom. + + Used when ``remove_online_sf_transpose=True`` so the kernel can skip the + runtime warp_transpose. Only valid when phys_block_kv == 128 (1 phys + block = 1 UTCCP atom of 128 tokens). + """ + assert sf_per_block.size(-1) == 128, ( + f"chunk layout requires atom_size=128 (phys_block_kv=128); got {sf_per_block.size(-1)}" + ) + new_shape = sf_per_block.shape[:-1] + (128,) + return ( + sf_per_block.reshape(*sf_per_block.shape[:-1], 4, 32) + .transpose(-1, -2) + .contiguous() + .reshape(*new_shape) + ) + + +def kv_cache_cast_to_fp4(x: torch.Tensor, remove_online_sf_transpose: bool = False): + # page_size here = phys_block_kv (tokens per physical KV page). + num_blocks, page_size, num_heads, head_dim = x.shape assert num_heads == 1 and head_dim == 128 + # Mirror the kernel's graceful fallback: chunk layout only valid for + # page_size == 128 (1 phys page = 1 UTCCP atom of 128 tokens). + if remove_online_sf_transpose and page_size != 128: + print( + f"[kv_cache_cast_to_fp4] remove_online_sf_transpose=True ignored: " + f"requires page_size (phys_block_kv) == 128, got {page_size}. " + f"Falling back to False." + ) + remove_online_sf_transpose = False + x_scaled, sf = per_token_cast_to_fp4( x.view(-1, head_dim), use_ue8m0=True, @@ -153,18 +189,25 @@ def kv_cache_cast_to_fp4(x: torch.Tensor): sf, gran_k=32, use_packed_ue8m0=True, - ).view(num_blocks, block_size, 1, head_dim) + ).view(num_blocks, page_size, 1, head_dim) x_fp4 = torch.empty( - (num_blocks, block_size * (head_dim // 2 + 4)), + (num_blocks, page_size * (head_dim // 2 + 4)), device=x.device, dtype=torch.uint8, ) - x_fp4[:, : block_size * head_dim // 2] = x_scaled.view( - num_blocks, block_size * head_dim // 2 + x_fp4[:, : page_size * head_dim // 2] = x_scaled.view( + num_blocks, page_size * head_dim // 2 ).view(torch.uint8) - x_fp4[:, block_size * head_dim // 2 :] = sf.view(num_blocks, block_size).view(torch.uint8) + + # Reorder SF tokens to UTCCP chunk layout (mirrors what the in-kernel + # warp_transpose would do at runtime), so the kernel can skip the + # runtime SMEM transpose when remove_online_sf_transpose=True. + sf_per_block = sf.view(num_blocks, page_size) # (num_blocks, 128) int32 + if remove_online_sf_transpose: + sf_per_block = _apply_utccp_chunk_layout(sf_per_block) + x_fp4[:, page_size * head_dim // 2 :] = sf_per_block.view(torch.uint8) return ( - x_fp4.view(num_blocks, block_size, num_heads, head_dim // 2 + 4), + x_fp4.view(num_blocks, page_size, num_heads, head_dim // 2 + 4), x_cast_back.to(x.dtype), ) @@ -375,7 +418,15 @@ def test_cute_dsl_fp4_paged_mqa_logits( ) # Quantize KV cache to fused FP4 layout. - kv_fused, kv_simulated = kv_cache_cast_to_fp4(kv_cache) + # Exercise the remove_online_sf_transpose path when supported (only valid + # for phys_block_kv=128). For other page sizes, both host helper and + # kernel silently fall back to False, but we skip enabling to avoid + # fallback print noise during the test sweep. + remove_online_sf_transpose = phys_block_kv == 128 + + kv_fused, kv_simulated = kv_cache_cast_to_fp4( + kv_cache, remove_online_sf_transpose=remove_online_sf_transpose + ) # Schedule metadata. DG c491439e requires 2D context_lens, and block_kv # must be 64 to align metadata SPLIT_KV (= block_kv*4) with the compute @@ -413,6 +464,7 @@ def test_cute_dsl_fp4_paged_mqa_logits( num_epi_subtiles=num_epi_subtiles, epi_dtype=epi_dtype, output_dtype=output_dtype, + remove_online_sf_transpose=remove_online_sf_transpose, ) assert logits.dtype == output_dtype @@ -641,6 +693,43 @@ def _generate_bench_data( } +def _choose_atom_split(batch, ctx, next_n, num_sms=148, split_kv_tokens=256, tie="max_na"): + """Pick (num_atoms, atom_size) decomposition of next_n minimizing wave count; + tie-break configurable via `tie`: + - "max_na": prefer LARGEST num_atoms = smallest atom = most SMs busy per + wave; pays HBM cost of num_atoms× KV re-reads. + - "max_atom": prefer LARGEST atom = smallest num_atoms = least HBM cost; + may leave SMs idle when ntask < num_sms. + + Kernel natively supports atom ∈ {1, 2, 3}. For next_n not divisible by any + of these (e.g. next_n=4), caller-side atom-split splits Q dim into + num_atoms groups of atom_size each; KV is read num_atoms× (1× HBM cost + per atom). + + Strategy: + 1. Enumerate (num_atoms, atom) with num_atoms * atom == next_n, + atom ∈ {1, 2, 3}. + 2. Compute waves = ceil(B * num_atoms * ceil(ctx / split_kv) / num_sms). + 3. Pick min waves; tie-break per `tie` param. + + Returns (num_atoms, atom).""" + cands = [] + for atom in (1, 2, 3): + if next_n % atom == 0: + na = next_n // atom + ntask = batch * na * ((ctx + split_kv_tokens - 1) // split_kv_tokens) + waves = (ntask + num_sms - 1) // num_sms + cands.append((waves, na, atom)) + if tie == "max_na": + cands.sort(key=lambda x: (x[0], -x[1])) # min waves, then MAX na + elif tie == "max_atom": + cands.sort(key=lambda x: (x[0], x[1])) # min waves, then MIN na (= max atom) + else: + raise ValueError(f"unknown tie={tie!r}; expected 'max_na' or 'max_atom'") + _, na, atom = cands[0] + return na, atom + + def benchmark_fp4_paged_mqa_logits( batch_sizes, next_ns, @@ -677,8 +766,10 @@ def benchmark_fp4_paged_mqa_logits( f"num_epi_subtiles={num_epi_subtiles} mode={mode_str} block_kv={block_kv}" ) hdr = ( - f"{'batch':>5s} {'ctx':>7s} {'next_n':>6s} {'nblk':>7s} | " - f"{'DSL(us)':>8s} {'DG(us)':>8s} {'DG/DSL':>7s}" + f"{'batch':>5s} {'ctx':>7s} {'next_n':>6s} {'nblk':>7s} {'ntask':>6s} | " + f"{'maxAtom':>7s} {'DSL(us)':>8s} | " + f"{'maxNa':>5s} {'DSL(us)':>8s} {'max_atom/max_na':>15s} | " + f"{'DG(us)':>8s} {'DG/DSL_max_atom':>15s} {'DG/DSL_max_na':>13s}" ) print(hdr) print("-" * len(hdr)) @@ -687,6 +778,30 @@ def benchmark_fp4_paged_mqa_logits( for context_len in context_lens: for batch_size in batch_sizes: nblk = batch_size * ((context_len + block_kv - 1) // block_kv) + SPLIT_KV_TOKENS = 256 + # Pick both atom-split strategies for A/B comparison: + # max_atom (baseline): min waves, tie-break max atom (least HBM) + # max_na (experimental): min waves, tie-break max num_atoms (more SMs busy) + na_base, atom_base = _choose_atom_split( + batch_size, + context_len, + next_n, + num_sms=num_sms, + split_kv_tokens=SPLIT_KV_TOKENS, + tie="max_atom", + ) + na_exp, atom_exp = _choose_atom_split( + batch_size, + context_len, + next_n, + num_sms=num_sms, + split_kv_tokens=SPLIT_KV_TOKENS, + tie="max_na", + ) + # ntask reflects baseline pick (matches existing log conventions). + ntask = ( + batch_size * na_base * ((context_len + SPLIT_KV_TOKENS - 1) // SPLIT_KV_TOKENS) + ) data = _generate_bench_data( batch_size, @@ -698,24 +813,38 @@ def benchmark_fp4_paged_mqa_logits( varlen=varlen, ) - # FP4 kernel natively supports next_n ∈ {1, 2, 3}. For - # next_n == 4 we apply DG's caller-side atom-split: - # [B,4,H,D//2] → [2B,2,H,D//2], context_lens / block_table - # are repeated 2× (mirrors DeepGEMM's kNextNAtom=2; 2× HBM). - # See study-deepseek-v4/DeepGEMM/tests/test_attention_post_merge.py - # for the reference pattern. - if next_n == 4: - exp_B = batch_size * 2 - dsl_q_fp4 = data["q_fp4"].reshape(exp_B, 2, num_heads, head_dim // 2) - dsl_sf_q = data["sf_q"].reshape(exp_B, 2, num_heads) - # weights [B*4, H] is layout-equivalent to [exp_B*2, H], unchanged. - dsl_ctx_lens = data["context_lens"].repeat_interleave(2) - dsl_block_table = data["block_table"].repeat_interleave(2, dim=0) - else: - dsl_q_fp4 = data["q_fp4"] - dsl_sf_q = data["sf_q"] - dsl_ctx_lens = data["context_lens"] - dsl_block_table = data["block_table"] + # Helper: reshape inputs per (na, atom). Returns dict of tensors. + # `data`, `batch_size`, `num_heads`, `head_dim` bound via default + # args for explicit early-binding (ruff F821 can't track deeply + # nested closure captures). + def _split( + na, + atom, + data=data, + batch_size=batch_size, + num_heads=num_heads, + head_dim=head_dim, + ): + if na > 1: + return { + "q": data["q_fp4"].reshape( + batch_size * na, atom, num_heads, head_dim // 2 + ), + "sf_q": data["sf_q"].reshape(batch_size * na, atom, num_heads), + "ctx_lens": data["context_lens"].repeat_interleave(na), + "block_table": data["block_table"].repeat_interleave(na, dim=0), + } + return { + "q": data["q_fp4"], + "sf_q": data["sf_q"], + "ctx_lens": data["context_lens"], + "block_table": data["block_table"], + } + + base_t = _split(na_base, atom_base) + # Experimental only differs from baseline when strategies diverge. + strats_diverge = (na_base, atom_base) != (na_exp, atom_exp) + exp_t = _split(na_exp, atom_exp) if strats_diverge else base_t # DG metadata: same convention as the FP8 bench. SPLIT_KV = # block_kv * 4 must equal DSL's compute SPLIT_KV = 256, so @@ -723,34 +852,60 @@ def benchmark_fp4_paged_mqa_logits( # 2D `(exp_B, 1)` context_lens gives num_next_n_atoms=1 — the # DSL kernel processes all real next_n positions in one atom. DG_METADATA_BLOCK_KV = 64 - dsl_schedule_meta = get_paged_mqa_logits_metadata( - dsl_ctx_lens.unsqueeze(-1), + dsl_schedule_meta_base = get_paged_mqa_logits_metadata( + base_t["ctx_lens"].unsqueeze(-1), DG_METADATA_BLOCK_KV, num_sms, ) - - def dsl_fn( - data=data, - dsl_q_fp4=dsl_q_fp4, - dsl_sf_q=dsl_sf_q, - dsl_ctx_lens=dsl_ctx_lens, - dsl_block_table=dsl_block_table, - ): - torch.ops.trtllm.cute_dsl_fp4_paged_mqa_logits( - dsl_q_fp4, - dsl_sf_q, - data["kv_fused"], - data["weights"], - dsl_ctx_lens, - dsl_block_table, - dsl_schedule_meta, - data["max_model_len"], - num_epi_subtiles=num_epi_subtiles, - epi_dtype=epi_dtype, - output_dtype=output_dtype, + dsl_schedule_meta_exp = ( + get_paged_mqa_logits_metadata( + exp_t["ctx_lens"].unsqueeze(-1), + DG_METADATA_BLOCK_KV, + num_sms, ) + if strats_diverge + else dsl_schedule_meta_base + ) - dsl_us = _bench_kineto(dsl_fn, "kernel_cutlass_kernel", num_iterations) * 1e6 + def _make_dsl_fn(t, schedule_meta, data=data): + def dsl_fn(t=t, schedule_meta=schedule_meta, data=data): + torch.ops.trtllm.cute_dsl_fp4_paged_mqa_logits( + t["q"], + t["sf_q"], + data["kv_fused"], + data["weights"], + t["ctx_lens"], + t["block_table"], + schedule_meta, + data["max_model_len"], + num_epi_subtiles=num_epi_subtiles, + epi_dtype=epi_dtype, + output_dtype=output_dtype, + ) + + return dsl_fn + + # baseline (max_atom) and experimental (max_na) timing. + base_us = ( + _bench_kineto( + _make_dsl_fn(base_t, dsl_schedule_meta_base), + "kernel_cutlass_kernel", + num_iterations, + ) + * 1e6 + ) + if strats_diverge: + exp_us = ( + _bench_kineto( + _make_dsl_fn(exp_t, dsl_schedule_meta_exp), + "kernel_cutlass_kernel", + num_iterations, + ) + * 1e6 + ) + else: + exp_us = base_us + strat_speedup = base_us / exp_us # >1 = max_na faster dg_us = None try: @@ -783,11 +938,17 @@ def dg_fn(data=data, dg_ctx_2d=dg_ctx_2d, q_fp4_dg=q_fp4_dg): except RuntimeError: pass - ratio_str = f"{dg_us / dsl_us:6.3f}x" if dg_us else " N/A " + # DG vs both DSL variants for direct comparison. + ratio_base_str = f"{dg_us / base_us:14.3f}x" if dg_us else " N/A " + ratio_exp_str = f"{dg_us / exp_us:12.3f}x" if dg_us else " N/A " dg_str = f"{dg_us:7.1f}" if dg_us else " N/A" + base_lab = f"{na_base}/{atom_base}" + exp_lab = f"{na_exp}/{atom_exp}" print( - f"{batch_size:5d} {context_len:7d} {next_n:6d} " - f"{nblk:7d} | {dsl_us:7.1f} {dg_str} {ratio_str}" + f"{batch_size:5d} {context_len:7d} {next_n:6d} {nblk:7d} {ntask:6d} | " + f"{base_lab:>7s} {base_us:8.1f} | " + f"{exp_lab:>5s} {exp_us:8.1f} {strat_speedup:14.3f}x | " + f"{dg_str} {ratio_base_str} {ratio_exp_str}" ) del data @@ -823,7 +984,7 @@ def dg_fn(data=data, dg_ctx_2d=dg_ctx_2d, q_fp4_dg=q_fp4_dg): "--context_len", type=int, nargs="+", - default=[4096, 8192, 16384, 32768, 65536, 131072], + default=[1024, 2048, 4096, 8192, 16384, 32768, 65536, 131072], help="context lengths (default: 4096 8192 16384 32768 65536 131072)", ) parser.add_argument( @@ -851,7 +1012,7 @@ def dg_fn(data=data, dg_ctx_2d=dg_ctx_2d, q_fp4_dg=q_fp4_dg): "--num_epi_subtiles", type=int, default=1, - choices=[1, 2, 4], + choices=[1, 2, 3, 4], help="epilogue sub-tile count (default: 1)", ) parser.add_argument( diff --git a/tests/unittest/_torch/attention/sparse/test_cute_dsl_fp8_paged_mqa_logits.py b/tests/unittest/_torch/attention/sparse/test_cute_dsl_fp8_paged_mqa_logits.py index 99cd2ceb7052..3ca2d837cbca 100644 --- a/tests/unittest/_torch/attention/sparse/test_cute_dsl_fp8_paged_mqa_logits.py +++ b/tests/unittest/_torch/attention/sparse/test_cute_dsl_fp8_paged_mqa_logits.py @@ -546,6 +546,35 @@ def _generate_bench_data( } +def _choose_atom_split( + batch, ctx, next_n, num_sms=148, split_kv_tokens=256, tie="max_na", kernel_atoms=(1, 2, 3, 4) +): + """Pick (num_atoms, atom_size) decomposition of next_n minimizing wave count; + tie-break configurable via `tie`: + - "max_na": prefer LARGEST num_atoms = smallest atom = most SMs busy per + wave; pays HBM cost of num_atoms× KV re-reads. + - "max_atom": prefer LARGEST atom = smallest num_atoms = least HBM cost. + + FP8 kernel natively supports atom ∈ {1, 2, 3, 4} (FP4 differs: {1, 2, 3}). + + Returns (num_atoms, atom).""" + cands = [] + for atom in kernel_atoms: + if next_n % atom == 0: + na = next_n // atom + ntask = batch * na * ((ctx + split_kv_tokens - 1) // split_kv_tokens) + waves = (ntask + num_sms - 1) // num_sms + cands.append((waves, na, atom)) + if tie == "max_na": + cands.sort(key=lambda x: (x[0], -x[1])) + elif tie == "max_atom": + cands.sort(key=lambda x: (x[0], x[1])) + else: + raise ValueError(f"unknown tie={tie!r}") + _, na, atom = cands[0] + return na, atom + + def benchmark_fp8_paged_mqa_logits( batch_sizes, next_ns, @@ -582,8 +611,10 @@ def benchmark_fp8_paged_mqa_logits( ) is_non_default = output_dtype != torch.float32 or num_epi_subtiles != 1 hdr = ( - f"{'batch':>5s} {'ctx':>7s} {'next_n':>6s} {'nblk':>7s} | " - f"{'DSL(us)':>8s} {'DG(fp32,us)':>12s} {'DG/DSL':>7s}" + f"{'batch':>5s} {'ctx':>7s} {'next_n':>6s} {'nblk':>7s} {'ntask':>6s} | " + f"{'maxAtom':>7s} {'DSL(us)':>8s} | " + f"{'maxNa':>5s} {'DSL(us)':>8s} {'max_atom/max_na':>15s} | " + f"{'DG(us)':>11s} {'DG/DSL_max_atom':>15s} {'DG/DSL_max_na':>13s}" ) if is_non_default: hdr += f" {'DSL(fp32,us)':>13s} {'DSL(fp32)/DSL':>13s}" @@ -594,6 +625,32 @@ def benchmark_fp8_paged_mqa_logits( for context_len in context_lens: for batch_size in batch_sizes: nblk = batch_size * ((context_len + block_kv - 1) // block_kv) + SPLIT_KV_TOKENS = 256 + # Pick both atom-split strategies for A/B comparison: + # max_atom (baseline): min waves, tie-break max atom (least HBM) + # max_na (experimental): min waves, tie-break max num_atoms (more SMs busy) + # FP8 kernel supports atom ∈ {1, 2, 3, 4}. + na_base, atom_base = _choose_atom_split( + batch_size, + context_len, + next_n, + num_sms=num_sms, + split_kv_tokens=SPLIT_KV_TOKENS, + tie="max_atom", + kernel_atoms=(1, 2, 3, 4), + ) + na_exp, atom_exp = _choose_atom_split( + batch_size, + context_len, + next_n, + num_sms=num_sms, + split_kv_tokens=SPLIT_KV_TOKENS, + tie="max_na", + kernel_atoms=(1, 2, 3, 4), + ) + ntask = ( + batch_size * na_base * ((context_len + SPLIT_KV_TOKENS - 1) // SPLIT_KV_TOKENS) + ) data = _generate_bench_data( batch_size, @@ -605,34 +662,91 @@ def benchmark_fp8_paged_mqa_logits( varlen=varlen, ) + # Helper: reshape Q + repeat ctx/block_table per (na, atom). + # weights [B*next_n, H] = [B*na*atom, H] needs no reshape. + # `data`, `batch_size`, `num_heads`, `head_dim` bound via default + # args (explicit early-binding; ruff F821 can't track deeply + # nested closure captures). + def _split( + na, + atom, + data=data, + batch_size=batch_size, + num_heads=num_heads, + head_dim=head_dim, + ): + if na > 1: + return { + "q": data["q_fp8"].reshape(batch_size * na, atom, num_heads, head_dim), + "ctx_lens": data["context_lens"].repeat_interleave(na), + "block_table": data["block_table"].repeat_interleave(na, dim=0), + } + return { + "q": data["q_fp8"], + "ctx_lens": data["context_lens"], + "block_table": data["block_table"], + } + + base_t = _split(na_base, atom_base) + strats_diverge = (na_base, atom_base) != (na_exp, atom_exp) + exp_t = _split(na_exp, atom_exp) if strats_diverge else base_t + # See `test_cute_dsl_fp8_paged_mqa_logits` for full reasoning # on the `block_kv = 64` choice. Short version: DG metadata # SPLIT_KV = block_kv * 4; we need SPLIT_KV = 256 (DSL # compute tile = 128 × kNumMathWarpGroups = 2), so pass 64. - # 2D `(B, 1)` context_lens forces num_next_n_atoms = 1. + # 2D `(B*na, 1)` context_lens forces num_next_n_atoms = 1. DG_METADATA_BLOCK_KV = 64 - dsl_schedule_meta = get_paged_mqa_logits_metadata( - data["context_lens"].unsqueeze(-1), + dsl_schedule_meta_base = get_paged_mqa_logits_metadata( + base_t["ctx_lens"].unsqueeze(-1), DG_METADATA_BLOCK_KV, num_sms, ) + dsl_schedule_meta_exp = ( + get_paged_mqa_logits_metadata( + exp_t["ctx_lens"].unsqueeze(-1), + DG_METADATA_BLOCK_KV, + num_sms, + ) + if strats_diverge + else dsl_schedule_meta_base + ) + + def _make_dsl_fn(t, schedule_meta, data=data): + def dsl_fn(t=t, schedule_meta=schedule_meta, data=data): + torch.ops.trtllm.cute_dsl_fp8_paged_mqa_logits( + t["q"], + data["kv_fused"], + data["weights"], + t["ctx_lens"], + t["block_table"], + schedule_meta, + data["max_model_len"], + num_epi_subtiles=num_epi_subtiles, + epi_dtype=output_dtype, + acc_dtype=output_dtype, + output_dtype=output_dtype, + ) + + return dsl_fn - def dsl_fn(data=data): - torch.ops.trtllm.cute_dsl_fp8_paged_mqa_logits( - data["q_fp8"], - data["kv_fused"], - data["weights"], - data["context_lens"], - data["block_table"], - dsl_schedule_meta, - data["max_model_len"], - num_epi_subtiles=num_epi_subtiles, - epi_dtype=output_dtype, - acc_dtype=output_dtype, - output_dtype=output_dtype, + base_us = _profile_kernel_us( + _make_dsl_fn(base_t, dsl_schedule_meta_base), + num_warmup, + num_iterations, + ) + if strats_diverge: + exp_us = _profile_kernel_us( + _make_dsl_fn(exp_t, dsl_schedule_meta_exp), + num_warmup, + num_iterations, ) + else: + exp_us = base_us + strat_speedup = base_us / exp_us # >1 = max_na faster - dsl_us = _profile_kernel_us(dsl_fn, num_warmup, num_iterations) + # Alias for fp32-variant code below (uses baseline schedule). + dsl_schedule_meta = dsl_schedule_meta_base dg_us = None try: @@ -691,15 +805,21 @@ def dsl_f32_fn(data=data): dsl_f32_us = _profile_kernel_us(dsl_f32_fn, num_warmup, num_iterations) - ratio_str = f"{dg_us / dsl_us:6.3f}x" if dg_us else " N/A " - dg_str = f"{dg_us:11.1f}" if dg_us else " N/A" + ratio_base_str = f"{dg_us / base_us:14.3f}x" if dg_us else " N/A " + ratio_exp_str = f"{dg_us / exp_us:12.3f}x" if dg_us else " N/A " + dg_str = f"{dg_us:10.1f}" if dg_us else " N/A" + base_lab = f"{na_base}/{atom_base}" + exp_lab = f"{na_exp}/{atom_exp}" line = ( f"{batch_size:5d} {context_len:7d} {next_n:6d} " - f"{nblk:7d} | {dsl_us:7.1f} {dg_str} {ratio_str}" + f"{nblk:7d} {ntask:6d} | " + f"{base_lab:>7s} {base_us:8.1f} | " + f"{exp_lab:>5s} {exp_us:8.1f} {strat_speedup:14.3f}x | " + f"{dg_str} {ratio_base_str} {ratio_exp_str}" ) if is_non_default: f32_str = f"{dsl_f32_us:12.1f}" if dsl_f32_us else " N/A" - f32_ratio = f"{dsl_f32_us / dsl_us:12.3f}x" if dsl_f32_us else " N/A " + f32_ratio = f"{dsl_f32_us / base_us:12.3f}x" if dsl_f32_us else " N/A " line += f" {f32_str} {f32_ratio}" print(line) @@ -734,7 +854,7 @@ def dsl_f32_fn(data=data): "--context_len", type=int, nargs="+", - default=[4096, 8192, 16384, 32768, 65536, 131072], + default=[1024, 2048, 4096, 8192, 16384, 32768, 65536, 131072], help="context lengths (default: 4096 8192 16384 32768 65536 131072)", ) parser.add_argument("--warmup", type=int, default=10, help="warmup iterations (default: 10)") @@ -750,7 +870,7 @@ def dsl_f32_fn(data=data): "--num_epi_subtiles", type=int, default=1, - choices=[1, 2, 4], + choices=[1, 2, 3, 4], help="epilogue sub-tile count (default: 1)", ) parser.add_argument( diff --git a/tests/unittest/_torch/attention/sparse/test_dsa_indexer.py b/tests/unittest/_torch/attention/sparse/test_dsa_indexer.py index 0aae4ff52d20..67b4d5641565 100644 --- a/tests/unittest/_torch/attention/sparse/test_dsa_indexer.py +++ b/tests/unittest/_torch/attention/sparse/test_dsa_indexer.py @@ -1221,7 +1221,8 @@ def test_indexer_decode_with_paged_kv_cache_fp4(batch_size, next_n, backend): * "dsl" -> torch.ops.trtllm.cute_dsl_fp4_paged_mqa_logits. The DSL FP4 kernel only supports next_n ∈ {1, 2, 3} natively, so the test caller-side reshapes [B, next_n, ...] -> [B*factor, - eff_next_n, ...] via `_pick_fp4_dsl_expand` when next_n > 3. + eff_next_n, ...] via `_pick_dsl_expand` when the wave-aware picker + chooses to split (factor > 1). - reference fed FP4-simulated bf16 inputs (dequantize the same FP4 bytes the kernel sees) so FP4 quantization noise cancels in the diff check. @@ -1236,7 +1237,7 @@ def test_indexer_decode_with_paged_kv_cache_fp4(batch_size, next_n, backend): from test_cute_dsl_fp4_paged_mqa_logits import cast_back_from_fp4 from tensorrt_llm._torch.attention_backend.sparse.dsa import \ - _pick_fp4_dsl_expand + _pick_dsl_expand from tensorrt_llm.deep_gemm import fp8_fp4_paged_mqa_logits use_dsl = backend == "dsl" @@ -1317,17 +1318,37 @@ def _force_direct_path(meta): meta.scheduler_metadata_buffer_full_next_n = torch.zeros( (meta.num_sms + 1, 2), device='cuda', dtype=torch.int32) - def _force_dsl_fp4_expand_setup(meta): - """Mirror DSAtrtllmAttentionMetadata.prepare's `expand_for_dsl_fp4` + def _force_dsl_expand_setup(meta): + """Mirror DSAtrtllmAttentionMetadata.prepare's `expand_for_dsl` populate so Indexer.prepare can build the schedule from - kv_lens_expanded_cuda. Only called for DSL backend with next_n > 3. + kv_lens_expanded_cuda. Calls the wave-aware picker; expand only + happens when factor > 1. + + Always sets `expand_for_dsl=True` (analogous to dsa.py's prepare) + and writes the picker decision to `dsl_expand_factor`/`dsl_atom` so + downstream Indexer.prepare + forward read the same values. """ if meta.num_generations == 0: return - factor, _ = _pick_fp4_dsl_expand(1 + meta.max_draft_tokens) - num_tokens = meta.num_generations * factor - meta.expand_for_dsl_fp4 = True + meta.expand_for_dsl = True + next_n = 1 + meta.max_draft_tokens + # FP4 kernel supports atoms ∈ {1, 2, 3}; matches dsa.py's + # `kernel_atoms = (1, 2, 3) if use_fp4 else (1, 2, 3, 4)` path. gen_kv_lens = meta.kv_lens[meta.num_contexts:meta.num_seqs] + max_ctx = int(gen_kv_lens.max().item()) if gen_kv_lens.numel() else 0 + factor, atom = _pick_dsl_expand( + next_n, + batch_size=meta.num_generations, + max_ctx=max_ctx, + num_sms=meta.num_sms, + kernel_atoms=(1, 2, 3), + ) + meta.dsl_expand_factor = factor + meta.dsl_atom = atom + if factor <= 1: + # Picker chose kernel-native; no buffer populate needed. + return + num_tokens = meta.num_generations * factor gen_kv_lens_expanded = gen_kv_lens.repeat_interleave(factor) meta.kv_lens_expanded_host[:num_tokens].copy_(gen_kv_lens_expanded) meta.kv_lens_expanded_cuda[:num_tokens].copy_( @@ -1392,8 +1413,11 @@ def _force_dsl_fp4_expand_setup(meta): ) if not use_dsl: _force_direct_path(metadata_gen) - elif next_n > 3: - _force_dsl_fp4_expand_setup(metadata_gen) + else: + # Unconditional for DSL: picker inside decides factor (1 = native, + # >1 = atom-split). Mirrors dsa.py's `if expand_for_dsl and + # num_generations > 0` block which runs for any next_n ≥ 2. + _force_dsl_expand_setup(metadata_gen) Indexer.prepare(metadata_gen) k_gen_fp4, k_gen_scale = torch.ops.trtllm.fused_cat_fp4( @@ -1421,9 +1445,9 @@ def _force_dsl_fp4_expand_setup(meta): if use_dsl: # DSL FP4 path: q tuple split into two args, q.dtype == uint8. - # For next_n > 3, reshape via _pick_fp4_dsl_expand and use the - # expanded metadata buffers populated above. The wiring under test - # is dsa.py's sparse_attn_indexer DSL FP4 branch. + # The picker decision (factor, atom) was cached on metadata_gen + # in `_force_dsl_expand_setup`; expand only when factor > 1. + # The wiring under test is dsa.py's sparse_attn_indexer DSL FP4 branch. dsl_q = q_fp4.view(torch.uint8) dsl_sf_q = sf_q dsl_context_lens = metadata_gen.kv_lens_cuda_runtime[0:batch_size] @@ -1431,8 +1455,9 @@ def _force_dsl_fp4_expand_setup(meta): 0:batch_size] dsl_schedule_meta = metadata_gen.scheduler_metadata_buffer - if next_n > 3: - factor, eff_next_n = _pick_fp4_dsl_expand(next_n) + if metadata_gen.dsl_expand_factor > 1: + factor = metadata_gen.dsl_expand_factor + eff_next_n = metadata_gen.dsl_atom exp_B = batch_size * factor dsl_q = dsl_q.reshape(exp_B, eff_next_n, heads, pe_dim) dsl_sf_q = dsl_sf_q.reshape(exp_B, eff_next_n, heads)