From 1e100952cbeee1bf41b0c7ede5840749236a8766 Mon Sep 17 00:00:00 2001 From: Akhil Goel Date: Fri, 21 Aug 2026 17:07:58 -0700 Subject: [PATCH 1/3] [Squash of #30805] DSv4 trtllm attention backend --- .../sglang/srt/arg_groups/deepseek_v4_hook.py | 51 ++ .../layers/attention/attention_registry.py | 11 +- .../layers/attention/deepseek_v4_backend.py | 399 ++++++++++-- .../attention/deepseek_v4_trtllm_backend.py | 568 ++++++++++++++++++ .../attention/dsv4/compressor_trtllm.py | 140 +++++ .../layers/attention/dsv4/compressor_v2.py | 94 +-- .../srt/mem_cache/deepseek_v4_memory_pool.py | 100 ++- python/sglang/srt/server_args.py | 16 + python/sglang/srt/speculative/draft_utils.py | 12 +- .../backends/test_dsv4_fp8_trtllm_backend.py | 305 ++++++++++ .../test_disaggregation_dsv4.py | 4 + .../test_deepseek_v4_flash_fp4_b200_trtllm.py | 236 ++++++++ test/registered/unit/test_model_overrides.py | 1 + 13 files changed, 1845 insertions(+), 92 deletions(-) create mode 100644 python/sglang/srt/layers/attention/deepseek_v4_trtllm_backend.py create mode 100644 python/sglang/srt/layers/attention/dsv4/compressor_trtllm.py create mode 100644 test/registered/backends/test_dsv4_fp8_trtllm_backend.py create mode 100644 test/registered/models_e2e/test_deepseek_v4_flash_fp4_b200_trtllm.py diff --git a/python/sglang/srt/arg_groups/deepseek_v4_hook.py b/python/sglang/srt/arg_groups/deepseek_v4_hook.py index 4569c9b627d3..4efd982c88b9 100644 --- a/python/sglang/srt/arg_groups/deepseek_v4_hook.py +++ b/python/sglang/srt/arg_groups/deepseek_v4_hook.py @@ -135,6 +135,36 @@ def apply_deepseek_v4_defaults(server_args: ServerArgs, model_arch: str) -> None run_post_process_pass(server_args, _deepseek_v4_kv_cache_dtype) + if server_args.dsv4_attn_backend == "trtllm": + from sglang.srt.utils.common import is_sm100_supported + + assert ( + server_args.device == "cuda" and is_sm100_supported() + ), "--dsv4-attn-backend trtllm requires an SM100/SM103 (Blackwell) GPU." + # "auto" is declared-but-unmaterialized here; the resolution pipeline + # (_deepseek_v4_kv_cache_dtype above) turns it into fp8_e4m3 on cuda. + assert server_args.kv_cache_dtype in ("auto", "fp8_e4m3"), ( + "--dsv4-attn-backend trtllm requires kv_cache_dtype=fp8_e4m3, " + f"got {server_args.kv_cache_dtype}." + ) + assert ( + not server_args.enable_hisparse + ), "--dsv4-attn-backend trtllm does not support enable_hisparse." + assert not ( + server_args.attn_cp_size > 1 + or server_args.dcp_size > 1 + or server_args.enable_prefill_cp + or server_args.enable_prefill_context_parallel + or server_args.enable_dsa_prefill_context_parallel + ), ( + "--dsv4-attn-backend trtllm does not support context parallelism " + "(prefill CP, attention CP, or decode CP)." + ) + logger.info( + "DeepSeek V4 attention: trtllm backend enabled " + "(uniform-FP8 KV pool, decode + sparse prefill)." + ) + if server_args.max_running_requests is None: server_args.max_running_requests = 256 logger.warning( @@ -151,6 +181,27 @@ def apply_deepseek_v4_defaults(server_args: ServerArgs, model_arch: str) -> None server_args.speculative_eagle_topk == 1 ), f"Only EAGLE speculative algorithm with topk == 1 is supported for {model_arch}" + # FIXME(follow-up): remove once the overlap+speculative corruption is + # root-caused (tracked in the PR #30805 follow-up list). + # TEMPORARY containment: trtllm + speculative decoding under the overlap + # scheduler intermittently corrupts an int32 table consumed by the + # trtllm-gen sparse kernel (illegal memory access in + # fmhaSm100fKernel...VarSeq during concurrent GSM8K-style bursts). This + # reproduces with both TP-only and DP-attention recipes; disabling overlap + # prevents the corruption while the root cause is investigated. + if ( + server_args.dsv4_attn_backend == "trtllm" + and server_args.speculative_algorithm is not None + and not server_args.disable_overlap_schedule + ): + logger.warning( + "Disabling the overlap scheduler for the trtllm DeepSeek-V4 " + "backend with speculative decoding (temporary " + "containment for an intermittent trtllm-gen kernel memory fault; " + "see the dsv4 trtllm PR discussion)." + ) + server_args.disable_overlap_schedule = True + def validate_deepseek_v4_cp(server_args: ServerArgs) -> None: """Validate DeepSeek V4 context-parallel configuration.""" diff --git a/python/sglang/srt/layers/attention/attention_registry.py b/python/sglang/srt/layers/attention/attention_registry.py index 1d176ce7b211..2749f485b78a 100644 --- a/python/sglang/srt/layers/attention/attention_registry.py +++ b/python/sglang/srt/layers/attention/attention_registry.py @@ -166,12 +166,15 @@ def create_dsv4_backend(runner): ) return DeepseekV4HipRadixBackend(runner) else: - from sglang.srt.layers.attention.deepseek_v4_backend import ( - DeepseekV4AttnBackend, + from sglang.srt.layers.attention.deepseek_v4_trtllm_backend import ( + create_deepseek_v4_attn_backend, ) - logger.info("Using DeepseekV4AttnBackend for dsv4 attention backend (CUDA).") - return DeepseekV4AttnBackend(runner) + backend = create_deepseek_v4_attn_backend(runner) + logger.info( + f"Using {type(backend).__name__} for dsv4 attention backend (CUDA)." + ) + return backend @register_attention_backend("triton") diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index 8b18534413e3..56c6b7ab2016 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -3,6 +3,7 @@ import enum import functools import logging +from collections import deque from dataclasses import dataclass, field from typing import ( TYPE_CHECKING, @@ -73,6 +74,7 @@ from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc from sglang.srt.speculative.ragged_verify import ( RaggedVerifyMode, + compute_ragged_extend_lengths, compute_target_verify_graph_key, compute_uniform_extend_lengths, read_ragged_verify_mode, @@ -190,6 +192,25 @@ class DSV4AttnMetadata: c128_page_indices: Optional[torch.Tensor] = None c128_topk_lengths_clamp1: Optional[torch.Tensor] = None + # trtllm decode: per-ratio combined sparse tables [num_tokens, 128 + w] + # and lens [num_tokens]. Layer-invariant content (SWA columns, the whole + # c128 table, the c0 lens) is written once per step by + # init_trtllm_sparse_buffers; the per-layer forward only writes the c4 + # tail + lens (indexer top-k). c0 layers pass swa_page_indices itself as + # the table. None for prefill-style metadata and when trtllm is off. + trtllm_swa_lens: Optional[torch.Tensor] = None + trtllm_c4_indices: Optional[torch.Tensor] = None + trtllm_c4_lens: Optional[torch.Tensor] = None + trtllm_c128_indices: Optional[torch.Tensor] = None + trtllm_c128_lens: Optional[torch.Tensor] = None + # trtllm prefill: per-chunk caches built lazily by _forward_trtllm_prefill + # (prefill runs eagerly). qmeta = (cum_seq_lens_q, max_q_len, sum_q, + # seq_lens_int32); c128 = (table, lens). + trtllm_prefill_qmeta: Optional[tuple] = None + trtllm_prefill_swa_lens: Optional[torch.Tensor] = None + trtllm_prefill_c4_indices: Optional[torch.Tensor] = None + trtllm_prefill_c128: Optional[tuple] = None + c1_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) c4_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) c128_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False) @@ -233,6 +254,11 @@ def copy_(self, other: DSV4AttnMetadata) -> None: "c4_sparse_topk_lengths", "c4_sparse_page_indices", "c4_sparse_raw_indices", + "trtllm_swa_lens", + "trtllm_c4_indices", + "trtllm_c4_lens", + "trtllm_c128_indices", + "trtllm_c128_lens", ], assign_fields=[ # Recomputed by the recorded init_forward_metadata_in_graph op @@ -241,6 +267,13 @@ def copy_(self, other: DSV4AttnMetadata) -> None: "c1_flashmla_metadata", "c4_flashmla_metadata", "c128_flashmla_metadata", + # Built lazily by the eager _forward_trtllm_prefill, so they + # are None on every metadata that reaches a graph replay; two + # of them are tuples, which content-copy cannot handle. + "trtllm_prefill_qmeta", + "trtllm_prefill_swa_lens", + "trtllm_prefill_c4_indices", + "trtllm_prefill_c128", ], ) @@ -258,6 +291,14 @@ def refresh_for_breakable_cuda_graph_replay_(self, other: DSV4AttnMetadata) -> N "c4_topk_lengths_raw", "c4_topk_lengths_clamp1", "c4_sparse_topk_lengths", + # trtllm tables: static parts are rebuilt per step by the host + # metadata; content-copy keeps the captured tensor objects alive + # (the c4 tail is refilled in place by the per-layer forward). + "trtllm_swa_lens", + "trtllm_c4_indices", + "trtllm_c4_lens", + "trtllm_c128_indices", + "trtllm_c128_lens", ] reference_assign_fields = [ "page_table", @@ -268,6 +309,14 @@ def refresh_for_breakable_cuda_graph_replay_(self, other: DSV4AttnMetadata) -> N "c1_flashmla_metadata", "c4_flashmla_metadata", "c128_flashmla_metadata", + # Per-chunk lazy caches built by the eager _forward_trtllm_prefill; + # MUST be reset from the fresh metadata (None) on every breakable + # replay refresh, or a reused metadata object serves stale + # sum_q/tables to a different batch shape. + "trtllm_prefill_qmeta", + "trtllm_prefill_swa_lens", + "trtllm_prefill_c4_indices", + "trtllm_prefill_c128", ] # Keep graph-captured tensor objects alive for fields that captured # kernels read by address; overwrite only their contents. @@ -396,6 +445,66 @@ def init_flashmla_related(self, is_prefill: bool = False): self.c4_flashmla_metadata = _create_flashmla_metadata() self.c128_flashmla_metadata = _create_flashmla_metadata() + def init_trtllm_sparse_buffers(self) -> None: + """Build the trtllm per-ratio combined sparse tables (decode). + + Kernel contract per row: columns [0:128) SWA physical token indices, + [128:) compressed-tier indices, -1 invalid, capacity % 4 == 0; lens + include the 128 SWA slots. Everything layer-invariant is written here, + once per step: the SWA columns of both tables, the ENTIRE c128 table + and lens (its page list is metadata-level), and the constant SWA-only + lens (c0 layers pass swa_page_indices itself as the table). The + per-layer decode forward only writes the c4 tail + lens. + """ + + num_tokens = self.seq_lens_casual.shape[0] + assert self.swa_page_indices.shape == (num_tokens, SWA_WINDOW) + + # TILE OVERRUN GUARD: the trtllm-gen VarSeq kernel processes query + # rows in 64-token tiles and reads per-token table rows up to the + # tile boundary, i.e. up to 63 rows past num_tokens. Allocate every + # per-token tensor it reads as a 64-row-aligned parent filled with + # INERT values (-1 indices, SWA-only lens, seq_len 1) and keep + # [:num_tokens] views, so over-reads land in mapped inert memory + # instead of past the allocation (observed as segment-boundary MMU + # faults / silent garbage depending on allocator layout). + n_pad = (num_tokens + 63) // 64 * 64 + + def _tile_padded(fill, src=None, width=None): + shape = (n_pad,) if width is None else (n_pad, width) + buf = torch.full(shape, fill, **self.cuda_int32_kwargs) + if src is not None: + buf[:num_tokens].copy_(src) + return buf[:num_tokens] + + if n_pad != num_tokens: + self.seq_lens_casual = _tile_padded(1, self.seq_lens_casual) + self.swa_page_indices = _tile_padded( + -1, self.swa_page_indices, width=SWA_WINDOW + ) + self.trtllm_swa_lens = _tile_padded(SWA_WINDOW) + if self.c4_sparse_page_indices is not None: + w4 = self.c4_sparse_page_indices.shape[-1] + assert w4 % 4 == 0, f"{w4=}" + # The c4 tail + lens are per-layer values written by the decode + # forward. Initialize them INERT (-1 = invalid index, lens = + # SWA-only) rather than torch.empty: any row the per-layer fill + # does not cover must not hand the kernel allocator garbage as + # indices/lengths. + self.trtllm_c4_indices = _tile_padded(-1, width=SWA_WINDOW + w4) + self.trtllm_c4_indices[:, :SWA_WINDOW].copy_(self.swa_page_indices) + self.trtllm_c4_lens = _tile_padded(SWA_WINDOW) + if self.c128_page_indices is not None: + w128 = self.c128_page_indices.shape[-1] + assert w128 % 4 == 0, f"{w128=}" + self.trtllm_c128_indices = _tile_padded(-1, width=SWA_WINDOW + w128) + self.trtllm_c128_indices[:, :SWA_WINDOW].copy_(self.swa_page_indices) + self.trtllm_c128_indices[:, SWA_WINDOW:].copy_(self.c128_page_indices) + self.trtllm_c128_lens = _tile_padded( + SWA_WINDOW, + (self.c128_topk_lengths_clamp1 + SWA_WINDOW).to(torch.int32), + ) + @dataclass class DSV4Metadata: @@ -507,6 +616,9 @@ class DeepseekV4AttnBackend( use_captured_forward_metadata_for_breakable_cuda_graph: bool = True supports_ragged_verify_graph: bool = True needs_cpu_seq_lens: bool = False + # True on DeepseekV4TrtllmAttnBackend (deepseek_v4_trtllm_backend.py), + # which dispatches attention through the trtllm-gen sparse MLA kernel. + trtllm_attn: bool = False def shared_read_ends(self, fm: ForwardMode) -> SharedReadEnds: # Breakable-graph verify rereads shared state across segments. @@ -605,15 +717,34 @@ def __init__( self.online_c128_mtp = OnlineC128MTPController(self) self.sparse_prefill_workspace = SparsePrefillWorkspace(self.device) spec_alg = model_runner.spec_algorithm + self._prep_in_cuda_graph: bool = self._compute_prep_in_cuda_graph(model_runner) self.needs_cpu_seq_lens = not spec_alg.is_dspark() and ( - not _is_cuda or self.online_c128_mtp.enabled() + not _is_cuda + or not self._prep_in_cuda_graph + or self.online_c128_mtp.enabled() ) self.is_dspark_draft = model_runner.is_draft_worker and spec_alg.is_dspark() self.is_draft_runner = model_runner.is_draft_worker self._verify_mask = None + # Pinned compress-plan staging buffers whose owning metadata was + # dropped by replay_cuda_graph_metadata_from (see + # _keep_plan_staging_alive). + self._plan_staging_keepalive: deque = deque(maxlen=8) self.cuda_graph_swa_out_cache_loc: Optional[torch.Tensor] = None + # Pins the metadata built inside each CUDA-graph capture; see the + # comment in init_forward_metadata_in_graph. + self._captured_full_metadata_refs: List[DSV4Metadata] = [] + + def _compute_prep_in_cuda_graph(self, model_runner: ModelRunner) -> bool: + """Whether attention metadata is prepared inside CUDA-graph capture. + + The trtllm subclass (deepseek_v4_trtllm_backend.py) overrides this to + prepare on the host for trtllm + speculative + DP attention. + """ + return True + def _move_to_device(self, x: List[int]) -> torch.Tensor: pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True) return pin_tensor.to(self.device, non_blocking=True) @@ -714,10 +845,38 @@ def init_forward_metadata_decode( req_pool_indices.shape[0] == seq_lens.shape[0] == out_cache_loc.shape[0] ), f"{req_pool_indices.shape=} {seq_lens.shape=} {out_cache_loc.shape=}" - return DSV4RawDecodeMetadata( + if self._prep_in_cuda_graph: + return DSV4RawDecodeMetadata( + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc, + ) + + core_attn_metadata = self.make_core_attn_metadata( + req_to_token=self.req_to_token, + req_pool_indices_repeated=req_pool_indices, + seq_lens_casual=seq_lens, + max_seq_len=max_seq_len, + out_loc=out_cache_loc, + need_compress=True, + ) + + indexer_metadata = self.init_forward_metadata_indexer(core_attn_metadata) + + create = functools.partial( + create_paged_compressor_data, + is_prefill=False, + token_to_kv_pool=self.token_to_kv_pool, + req_to_token=self.req_to_token, req_pool_indices=req_pool_indices, seq_lens=seq_lens, - out_cache_loc=out_cache_loc, + ) + + return DSV4Metadata( + core_attn_metadata, + indexer_metadata, + c4_compress_metadata=create(compress_ratio=4), + c128_compress_metadata=create(compress_ratio=128), ) def init_forward_metadata_prefill( @@ -836,44 +995,108 @@ def init_forward_metadata_target_verify( online_c128_state_slot_offset: int = 0, ragged_layout: Optional[RaggedVerifyLayout] = None, ) -> Union[DSV4Metadata, DSV4RawVerifyMetadata]: - assert out_cache_loc is not None - bs = len(seq_lens) - if self.needs_cpu_seq_lens: - assert seq_lens_cpu is not None - seq_lens_cpu_list = seq_lens_cpu.tolist() + if self._prep_in_cuda_graph: + assert out_cache_loc is not None + bs = len(seq_lens) + if self.needs_cpu_seq_lens: + assert seq_lens_cpu is not None + seq_lens_cpu_list = seq_lens_cpu.tolist() + else: + seq_lens_cpu_list = None + if ragged_layout is None: + self.extend_seq_lens_buffer[:bs].fill_( + self.speculative_num_draft_tokens + ) + extend_seq_lens = self.extend_seq_lens_buffer[:bs] + extend_start_loc = None + verify_lens = None + total_verify_tokens = self.speculative_num_draft_tokens * bs + else: + self.extend_seq_lens_buffer[:bs].copy_(ragged_layout.verify_lens) + self.extend_start_loc_buffer[:bs].copy_(ragged_layout.extend_start_loc) + extend_seq_lens = self.extend_seq_lens_buffer[:bs] + extend_start_loc = self.extend_start_loc_buffer[:bs] + verify_lens = self.extend_seq_lens_buffer[:bs] + total_verify_tokens = ragged_layout.graph_num_tokens + + return DSV4RawVerifyMetadata( + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc, + extend_seq_lens=extend_seq_lens, + seq_lens_cpu=seq_lens_cpu_list, + c128_compress_metadata=self._make_target_verify_c128_metadata( + req_pool_indices, + seq_lens, + seq_lens_cpu_list, + extend_seq_lens, + use_prefill_cuda_graph, + online_c128_state_slot_offset, + ), + extend_start_loc=extend_start_loc, + verify_lens=verify_lens, + total_verify_tokens=total_verify_tokens, + ) else: - seq_lens_cpu_list = None + seq_lens_cpu_list = ( + seq_lens_cpu.tolist() if seq_lens_cpu is not None else seq_lens.tolist() + ) + return self.init_forward_metadata_target_verify_old( + max_seq_len=max_seq_len, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu_list, + out_cache_loc=out_cache_loc, + use_prefill_cuda_graph=use_prefill_cuda_graph, + online_c128_state_slot_offset=online_c128_state_slot_offset, + ragged_layout=ragged_layout, + ) + + def init_forward_metadata_target_verify_old( + self, + max_seq_len: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + seq_lens_cpu: Optional[List[int]] = None, + out_cache_loc: Optional[torch.Tensor] = None, + use_prefill_cuda_graph: bool = False, + online_c128_state_slot_offset: int = 0, + ragged_layout: Optional[RaggedVerifyLayout] = None, + ) -> DSV4Metadata: if ragged_layout is None: - self.extend_seq_lens_buffer[:bs].fill_(self.speculative_num_draft_tokens) - extend_seq_lens = self.extend_seq_lens_buffer[:bs] - extend_start_loc = None - verify_lens = None - total_verify_tokens = self.speculative_num_draft_tokens * bs + lengths = compute_uniform_extend_lengths( + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + extend_len=self.speculative_num_draft_tokens, + ) + extend_seq_lens = self._move_to_device(lengths.extend_seq_lens_cpu) else: - self.extend_seq_lens_buffer[:bs].copy_(ragged_layout.verify_lens) - self.extend_start_loc_buffer[:bs].copy_(ragged_layout.extend_start_loc) - extend_seq_lens = self.extend_seq_lens_buffer[:bs] - extend_start_loc = self.extend_start_loc_buffer[:bs] - verify_lens = self.extend_seq_lens_buffer[:bs] - total_verify_tokens = ragged_layout.graph_num_tokens - - return DSV4RawVerifyMetadata( + lengths = compute_ragged_extend_lengths( + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + ragged_layout=ragged_layout, + ) + extend_seq_lens = ragged_layout.verify_lens + seq_lens = lengths.seq_lens_extended + seq_lens_cpu = lengths.seq_lens_cpu_extended + extend_seq_lens_cpu = lengths.extend_seq_lens_cpu + num_tokens = lengths.num_tokens + extend_start_loc = lengths.extend_start_loc + if out_cache_loc is None: + out_cache_loc = seq_lens.new_zeros(num_tokens) + return self.init_forward_metadata_prefill( + max_seq_len=max_seq_len, req_pool_indices=req_pool_indices, seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, out_cache_loc=out_cache_loc, + num_tokens=num_tokens, extend_seq_lens=extend_seq_lens, - seq_lens_cpu=seq_lens_cpu_list, - c128_compress_metadata=self._make_target_verify_c128_metadata( - req_pool_indices, - seq_lens, - seq_lens_cpu_list, - extend_seq_lens, - use_prefill_cuda_graph, - online_c128_state_slot_offset, - ), + extend_seq_lens_cpu=extend_seq_lens_cpu, extend_start_loc=extend_start_loc, - verify_lens=verify_lens, - total_verify_tokens=total_verify_tokens, + need_compress=True, + use_prefill_cuda_graph=use_prefill_cuda_graph, + online_c128_state_slot_offset=online_c128_state_slot_offset, ) def init_forward_metadata_dspark_draft_block( @@ -1043,6 +1266,12 @@ def init_forward_metadata_draft_extend( req_pool_indices=req_pool_indices, ) ) + # DP-padded rows carry the graph seq-len fill value (1); expanding a + # 1-length row over num_tokens_per_req query tokens yields 0 (or + # negative) per-token lens for its leading tokens. Keep the >=1 floor + # kernels assume (same convention as *_topk_lengths_clamp1) -- real + # rows are unaffected (their first token's KV length is >= 1). + seq_lens_casual = seq_lens_casual.clamp(min=1) core_attn_metadata = self.make_core_attn_metadata( req_to_token=self.req_to_token, req_pool_indices_repeated=req_pool_indices_repeated, @@ -1086,8 +1315,8 @@ def _fill_cuda_graph_swa_out_cache_loc( def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: # Upgrade Raw->Full so the c4/c128 compress + core_attn + indexer - # materialization is recorded inside the cuda graph; a no-op (Full - # already) when PREP_IN_CUDA_GRAPH=0. + # materialization is recorded inside the cuda graph; this is a no-op + # when the trtllm+MTP+DP exception prepared Full metadata on the host. if isinstance(self.forward_metadata, DSV4RawVerifyMetadata): self.forward_metadata = self.make_forward_metadata_from_raw_verify( raw_metadata=self.forward_metadata, @@ -1097,6 +1326,21 @@ def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: self.forward_metadata = self.make_forward_metadata_from_raw_decode( raw_metadata=self.forward_metadata, ) + # Tensors created while capturing a CUDA graph come from the memory + # pool shared by all captured graphs. If the last reference is + # dropped after capture, the pool hands the block to the next + # graph's capture, and two graphs' recorded kernels then address the + # same memory -- replaying one bucket's graph can overwrite metadata + # another bucket's graph reads on its next replay (observed victim: + # page_table, which also feeds the indexer, turning the corruption + # into silently wrong top-k KV selection). Holding one reference per + # captured graph prevents the pool from ever reusing these blocks. + # Under the trtllm+MTP+DP host-side exception, the metadata reaching capture is + # already Full, but tensors created below (swa_out_cache_loc) still + # come from the capture pool and are reference-replaced on the pinned + # object by copy_'s assign_fields on every replay-prep. + if torch.cuda.is_current_stream_capturing(): + self._captured_full_metadata_refs.append(self.forward_metadata) # Compute the SWA KV-store write target once per forward and cache it on # the metadata for every layer's store. This is recorded inside the cuda @@ -1127,6 +1371,17 @@ def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: torch.int32 ) ) + if torch.cuda.is_current_stream_capturing(): + # The recorded translate/store ops address THIS tensor at + # every replay, but copy_'s assign_fields replace the field + # on the (pinned) captured metadata each replay-prep, + # dropping the only python reference. Pin the tensor itself + # so the capture pool can never release its block (a released + # block unmaps -> replay MMU-faults at the baked address; + # a recycled block silently corrupts). + self._captured_full_metadata_refs.append( + metadata.core_attn_metadata.swa_out_cache_loc + ) if self.is_dspark_draft and forward_batch.forward_mode.is_target_verify(): block_size = int(forward_batch.spec_info.draft_token_num) @@ -1541,8 +1796,36 @@ def replay_cuda_graph_metadata_from( self.forward_metadata = temp_metadata return chosen_metadata.copy_(temp_metadata) + # temp_metadata is dropped after the content copy above, and it is + # the only owner of the compress planner's pinned staging buffer. + # That buffer was filled on the CPU and pushed with a raw + # cudaMemcpyAsync (no CachingHostAllocator event, see + # CompressorPrefillPlan.generate), so freeing it here lets the block + # be recycled and rewritten while the DMA is still in flight -- + # poisoned plan entries under overlap-scheduler CPU run-ahead. Retain + # the pinned blocks for a few iterations to outlive any run-ahead. + self._keep_plan_staging_alive(temp_metadata) self.forward_metadata = chosen_metadata + def _keep_plan_staging_alive( + self, + metadata: Union[ + DSV4Metadata, + DSV4RawVerifyMetadata, + DSV4RawDecodeMetadata, + ], + ) -> None: + if not isinstance(metadata, DSV4Metadata): + return + core = metadata.core_attn_metadata + for fused in ( + getattr(core, "c4_compress_metadata", None), + getattr(core, "c128_compress_metadata", None), + ): + pin = getattr(getattr(fused, "plan", None), "pin_buffer", None) + if pin is not None and pin.numel() > 0: + self._plan_staging_keepalive.append(pin) + def get_cuda_graph_seq_len_fill_value(self): return 1 @@ -1556,8 +1839,8 @@ def on_after_cuda_graph_warmup(self): core.c4_flashmla_metadata = _create_flashmla_metadata() core.c128_flashmla_metadata = _create_flashmla_metadata() - # PREP_IN_CUDA_GRAPH=True: warmup upgraded raw->full on the host; - # restore raw so capture re-runs the upgrade inside the graph. + # If warmup upgraded raw->full on the host, restore raw so capture + # re-runs the upgrade inside the graph. current_raw = getattr(self, "_current_capture_raw", None) if current_raw is not None: self.forward_metadata = current_raw @@ -1669,6 +1952,22 @@ def match_num_queries(x, value): extra_indices = match_num_queries(extra_indices, value=-1) extra_topk_lengths = match_num_queries(extra_topk_lengths, value=1) + if self.trtllm_attn: + # The uniform-FP8 pool is only readable by the trtllm-gen + # kernel; the subclass owns the whole dispatch + # (deepseek_v4_trtllm_backend.DeepseekV4TrtllmAttnBackend). + return self._forward_trtllm( + q=q, + layer=layer, + compress_ratio=compress_ratio, + core_attn_metadata=core_attn_metadata, + forward_batch=forward_batch, + attn_sink=attn_sink, + swa_page_indices=swa_page_indices, + extra_indices=extra_indices, + extra_topk_lengths=extra_topk_lengths, + ) + if q.ndim == 3: q = q.unsqueeze(1) if swa_page_indices.ndim == 2: @@ -2223,6 +2522,8 @@ def make_core_attn_metadata( if need_compress: core_attn_metadata.init_compression_metadata(num_tokens) core_attn_metadata.init_flashmla_related(is_prefill=is_prefill) + if self.trtllm_attn: + core_attn_metadata.init_trtllm_sparse_buffers() else: core_attn_metadata.c4_sparse_topk_lengths = None core_attn_metadata.c4_sparse_page_indices = None @@ -2230,6 +2531,9 @@ def make_core_attn_metadata( core_attn_metadata.c1_flashmla_metadata = _create_flashmla_metadata() core_attn_metadata.c4_flashmla_metadata = None core_attn_metadata.c128_flashmla_metadata = None + if self.trtllm_attn: + # SWA-only capacity (draft-extend metadata skips compression). + core_attn_metadata.init_trtllm_sparse_buffers() return core_attn_metadata def get_dspark_swa_page_indices( @@ -2272,14 +2576,19 @@ def __init__( self.speculative_num_steps = speculative_num_steps self.attn_backends: List[DeepseekV4AttnBackend] = [] for i in range(self.speculative_num_steps): - self.attn_backends.append( - DeepseekV4AttnBackend( - model_runner, - speculative_step_id=i, - topk=self.topk, - speculative_num_steps=self.speculative_num_steps, - ) - ) + self.attn_backends.append(self._make_step_backend(model_runner, i)) + + def _make_step_backend( + self, model_runner: ModelRunner, step_id: int + ) -> DeepseekV4AttnBackend: + # Overridden by DeepseekV4TrtllmMultiStepBackend + # (deepseek_v4_trtllm_backend.py) to build trtllm per-step backends. + return DeepseekV4AttnBackend( + model_runner, + speculative_step_id=step_id, + topk=self.topk, + speculative_num_steps=self.speculative_num_steps, + ) def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: for attn_backend in self.attn_backends: diff --git a/python/sglang/srt/layers/attention/deepseek_v4_trtllm_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_trtllm_backend.py new file mode 100644 index 000000000000..e70cc0b4c348 --- /dev/null +++ b/python/sglang/srt/layers/attention/deepseek_v4_trtllm_backend.py @@ -0,0 +1,568 @@ +"""DeepSeek V4 attention backend for the flashinfer trtllm-gen sparse MLA +kernel (``--dsv4-attn-backend trtllm``, SM100/SM103). + +Subclasses :class:`DeepseekV4AttnBackend` (FlashMLA) and overrides only the +kernel dispatch: decode / target-verify / draft-extend and varlen prefill go +through ``trtllm_batch_decode_sparse_mla_dsv4`` against the uniform 512-dim +FP8 KV pools. Metadata construction (including the trtllm combined sparse +tables, ``DSV4AttnMetadata.init_trtllm_sparse_buffers``) stays on the shared +metadata class so BCG capture/replay ``copy_``/``assign_fields`` semantics +are identical for both backends. +""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Literal, Optional, Tuple + +import torch + +from sglang.srt.environ import envs +from sglang.srt.layers.attention.deepseek_v4_backend import ( + SWA_WINDOW, + DeepseekV4AttnBackend, + DeepseekV4MultiStepBackend, +) +from sglang.srt.runtime_context import get_exec, get_parallel + +if TYPE_CHECKING: + from sglang.srt.layers.attention.deepseek_v4_backend import DSV4AttnMetadata + from sglang.srt.layers.radix_attention import RadixAttention + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + from sglang.srt.model_executor.model_runner import ModelRunner + +logger = logging.getLogger(__name__) + +# trtllm decode workspace: one zero-initialized int8 buffer, allocated once +# through the RuntimeContext persistent-buffer lifecycle and shared by all +# backend instances (same pattern as trtllm_mla_backend's +# "trtllm_mla_zero_workspace"). +_TRTLLM_GEN_WORKSPACE_SIZE_MB = 128 + + +def _get_trtllm_workspace_buffer(device: torch.device) -> torch.Tensor: + from sglang.srt.runtime_context import get_buffer + + return get_buffer( + "trtllm_dsv4_zero_workspace", + lambda: torch.zeros( + _TRTLLM_GEN_WORKSPACE_SIZE_MB * 1024 * 1024, + dtype=torch.int8, + device=device, + ), + ) + + +_trtllm_semaphore_installed = False + + +def _install_persistent_trtllm_semaphores() -> None: + """WAR for flashinfer's trtllm-gen multi-CTA KV counter (semaphore) + buffer handling. flashinfer sizes it batch_size (= requests) x heads and + allocates it per call, but the kernel indexes semaphores as + [batch, numCtasForAllHeads, maxNumCtasQ] with maxNumCtasQ == seqLenQ for + the DSv4 sparse kernels (trtllm-gen TmemCorr.h counterOffset) -- i.e. + under-allocated for any multi-token launch. Remove once flashinfer + accepts a caller-provided persistent buffer with tile-aware sizing.""" + global _trtllm_semaphore_installed + if _trtllm_semaphore_installed: + return + import flashinfer.mla._core as _fi_core + + _orig = _fi_core._get_trtllm_gen_multi_ctas_kv_counter_buffer + # Single persistent counter buffer, allocated once OUTSIDE any graph + # capture and shared by every launch -- the same design TRT-LLM uses + # (AttentionOp::mMultiBlockSemaphores). The kernel self-resets counters + # at the end of each launch and launches are stream-ordered, so sharing + # is safe. This removes both failure modes of per-call allocation: + # (a) under-sizing for VarSeq multi-token launches (kernel indexes + # counters per q-tile, TmemCorr.h counterOffset formula), sized here for + # 16384 query rows; (b) capture-pool lifetime bugs (a buffer allocated + # inside capture whose reference dies is recycled by later captures, + # letting replays scribble counters over other graphs' tensors). + state: dict = {} + + def _patched(batch_size, num_qo_heads, sm_count, device): + buf = state.get("buf") + if buf is None or buf.device != device: + assert not torch.cuda.is_current_stream_capturing(), ( + "persistent trtllm semaphore buffer must be created outside " + "graph capture (first call is expected during eager warmup)" + ) + buf = _orig(16384, num_qo_heads, sm_count, device) + state["buf"] = buf + return buf + + _fi_core._get_trtllm_gen_multi_ctas_kv_counter_buffer = _patched + _fi_core._trtllm_semaphore_state = state + _trtllm_semaphore_installed = True + logger.info( + "trtllm-gen multi-CTA semaphores: single persistent 16384-row buffer " + "shared across launches (flashinfer sizing WAR)." + ) + + +class DeepseekV4TrtllmAttnBackend(DeepseekV4AttnBackend): + """DSV4 attention through the trtllm-gen sparse MLA kernel.""" + + trtllm_attn: bool = True + + def __init__( + self, + model_runner: ModelRunner, + skip_prefill: bool = False, + speculative_step_id=0, + topk=0, + speculative_num_steps=0, + ): + _install_persistent_trtllm_semaphores() + super().__init__( + model_runner, + skip_prefill=skip_prefill, + speculative_step_id=speculative_step_id, + topk=topk, + speculative_num_steps=speculative_num_steps, + ) + assert ( + self.token_to_kv_pool.uniform_fp8 + ), "the trtllm backend requires the uniform-FP8 DSv4 KV pool." + assert not envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get(), ( + "--dsv4-attn-backend trtllm does not support " + "SGLANG_OPT_USE_ONLINE_COMPRESS yet." + ) + # The varlen prefill packs query tokens contiguously per request + # (cum_seq_lens_q); CP round-robin reindexing breaks that packing. + assert get_parallel().attn_cp_size == 1, ( + "--dsv4-attn-backend trtllm does not support " + "context parallelism (attn_cp_size > 1) yet." + ) + self.trtllm_workspace_buffer = _get_trtllm_workspace_buffer(self.device) + + def _compute_prep_in_cuda_graph(self, model_runner: ModelRunner) -> bool: + # trtllm + speculative + DP attention prepares metadata on the host: + # in-graph prep degrades draft acceptance under DP's padded/idle-rank + # batches. Otherwise prep metadata in-graph. + return not (self.mtp_enabled and model_runner.server_args.enable_dp_attention) + + def _forward_trtllm( + self, + *, + q: torch.Tensor, + layer: RadixAttention, + compress_ratio: Literal[0, 4, 128], + core_attn_metadata: DSV4AttnMetadata, + forward_batch: ForwardBatch, + attn_sink: torch.Tensor, + swa_page_indices: torch.Tensor, + extra_indices: Optional[torch.Tensor], + extra_topk_lengths: Optional[torch.Tensor], + ) -> torch.Tensor: + # The uniform-FP8 pool is only readable by the trtllm + # backend. Decode runs one row per request; target-verify and + # draft-extend run one row per query token against the + # per-token metadata rows (seq_lens_casual and the per-token + # index tables built by the caller), exactly like the flashmla path. + assert attn_sink is not None + if ( + forward_batch.forward_mode.is_decode_or_idle() + or forward_batch.forward_mode.is_target_verify() + or forward_batch.forward_mode.is_draft_extend_v2() + ): + return self._forward_trtllm_decode( + q=q, + layer=layer, + compress_ratio=compress_ratio, + core_attn_metadata=core_attn_metadata, + attn_sink=attn_sink, + swa_page_indices=swa_page_indices, + extra_indices=extra_indices, + extra_topk_lengths=extra_topk_lengths, + ) + assert forward_batch.forward_mode.is_extend_without_speculative(), ( + "uniform-FP8 pool cannot be read by the packed FlashMLA " + f"kernels; unsupported forward mode " + f"{forward_batch.forward_mode} under " + "--dsv4-attn-backend trtllm" + ) + return self._forward_trtllm_prefill( + q=q, + layer=layer, + compress_ratio=compress_ratio, + forward_batch=forward_batch, + attn_sink=attn_sink, + swa_page_indices=swa_page_indices, + extra_indices=extra_indices, + extra_topk_lengths=extra_topk_lengths, + ) + + def _get_trtllm_bmm_scales(self, layer: RadixAttention) -> Tuple[float, float]: + """Fixed (bmm1, bmm2) = (softmax_scale, 1.0) as host floats. + + The kv dequant scale must be 1.0 because the uniform-FP8 store + quantizes at 1.0. We use host floats because tensor scales corrupt split-KV + reduction on flashinfer <0.6.13, and these are per-layer constants. + """ + + assert layer.k_scale_float is None or layer.k_scale_float == 1.0, ( + "--dsv4-attn-backend trtllm stores KV with a " + "fixed per-tensor scale of 1.0; a non-unit checkpoint kv-cache " + f"scale (k_scale_float={layer.k_scale_float}) is not supported yet." + ) + return (self.softmax_scale, 1.0) + + def _trtllm_kv_cache_views( + self, layer_id: int, compress_ratio: Literal[0, 4, 128] + ) -> Tuple[torch.Tensor, torch.Tensor]: + """HND paged views ``[pages, 1, page_size, 512]`` (e4m3) of the + uniform-FP8 SWA and compressed-tier pools. + + The kernel requires a compressed cache tensor even for SWA-only + (``compress_ratio == 0``) layers; the SWA pool is passed there and + the compressed region stays fully masked via ``sparse_topk_lens``. + """ + + token_to_kv_pool = self.token_to_kv_pool + swa_buf = token_to_kv_pool.get_swa_key_buffer_radix(layer_id) + swa_page_size = token_to_kv_pool.swa_kv_pool.page_size + swa_kv_cache = swa_buf.view(swa_buf.shape[0], 1, swa_page_size, 512) + if compress_ratio == 0: + compressed_kv_cache = swa_kv_cache + else: + extra_buf = token_to_kv_pool.get_extra_key_buffer(layer_id) + extra_page_size = token_to_kv_pool.get_extra_key_page_size(layer_id) + compressed_kv_cache = extra_buf.view( + extra_buf.shape[0], 1, extra_page_size, 512 + ) + return swa_kv_cache, compressed_kv_cache + + def _forward_trtllm_decode( + self, + *, + q: torch.Tensor, + layer: RadixAttention, + compress_ratio: Literal[0, 4, 128], + core_attn_metadata: DSV4AttnMetadata, + attn_sink: torch.Tensor, + swa_page_indices: torch.Tensor, + extra_indices: Optional[torch.Tensor], + extra_topk_lengths: Optional[torch.Tensor], + ) -> torch.Tensor: + """Sparse MLA decode via ``trtllm_batch_decode_sparse_mla_dsv4``. + + The combined sparse table lives in preallocated metadata buffers. + """ + + from flashinfer.mla import trtllm_batch_decode_sparse_mla_dsv4 + + bs, num_heads, head_dim = q.shape + assert head_dim == 512 + + # Pre-pad-planned metadata: draft-extend plans its metadata BEFORE + # prepare_mlp_sync_batch DP-pads the batch (prepare_for_draft_extend + # marks the plan non-replannable, see #27091), so under DP MAX_LEN + # padding q can carry MORE rows than the metadata. The extra rows are + # padding whose outputs are discarded downstream; run the kernel on + # the metadata-covered rows only and zero-fill the tail (finite, so + # nothing NaN-propagates) -- same recipe as the padded-prefill path. + # FlashMLA tolerates this overhang implicitly; the trtllm tables are + # exact-rows, hence the explicit slice. + n_meta_rows = core_attn_metadata.seq_lens_casual.shape[0] + out_pad_tail = None + if n_meta_rows < bs: + out_pad_tail = torch.zeros( + (bs, num_heads, 512), dtype=torch.bfloat16, device=q.device + ) + q = q[:n_meta_rows] + swa_page_indices = swa_page_indices[:n_meta_rows] + if extra_indices is not None: + extra_indices = extra_indices[:n_meta_rows] + if extra_topk_lengths is not None: + extra_topk_lengths = extra_topk_lengths[:n_meta_rows] + bs = n_meta_rows + + # Per-ratio tables: layer-invariant content (SWA columns, the whole + # c128 table, the c0 lens) was written once per step by + # init_trtllm_sparse_buffers; only the c4 tail + lens (indexer + # top-k) are written here. + assert swa_page_indices.shape == (bs, SWA_WINDOW) + if compress_ratio == 0: + # swa_page_indices is itself a valid combined table (capacity + # 128, all-SWA); no fill needed. Use the METADATA's table (a + # [:n] view of a 64-row-aligned inert parent), not the + # match_num_queries-processed argument: the arg is an exact-row + # tensor and the VarSeq kernel reads rows to the 64-token tile + # boundary (tile-overrun guard, see init_trtllm_sparse_buffers). + sparse_indices = core_attn_metadata.swa_page_indices + sparse_topk_lens = core_attn_metadata.trtllm_swa_lens + elif compress_ratio == 128: + sparse_indices = core_attn_metadata.trtllm_c128_indices + sparse_topk_lens = core_attn_metadata.trtllm_c128_lens + else: + sparse_indices = core_attn_metadata.trtllm_c4_indices + sparse_topk_lens = core_attn_metadata.trtllm_c4_lens + assert sparse_indices is not None and sparse_topk_lens is not None, ( + "trtllm decode requires metadata built with " + "init_trtllm_sparse_buffers (decode-mode DSV4AttnMetadata)" + ) + if sparse_indices.shape[0] != bs: + # Metadata may be built against a padded batch; slice to the live + # rows (views -- the c4 fill below stays in place). + assert sparse_indices.shape[0] > bs, f"{sparse_indices.shape=} {bs=}" + sparse_indices = sparse_indices[:bs] + if sparse_topk_lens.shape[0] != bs: + assert sparse_topk_lens.shape[0] > bs, f"{sparse_topk_lens.shape=}" + sparse_topk_lens = sparse_topk_lens[:bs] + + if compress_ratio == 4: + assert extra_indices is not None and extra_topk_lengths is not None + width = extra_indices.shape[-1] + assert ( + SWA_WINDOW + width == sparse_indices.shape[1] + ), f"{width=} {sparse_indices.shape=}" + sparse_indices[:, SWA_WINDOW:].copy_(extra_indices) + # Lens include the fixed 128 SWA slots; SWA validity itself is + # derived from seq_lens inside the kernel. + sparse_topk_lens.copy_(extra_topk_lengths) + sparse_topk_lens.add_(SWA_WINDOW) + + swa_kv_cache, compressed_kv_cache = self._trtllm_kv_cache_views( + layer.layer_id, compress_ratio + ) + + # RoPE is already applied upstream; at per-tensor scale 1.0 the FP8 + # quantization is a plain e4m3 cast. + q_fp8 = q.to(torch.float8_e4m3fn).view(bs, 1, num_heads, 512) + + bmm1_scale, bmm2_scale = self._get_trtllm_bmm_scales(layer) + + seq_lens = core_attn_metadata.seq_lens_casual + if seq_lens.shape[0] != bs: + assert seq_lens.shape[0] > bs, f"{seq_lens.shape=} {bs=}" + seq_lens = seq_lens[:bs] + assert attn_sink.dtype == torch.float32 + assert self.trtllm_workspace_buffer is not None + + out = trtllm_batch_decode_sparse_mla_dsv4( + query=q_fp8, + swa_kv_cache=swa_kv_cache, + workspace_buffer=self.trtllm_workspace_buffer, + sparse_indices=sparse_indices, + compressed_kv_cache=compressed_kv_cache, + sparse_topk_lens=sparse_topk_lens, + seq_lens=seq_lens, + bmm1_scale=bmm1_scale, + bmm2_scale=bmm2_scale, + sinks=attn_sink, + kv_layout="HND", + ) + if out_pad_tail is not None: + out_pad_tail[:bs] = out.view(bs, num_heads, 512) + return out_pad_tail + return out.view(bs, num_heads, 512) + + def _forward_trtllm_prefill( + self, + *, + q: torch.Tensor, + layer: RadixAttention, + compress_ratio: Literal[0, 4, 128], + forward_batch: ForwardBatch, + attn_sink: torch.Tensor, + swa_page_indices: torch.Tensor, + extra_indices: Optional[torch.Tensor], + extra_topk_lengths: Optional[torch.Tensor], + ) -> torch.Tensor: + """Sparse MLA varlen prefill: the decode kernel driven with + multi-token queries (``cum_seq_lens_q``/``max_q_len``). + + The sparse table has one row per query token: the token's own causal + SWA window in columns ``[0:128)`` and its compressed tier after. + ``seq_lens`` must be the per-request TOTAL KV length including any + cached prefix (chunked prefill / cache-hit extends); the kernel + derives each token's causal SWA validity from it, so no masks are + built here. Runs eagerly, so per-call allocations are fine. + """ + + from flashinfer.mla import trtllm_batch_decode_sparse_mla_dsv4 + + assert q.ndim == 3, f"{q.shape=}" + num_qo_padded, num_heads, head_dim = q.shape + assert head_dim == 512 + + # Varlen query structure, from the same host-side extend lens that + # produced this metadata (init_forward_metadata_prefill / + # expand_prefill_casually). + core = self.forward_metadata.core_attn_metadata + if core.trtllm_prefill_qmeta is None: + extend_seq_lens_cpu = forward_batch.extend_seq_lens_cpu + assert extend_seq_lens_cpu is not None and len(extend_seq_lens_cpu) > 0 + batch_size = len(extend_seq_lens_cpu) + cum_lens = [0] * (batch_size + 1) + for i, extend_len in enumerate(extend_seq_lens_cpu): + cum_lens[i + 1] = cum_lens[i] + int(extend_len) + sum_q = cum_lens[-1] + max_q_len = max(int(x) for x in extend_seq_lens_cpu) + # q (and the per-token metadata rows, via match_num_queries) may + # be padded past the real extend tokens; the pad rows sit at the + # end. + assert 0 < sum_q <= num_qo_padded, f"{sum_q=} {num_qo_padded=}" + # Per-request TOTAL KV length (cached prefix + extend tokens). + seq_lens_i32 = forward_batch.seq_lens.to(torch.int32) + assert seq_lens_i32.shape == (batch_size,), f"{seq_lens_i32.shape=}" + core.trtllm_prefill_qmeta = ( + self._move_to_device(cum_lens), + max_q_len, + sum_q, + seq_lens_i32, + ) + cum_seq_lens_q, max_q_len, sum_q, seq_lens = core.trtllm_prefill_qmeta + + # Combined per-token sparse table (physical indices, -1 invalid). + # Layer-invariant parts are cached per chunk on the metadata: the c0 + # lens, the c4 table's SWA half (per-layer: the indexer top-k tail + + # lens), and the whole c128 table + lens (its page list is + # metadata-level). + # TILE OVERRUN GUARD (see init_trtllm_sparse_buffers): the VarSeq + # kernel reads per-token table rows up to the 64-token tile boundary, + # so allocate every per-token tensor as a 64-row-aligned inert parent + # and hand the kernel [:sum_q] views. + sum_q_pad = (sum_q + 63) // 64 * 64 + + def _tile_padded_pf(fill, src=None, width=None): + shape = (sum_q_pad,) if width is None else (sum_q_pad, width) + buf = torch.full(shape, fill, **self.cuda_int32_kwargs) + if src is not None: + buf[:sum_q].copy_(src) + return buf[:sum_q] + + swa_indices = _tile_padded_pf(-1, swa_page_indices[:sum_q], width=SWA_WINDOW) + assert swa_indices.shape == (sum_q, SWA_WINDOW), f"{swa_indices.shape=}" + if extra_indices is None: + # SWA-only (compress_ratio == 0) layer. SWA_WINDOW satisfies the + # kernel's capacity constraints (>= 128, % 4 == 0). + sparse_indices = swa_indices + if core.trtllm_prefill_swa_lens is None: + core.trtllm_prefill_swa_lens = _tile_padded_pf(SWA_WINDOW) + sparse_topk_lens = core.trtllm_prefill_swa_lens + elif compress_ratio == 128: + if core.trtllm_prefill_c128 is None: + width = extra_indices.shape[-1] + assert width % 4 == 0, f"{width=}" + table = _tile_padded_pf(-1, width=SWA_WINDOW + width) + table[:, :SWA_WINDOW].copy_(swa_indices) + table[:, SWA_WINDOW:].copy_(extra_indices[:sum_q]) + assert extra_topk_lengths is not None + lens = _tile_padded_pf( + SWA_WINDOW, + extra_topk_lengths[:sum_q].to(torch.int32) + SWA_WINDOW, + ) + core.trtllm_prefill_c128 = (table, lens) + sparse_indices, sparse_topk_lens = core.trtllm_prefill_c128 + else: + assert extra_topk_lengths is not None + width = extra_indices.shape[-1] + # c4 index tables are padded to multiples of 64 upstream + # (_pad_last_dim), so the combined capacity satisfies % 4 == 0. + assert width % 4 == 0, f"{width=}" + if core.trtllm_prefill_c4_indices is None: + core.trtllm_prefill_c4_indices = _tile_padded_pf( + -1, width=SWA_WINDOW + width + ) + core.trtllm_prefill_c4_indices[:, :SWA_WINDOW].copy_(swa_indices) + sparse_indices = core.trtllm_prefill_c4_indices + assert sparse_indices.shape == ( + sum_q, + SWA_WINDOW + width, + ), f"{sparse_indices.shape=} {width=}" + sparse_indices[:, SWA_WINDOW:].copy_(extra_indices[:sum_q]) + # Total lens include the fixed 128 SWA slots; SWA validity itself + # is derived from seq_lens/cum_seq_lens_q inside the kernel. + sparse_topk_lens = _tile_padded_pf( + SWA_WINDOW, + extra_topk_lengths[:sum_q].to(torch.int32) + SWA_WINDOW, + ) + + # FP8 query: RoPE already applied upstream; per-tensor scale 1.0 makes + # quantization a plain e4m3 cast (same recipe as the decode branch). + q_fp8 = q[:sum_q].to(torch.float8_e4m3fn) + + swa_kv_cache, compressed_kv_cache = self._trtllm_kv_cache_views( + layer.layer_id, compress_ratio + ) + bmm1_scale, bmm2_scale = self._get_trtllm_bmm_scales(layer) + assert attn_sink.dtype == torch.float32 + assert self.trtllm_workspace_buffer is not None + + out_padded = None + out_arg = None + if num_qo_padded != sum_q: + # Padded prefill: run the kernel over the real tokens only and + # zero the pad rows (their outputs are discarded downstream, but + # keep them finite so nothing NaN-propagates). + out_padded = torch.zeros( + (num_qo_padded, num_heads, 512), + dtype=torch.bfloat16, + device=q.device, + ) + out_arg = out_padded[:sum_q] + + out = trtllm_batch_decode_sparse_mla_dsv4( + query=q_fp8, + swa_kv_cache=swa_kv_cache, + workspace_buffer=self.trtllm_workspace_buffer, + sparse_indices=sparse_indices, + compressed_kv_cache=compressed_kv_cache, + sparse_topk_lens=sparse_topk_lens, + seq_lens=seq_lens, + out=out_arg, + bmm1_scale=bmm1_scale, + bmm2_scale=bmm2_scale, + sinks=attn_sink, + kv_layout="HND", + cum_seq_lens_q=cum_seq_lens_q, + max_q_len=max_q_len, + ) + return out_padded if out_padded is not None else out + + +class DeepseekV4TrtllmMultiStepBackend( + DeepseekV4MultiStepBackend, DeepseekV4TrtllmAttnBackend +): + """Multi-step draft wrapper whose per-step backends are trtllm.""" + + def _make_step_backend( + self, model_runner: ModelRunner, step_id: int + ) -> DeepseekV4AttnBackend: + return DeepseekV4TrtllmAttnBackend( + model_runner, + speculative_step_id=step_id, + topk=self.topk, + speculative_num_steps=self.speculative_num_steps, + ) + + +def is_dsv4_trtllm_attn_enabled() -> bool: + return get_exec().kernel.dsv4_attn_backend == "trtllm" + + +def create_deepseek_v4_attn_backend( + model_runner: ModelRunner, **kwargs +) -> DeepseekV4AttnBackend: + """Construct the DSV4 backend matching --dsv4-attn-backend.""" + cls = ( + DeepseekV4TrtllmAttnBackend + if is_dsv4_trtllm_attn_enabled() + else DeepseekV4AttnBackend + ) + return cls(model_runner, **kwargs) + + +def create_deepseek_v4_multistep_backend( + model_runner: ModelRunner, topk: int, speculative_num_steps: int +) -> DeepseekV4MultiStepBackend: + cls = ( + DeepseekV4TrtllmMultiStepBackend + if is_dsv4_trtllm_attn_enabled() + else DeepseekV4MultiStepBackend + ) + return cls(model_runner, topk=topk, speculative_num_steps=speculative_num_steps) diff --git a/python/sglang/srt/layers/attention/dsv4/compressor_trtllm.py b/python/sglang/srt/layers/attention/dsv4/compressor_trtllm.py new file mode 100644 index 000000000000..26d2c95e3b11 --- /dev/null +++ b/python/sglang/srt/layers/attention/dsv4/compressor_trtllm.py @@ -0,0 +1,140 @@ +"""Uniform-FP8 (trtllm backend) DSV4 compressor store path. + +The fused JIT compressor epilogue writes the packed FlashMLA cache layout +(448-dim FP8 NoPE + 64-dim BF16 RoPE + block scales); the trtllm backend's +uniform 512-dim FP8 pool needs a different epilogue. This module carries a +standalone unfused pipeline for that pool -- compress, invalid-row masking, +norm + RoPE, then a plain e4m3 cast store -- so the shared +``compressor_v2.py`` (FlashMLA / HIP) stays untouched. + +The pipeline intentionally duplicates the compress/norm/RoPE steps of +``CompressorBackendMixin._forward_unified_hip`` rather than refactoring +them out of the shared file. A follow-up fused uniform-FP8 store (see PR +#32975) replaces this module wholesale. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from sglang.kernels.ops.attention.dsv4 import compress_forward + +if TYPE_CHECKING: + from sglang.srt.layers.attention.dsv4.compressor import Compressor + from sglang.srt.layers.attention.dsv4.compressor_v2 import CompressorBackendMixin + from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool + + +def _mask_invalid_prefill_compress_rows( + kv_compressed: torch.Tensor, + plan_raw: torch.Tensor, + out_loc: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """Make padded prefill-plan rows inert before the unfused store. + + ``PlanC::invalid()`` uses ``seq_len == -1`` and a ``ragged_id`` of + ``0xffff``. The fused epilogue checks ``is_invalid()`` before reading + ``out_loc``; the unfused uniform-FP8 path must provide the same guard + without introducing a dynamic-shape mask (it is also used by BCG). + """ + valid = plan_raw[:, 0] != -1 + kv_compressed = torch.where( + valid.unsqueeze(-1), kv_compressed, torch.zeros_like(kv_compressed) + ) + + ragged_ids = plan_raw[:, 1].to(torch.int32) & 0xFFFF + safe_ragged_ids = torch.where(valid, ragged_ids, torch.zeros_like(ragged_ids)) + mapped_out_loc = out_loc[safe_ragged_ids.long()] + # Slot 0 is the allocator's reserved padding sink, so duplicate invalid + # writes cannot collide with a live cache entry. + out_loc_to_store = torch.where( + valid, mapped_out_loc, torch.zeros_like(mapped_out_loc) + ) + return kv_compressed, out_loc_to_store + + +def forward_compress_uniform_fp8( + backend: CompressorBackendMixin, + *, + token_to_kv_pool: DeepSeekV4TokenToKVPool, + kv_score_input: torch.Tensor, + state_pool, + compressor: Compressor, + layer_id: int, +) -> None: + """Unfused compress + norm + RoPE + e4m3 store for the uniform-FP8 pool. + + The compression math is the same JIT kernel as the fused path; only the + epilogue differs (plain FP8 cast into the 512-dim uniform layout). + """ + from sglang.kernels.ops.attention.deepseek_v4_rope import ( + fused_norm_rope_inplace_triton, + ) + from sglang.srt.layers.attention.dsv4.compressor_v2 import ( + _extract_positions_from_plan, + _use_online_compress, + is_overlap_compress, + ) + + assert not compressor.is_in_indexer + assert compressor.head_dim == 512, f"{compressor.head_dim=}" + assert not _use_online_compress(compressor.ratio), ( + "SGLANG_OPT_USE_ONLINE_COMPRESS is not supported with the " + "uniform-FP8 KV layout yet." + ) + + compress_ratio = compressor.ratio + head_dim = compressor.head_dim + + plan = backend._get_paged_compress_metadata(compress_ratio) + out_loc = backend._get_out_loc(compress_ratio) + + coff = 2 if is_overlap_compress(compress_ratio) else 1 + kv_score_buffer = state_pool.kv_score_buffer.kv_score.view( + -1, compress_ratio, 2 * head_dim * coff + ) + + kv_compressed = compress_forward( + kv_score_buffer=kv_score_buffer, + kv_score_input=kv_score_input, + ape=compressor.ape.view(-1, head_dim), + plan=plan, + compress_ratio=compress_ratio, + head_dim=head_dim, + is_online=False, + ) + if kv_compressed.shape[0] == 0: + return + + plan_raw = plan[1].view(torch.int32) + if plan.is_decode: + # Zero out non-boundary tokens to prevent corrupting kvcache loc 0. + seq_lens_plan = plan_raw[:, 0].to(torch.int32) + is_boundary = (seq_lens_plan % compress_ratio == 0).unsqueeze(-1) + kv_compressed = torch.where( + is_boundary, kv_compressed, torch.zeros_like(kv_compressed) + ) + out_loc_to_store = out_loc + else: + kv_compressed, out_loc_to_store = _mask_invalid_prefill_compress_rows( + kv_compressed, + plan_raw, + out_loc, + ) + + positions = _extract_positions_from_plan(plan, compress_ratio).clamp(min=0) + fused_norm_rope_inplace_triton( + kv_compressed, + compressor.norm.weight, + compressor.norm.variance_epsilon, + compressor.freqs_cis, + positions=positions, + ) + + token_to_kv_pool.set_extra_key_buffer_fused( + layer_id=layer_id, + loc=out_loc_to_store, + cache_k=kv_compressed, + ) diff --git a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py index b9268e8c6eeb..ebaf4c037373 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py @@ -218,45 +218,65 @@ def forward_unified( is_unified_kv_triton, ) - out_loc = self._get_out_loc(compressor.ratio) - use_fp4_indexer = ( - compressor.is_in_indexer and self.enable_deepseek_v4_fp4_indexer - ) - bf16_store = False - if compressor.is_in_indexer: - kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer(layer_id) - page_size = token_to_kv_pool.get_index_k_page_size() - elif is_unified_kv_triton(): - kv_cache = token_to_kv_pool.get_unified_kv(layer_id) - page_size = 1 - out_loc = getattr( - self.forward_metadata.core_metadata.unified, - f"c{compressor.ratio}_out_loc", + if token_to_kv_pool.uniform_fp8 and not compressor.is_in_indexer: + # The fused epilogue only writes the packed layout; the trtllm + # backend's uniform-FP8 pool stores through its standalone + # pipeline (compressor_trtllm.py). The indexer compressor keeps + # its own (blockwise-FP8) path below. + from sglang.srt.layers.attention.dsv4.compressor_trtllm import ( + forward_compress_uniform_fp8, + ) + + forward_compress_uniform_fp8( + self, + token_to_kv_pool=token_to_kv_pool, + kv_score_input=kv_score_input, + state_pool=state_pool, + compressor=compressor, + layer_id=layer_id, ) - bf16_store = True else: - _, _, compress_kv_pool = token_to_kv_pool.layer_mapping[layer_id] - assert compress_kv_pool is not None - kv_cache = token_to_kv_pool.get_extra_key_buffer(layer_id) - page_size = token_to_kv_pool.get_extra_key_page_size(layer_id) - if hasattr(compress_kv_pool, "translate_loc_to_hisparse_device"): - out_loc = compress_kv_pool._translate_loc_to_hisparse_device(out_loc) - self._forward_compress_all_in_one( - kv_score_buffer=state_pool.kv_score_buffer.kv_score, - kv_score_input=kv_score_input, - ape=compressor.ape, - head_dim=compressor.head_dim, - norm=compressor.norm, - freqs_cis_cache=compressor.freqs_cis, - kv_cache=kv_cache.view(dtype=torch.uint8), - is_indexer=compressor.is_in_indexer, - rotate=compressor.rotate, - compress_ratio=compressor.ratio, - page_size=page_size, - out_loc=out_loc, - use_fp4_indexer=use_fp4_indexer, - bf16_store=bf16_store, - ) + out_loc = self._get_out_loc(compressor.ratio) + use_fp4_indexer = ( + compressor.is_in_indexer and self.enable_deepseek_v4_fp4_indexer + ) + bf16_store = False + if compressor.is_in_indexer: + kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer(layer_id) + page_size = token_to_kv_pool.get_index_k_page_size() + elif is_unified_kv_triton(): + kv_cache = token_to_kv_pool.get_unified_kv(layer_id) + page_size = 1 + out_loc = getattr( + self.forward_metadata.core_metadata.unified, + f"c{compressor.ratio}_out_loc", + ) + bf16_store = True + else: + _, _, compress_kv_pool = token_to_kv_pool.layer_mapping[layer_id] + assert compress_kv_pool is not None + kv_cache = token_to_kv_pool.get_extra_key_buffer(layer_id) + page_size = token_to_kv_pool.get_extra_key_page_size(layer_id) + if hasattr(compress_kv_pool, "translate_loc_to_hisparse_device"): + out_loc = compress_kv_pool._translate_loc_to_hisparse_device( + out_loc + ) + self._forward_compress_all_in_one( + kv_score_buffer=state_pool.kv_score_buffer.kv_score, + kv_score_input=kv_score_input, + ape=compressor.ape, + head_dim=compressor.head_dim, + norm=compressor.norm, + freqs_cis_cache=compressor.freqs_cis, + kv_cache=kv_cache.view(dtype=torch.uint8), + is_indexer=compressor.is_in_indexer, + rotate=compressor.rotate, + compress_ratio=compressor.ratio, + page_size=page_size, + out_loc=out_loc, + use_fp4_indexer=use_fp4_indexer, + bf16_store=bf16_store, + ) online_c128_mtp = getattr(self, "online_c128_mtp", None) if online_c128_mtp is not None: online_c128_mtp.write_prefix_states( diff --git a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py index 54eabcb7e878..264b9ff88490 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -174,6 +174,67 @@ def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]: raise NotImplementedError("Use get_key_buffer instead.") +class DeepSeekV4UniformFP8KVPool(DeepSeekV4SingleKVPool): + """Uniform 512-dim FP8 (e4m3) variant of the DSv4 single-KV pool. + + Layout required by the trtllm-gen sparse MLA kernel: every token is 512 + contiguous e4m3 values (448 nope + 64 rope, both FP8), with no in-cache + scales and no per-page padding -- the kernel addresses the cache at token + granularity as ``flat_index * 512`` bytes. The per-tensor dequant scale + is delivered externally via the kernel's bmm scales. + """ + + def get_bytes_per_token(self) -> int: + return self.qk_nope_head_dim + self.qk_rope_head_dim + + def create_buffer(self, *, num_pages: int): + bytes_per_token = self.get_bytes_per_token() + assert bytes_per_token == 512, ( + "DSV4 uniform-FP8 KV layout: qk_nope_head_dim (448) + " + "qk_rope_head_dim (64), all e4m3 = 512 bytes/token" + ) + self.kv_cache_total_dim = bytes_per_token + self.bytes_per_page_padded = self.page_size * bytes_per_token + + return torch.zeros( + num_pages, + self.page_size * bytes_per_token, + dtype=torch.float8_e4m3fn, + device=self.device, + ) + + def get_key_buffer(self, layer_id: int): + return self.kv_buffer[layer_id - self.start_layer] + + def set_key_buffer( + self, + layer_id: int, + loc: torch.Tensor, + cache_nope_fp8_rope_bf16_pack: NopeFp8RopeBf16Pack, + ): + raise NotImplementedError( + "The packed NopeFp8RopeBf16Pack store does not apply to the " + "uniform-FP8 pool; use set_key_buffer_fused." + ) + + def set_key_buffer_fused( + self, + layer_id: int, + loc: torch.Tensor, + cache_k: torch.Tensor, + ) -> None: + """Store already-normed/roped rows as plain e4m3 with a fixed scale of + 1.0 (the backend's bmm scales assume this - change both together). + + Uses uint8 views to work around index_put not supporting fp8 dtypes. + """ + + assert cache_k.dim() == 2 and cache_k.shape[1] == self.kv_cache_total_dim + self.kv_buffer[layer_id].view(torch.uint8).view(-1, self.kv_cache_total_dim)[ + loc.long() + ] = cache_k.to(torch.float8_e4m3fn).view(torch.uint8) + + class HiSparseC4DevicePool(DeepSeekV4SingleKVPool): def __init__( @@ -578,6 +639,10 @@ def __init__( ) self._unified_kv = is_unified_kv_triton() + # Uniform 512-dim e4m3 layout for the trtllm attention backend + self.uniform_fp8 = ( + not self._unified_kv + ) and get_exec().kernel.dsv4_attn_backend == "trtllm" if self._unified_kv: self.swa_kv_pool = None @@ -606,6 +671,13 @@ def __init__( self.unified_swa_pages = self.unified_kv_pool.swa_pages else: self.unified_kv_pool = None + kv_pool_cls: type = DeepSeekV4SingleKVPool + if self.uniform_fp8: + assert dtype == torch.float8_e4m3fn, ( + "--dsv4-attn-backend trtllm requires " + f"kv_cache_dtype=fp8_e4m3, got {dtype}" + ) + kv_pool_cls = DeepSeekV4UniformFP8KVPool self.swa_kv_pool = self._make_kv_pool( size=swa_size, page_size=swa_page_size, @@ -614,10 +686,15 @@ def __init__( device=device, enable_memory_saver=enable_memory_saver, global_page_size=swa_page_size, + cls=kv_pool_cls, ) - c4_kv_pool_type = DeepSeekV4SingleKVPool + c4_kv_pool_type = kv_pool_cls if enable_hisparse: + assert not self.uniform_fp8, ( + "enable_hisparse is not supported with " + "--dsv4-attn-backend trtllm." + ) c4_kv_pool_type = HiSparseC4DevicePool self.c4_kv_pool = self._make_kv_pool( size=c4_size, @@ -638,6 +715,7 @@ def __init__( device=device, enable_memory_saver=enable_memory_saver, global_page_size=page_size, + cls=kv_pool_cls, ) indexer_size = self.c4_logical_size @@ -1187,6 +1265,26 @@ def set_swa_key_buffer_radix_fused_norm_rope( freqs_cis: torch.Tensor, positions: torch.Tensor, ) -> None: + if self.uniform_fp8: + # Uniform-FP8 (trtllm-gen) layout: norm + RoPE with the existing + # Triton kernel (in-place on kv; safe -- kv is not read again), + # then a plain e4m3 cast + scatter in the pool setter (per-tensor + # scale 1.0). Fusing the store is deferred to the perf phase. + from sglang.kernels.ops.attention.deepseek_v4_rope import ( + fused_norm_rope_inplace_triton, + ) + + fused_norm_rope_inplace_triton( + kv, + kv_weight, + eps, + freqs_cis, + positions=positions, + ) + self.swa_kv_pool.set_key_buffer_fused( + self._swa_local_layer_id(layer_id), swa_loc, kv + ) + return fused_k_norm_rope_flashmla( kv=kv, kv_weight=kv_weight, diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index e7c8d4c0540e..8c70d0b49f73 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -359,6 +359,8 @@ def resolve_encoder_transfer_backend( ] NSA_CHOICES = DSA_CHOICES # deprecated alias +DSV4_ATTN_BACKEND_CHOICES = ["auto", "flashmla", "trtllm"] + DSV4_PREFILL_BACKEND_CHOICES = [ "auto", "flashmla_sparse", @@ -1848,6 +1850,20 @@ class ServerArgs: ), NS("exec.kernel"), ] = "sgl-kernel" + dsv4_attn_backend: A[ + str, + Arg( + help="DeepSeek V4 attention backend. 'auto' (default) resolves to " + "'flashmla'. 'trtllm' (opt-in, SM100/SM103 with FP8 KV cache) " + "switches the SWA/compressed KV pools to a " + "uniform 512-dim FP8 layout and runs decode and sparse prefill " + "through the flashinfer trtllm-gen sparse MLA kernel. The backend " + "choice is shared by prefill and decode.", + choices=DSV4_ATTN_BACKEND_CHOICES, + resolvable=True, + ), + NS("exec.kernel"), + ] = "auto" disable_flashinfer_autotune: A[ bool, "Disable FlashInfer autotuning.", NS("exec.kernel") ] = False diff --git a/python/sglang/srt/speculative/draft_utils.py b/python/sglang/srt/speculative/draft_utils.py index 12bebcc1ce5b..1adc6a54d200 100644 --- a/python/sglang/srt/speculative/draft_utils.py +++ b/python/sglang/srt/speculative/draft_utils.py @@ -381,8 +381,8 @@ def _create_dsv4_decode_backend(self): DeepseekV4MultiStepBackend, ) else: - from sglang.srt.layers.attention.deepseek_v4_backend import ( - DeepseekV4MultiStepBackend, + from sglang.srt.layers.attention.deepseek_v4_trtllm_backend import ( + create_deepseek_v4_multistep_backend as DeepseekV4MultiStepBackend, ) return ( @@ -521,11 +521,13 @@ def _create_dsv4_prefill_backend(self): "dsv4", DeepseekV4HipRadixBackend(self.draft_model_runner, skip_prefill=False), ) - from sglang.srt.layers.attention.deepseek_v4_backend import ( - DeepseekV4AttnBackend, + from sglang.srt.layers.attention.deepseek_v4_trtllm_backend import ( + create_deepseek_v4_attn_backend, ) return ( "dsv4", - DeepseekV4AttnBackend(self.draft_model_runner, skip_prefill=False), + create_deepseek_v4_attn_backend( + self.draft_model_runner, skip_prefill=False + ), ) diff --git a/test/registered/backends/test_dsv4_fp8_trtllm_backend.py b/test/registered/backends/test_dsv4_fp8_trtllm_backend.py new file mode 100644 index 000000000000..a60ea9dc3e9d --- /dev/null +++ b/test/registered/backends/test_dsv4_fp8_trtllm_backend.py @@ -0,0 +1,305 @@ +"""DeepSeek-V4 uniform-FP8 trtllm backend (SM100/SM103 Blackwell only). + +Validates --dsv4-attn-backend trtllm — DSv4 decode and sparse +varlen prefill through flashinfer's ``trtllm_batch_decode_sparse_mla_dsv4`` +on a uniform 512-dim FP8-e4m3 KV cache with an FP8 query — against the +default packed-FP8 FlashMLA path: + +1. A short greedy-output baseline is collected from a FlashMLA server, then + the same prompts are replayed on a trtllm server and compared (output-level). +2. Long multi-k-token prompts do the same comparison for the varlen prefill + path (mixed lengths fired concurrently to exercise cum_seq_lens_q + packing; a repeat run exercises the radix-cache-hit / cached-prefix + extend, and the longest prompt exceeds --chunked-prefill-size so chunked + prefill is exercised too). +3. Decode-correctness probes + a GSM8K sanity eval run on the trtllm + server. +4. A CUDA-graph capture/replay smoke: concurrent decode batches of varying + size (replaying different captured decode-graph buckets) must stay + consistent with a single-request greedy run. +""" + +import concurrent.futures +import difflib +import unittest +from types import SimpleNamespace + +import requests +import torch + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.basic_decode_correctness_kit import BasicDecodeCorrectnessMixin +from sglang.test.run_eval import run_eval +from sglang.test.test_utils import ( + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=1200, stage="base-c", runner_config="4-gpu-b200") + +DSV4_FLASH_MODEL_PATH = "deepseek-ai/DeepSeek-V4-Flash" +SERVER_LAUNCH_TIMEOUT = 3600 +DSV4_BASE_ENV = { + "SGLANG_JIT_DEEPGEMM_FAST_WARMUP": "1", +} + +SERVER_ARGS = [ + "--trust-remote-code", + "--tp", + "4", + "--max-running-requests", + "32", + "--mem-fraction-static", + "0.85", + "--chunked-prefill-size", + "4096", + # V4-Flash shis MXFP4 routed experts, and the auto-selected Triton MoE runner + # cannot consume the packed layout. Matches the B200 Flash cookbook recipe. + "--moe-runner-backend", + "flashinfer_mxfp4", + "--disable-flashinfer-autotune", +] + +COMPARE_PROMPTS = [ + "The capital of France is", + "In one sentence, explain why the sky is blue.", + "List the first five prime numbers:", + "Water boils at", + "The author of Romeo and Juliet is", + "Translate 'good morning' to French:", + "2 + 2 * 3 =", + "Photosynthesis is the process by which", +] +COMPARE_MAX_NEW_TOKENS = 64 +# FP8 formats differ (packed per-block scales + BF16 rope vs uniform +# per-tensor e4m3), so greedy outputs may diverge after some tokens; require +# strong average prefix similarity rather than exact equality. +COMPARE_MIN_MEAN_SIMILARITY = 0.6 + +# Long multi-k-token prompts that exercise the trtllm-gen varlen +# prefill: real c4 indexer top-k selection needs >~2k tokens of context and +# the c128 far tier needs whole 128-token pages; the mixed lengths also +# exercise cum_seq_lens_q packing when fired concurrently, and the longest +# exceeds --chunked-prefill-size (4096) so it prefills in multiple chunks. +_FILLER_SENTENCES = [ + "The expedition recorded water temperature, salinity, and current speed " + "at every station along the transect. ", + "Archival records from the observatory describe decades of nightly " + "measurements taken with remarkable consistency. ", + "Each greenhouse module recycles condensate through a gravel bed before " + "returning it to the irrigation loop. ", + "The survey team catalogued the masonry of the aqueduct arch by arch, " + "noting repairs from three distinct centuries. ", +] +_LONG_PROMPT_QUESTION = ( + "\n\nIn one short sentence, what kind of activity do the paragraphs " + "above describe?" +) + + +def _make_long_prompt(idx: int, target_chars: int) -> str: + sentence = _FILLER_SENTENCES[idx % len(_FILLER_SENTENCES)] + body = "" + n = 0 + while len(body) < target_chars: + body += f"[Entry {idx}-{n}] " + sentence + n += 1 + return body + _LONG_PROMPT_QUESTION + + +# ~4 chars/token: roughly 2.5k, 4.5k, and 7k tokens. +LONG_PROMPTS = [ + _make_long_prompt(0, 10_000), + _make_long_prompt(1, 18_000), + _make_long_prompt(2, 28_000), +] +LONG_MAX_NEW_TOKENS = 32 +LONG_MIN_MEAN_SIMILARITY = 0.6 + +GSM8K_NUM_EXAMPLES = 200 +GSM8K_MIN_SCORE = 0.90 + +_REQUEST_TIMEOUT = 600 + + +def _is_sm100() -> bool: + if not torch.cuda.is_available(): + return False + return torch.cuda.get_device_capability() in ((10, 0), (10, 3)) + + +def _launch(backend: str): + return popen_launch_server( + DSV4_FLASH_MODEL_PATH, + DEFAULT_URL_FOR_TEST, + timeout=SERVER_LAUNCH_TIMEOUT, + other_args=SERVER_ARGS + ["--dsv4-attn-backend", backend], + env=dict(DSV4_BASE_ENV), + ) + + +def _greedy_generate(base_url: str, prompt: str, max_new_tokens: int) -> str: + resp = requests.post( + base_url + "/generate", + json={ + "text": prompt, + "sampling_params": { + "temperature": 0.0, + "max_new_tokens": max_new_tokens, + }, + }, + timeout=_REQUEST_TIMEOUT, + ) + resp.raise_for_status() + return resp.json()["text"] + + +class TestDSV4Fp8TrtllmBackend(BasicDecodeCorrectnessMixin, CustomTestCase): + """TP4 DSv4-Flash-FP8 with --dsv4-attn-backend trtllm.""" + + @classmethod + def setUpClass(cls): + if not _is_sm100(): + raise unittest.SkipTest( + "DSv4 trtllm uniform-FP8 attention requires SM100/SM103 (Blackwell)" + ) + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = None + + # Collect the packed-FP8 FlashMLA greedy baselines (short + # decode-focused prompts + long prefill-focused prompts). + baseline_process = _launch("flashmla") + try: + cls.flashmla_outputs = [ + _greedy_generate(cls.base_url, p, COMPARE_MAX_NEW_TOKENS) + for p in COMPARE_PROMPTS + ] + cls.flashmla_long_outputs = [ + _greedy_generate(cls.base_url, p, LONG_MAX_NEW_TOKENS) + for p in LONG_PROMPTS + ] + finally: + kill_process_tree(baseline_process.pid) + + # The server under test (kept alive for all test methods). + cls.process = _launch("trtllm") + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process is not None: + kill_process_tree(cls.process.pid) + + def test_greedy_outputs_match_flashmla(self): + """Output-level comparison vs the packed-FP8 FlashMLA path.""" + similarities = [] + for prompt, ref_out in zip(COMPARE_PROMPTS, self.flashmla_outputs): + out = _greedy_generate(self.base_url, prompt, COMPARE_MAX_NEW_TOKENS) + sim = difflib.SequenceMatcher(None, ref_out, out).ratio() + similarities.append(sim) + print( + f"[compare] sim={sim:.3f} prompt={prompt!r}\n" + f" flashmla : {ref_out!r}\n" + f" trtllm : {out!r}" + ) + mean_sim = sum(similarities) / len(similarities) + self.assertGreater( + mean_sim, + COMPARE_MIN_MEAN_SIMILARITY, + f"trtllm greedy outputs diverge from flashmla: " + f"mean similarity {mean_sim:.3f}, per-prompt {similarities}", + ) + + def test_long_prompt_prefill_matches_flashmla(self): + """Varlen trtllm-gen prefill vs FlashMLA on multi-k-token prompts. + + The three prompts are fired concurrently (mixed extend lengths in + one batch exercise cum_seq_lens_q packing and per-token sparse-table + construction), then the longest is re-sent alone (radix-cache hit → + cached-prefix extend, where seq_lens > extend len). The longest + prompt also exceeds --chunked-prefill-size, covering chunked + prefill. + """ + + with concurrent.futures.ThreadPoolExecutor(len(LONG_PROMPTS)) as pool: + outs = list( + pool.map( + lambda p: _greedy_generate(self.base_url, p, LONG_MAX_NEW_TOKENS), + LONG_PROMPTS, + ) + ) + + similarities = [] + for i, (ref_out, out) in enumerate(zip(self.flashmla_long_outputs, outs)): + sim = difflib.SequenceMatcher(None, ref_out, out).ratio() + similarities.append(sim) + print( + f"[long-compare] sim={sim:.3f} prompt_chars={len(LONG_PROMPTS[i])}\n" + f" flashmla : {ref_out!r}\n" + f" trtllm : {out!r}" + ) + mean_sim = sum(similarities) / len(similarities) + self.assertGreater( + mean_sim, + LONG_MIN_MEAN_SIMILARITY, + f"trtllm long-prompt (varlen prefill) outputs diverge from " + f"flashmla: mean similarity {mean_sim:.3f}, per-prompt {similarities}", + ) + + # Cached-prefix extend: re-run the longest prompt; the radix cache + # holds its prefix, so this prefill extends from cached tokens + # (seq_lens total > extend tokens). Greedy output must be unchanged. + rerun = _greedy_generate(self.base_url, LONG_PROMPTS[-1], LONG_MAX_NEW_TOKENS) + self.assertEqual( + outs[-1], + rerun, + "greedy long-prompt output changed on the cached-prefix " + "(radix-cache hit) extend path", + ) + + def test_gsm8k_sanity(self): + args = SimpleNamespace( + base_url=self.base_url, + model=DSV4_FLASH_MODEL_PATH, + eval_name="gsm8k", + api="completion", + max_tokens=512, + num_examples=GSM8K_NUM_EXAMPLES, + num_threads=64, + ) + metrics = run_eval(args) + print(f"GSM8K sanity on trtllm decode: {metrics=}") + self.assertGreater(metrics["score"], GSM8K_MIN_SCORE) + + def test_cuda_graph_capture_replay_smoke(self): + """Exercise decode CUDA-graph replay across batch-size buckets. + + Bursts of concurrent requests at varying concurrency replay different + captured decode-graph buckets; a repeated greedy single must remain + identical to its first run afterwards (address-stable trtllm sparse + buffers, no capture/replay corruption). + """ + anchor_prompt = "Q: What is the capital of France?\nA:" + anchor_out = _greedy_generate(self.base_url, anchor_prompt, 32) + + for concurrency in (2, 4, 8, 16): + prompts = [f"Count from {i} to {i + 5}: " for i in range(concurrency)] + with concurrent.futures.ThreadPoolExecutor(concurrency) as pool: + outs = list( + pool.map(lambda p: _greedy_generate(self.base_url, p, 32), prompts) + ) + self.assertEqual(len(outs), concurrency) + for out in outs: + self.assertGreater(len(out), 0) + + anchor_out_replayed = _greedy_generate(self.base_url, anchor_prompt, 32) + self.assertEqual( + anchor_out, + anchor_out_replayed, + "greedy output changed after batched decode-graph replays", + ) + + +if __name__ == "__main__": + unittest.main(verbosity=3) diff --git a/test/registered/disaggregation/test_disaggregation_dsv4.py b/test/registered/disaggregation/test_disaggregation_dsv4.py index e930eac41f82..97b9287c53d8 100644 --- a/test/registered/disaggregation/test_disaggregation_dsv4.py +++ b/test/registered/disaggregation/test_disaggregation_dsv4.py @@ -75,6 +75,8 @@ def start_prefill(cls): DEEPEP_CONFIG, "--cuda-graph-max-bs-decode", "128", + "--mem-fraction-static", + "0.9", "--max-running-requests", "128", *_EAGLE_SPEC_ARGS, @@ -113,6 +115,8 @@ def start_decode(cls): DEEPEP_CONFIG, "--cuda-graph-max-bs-decode", "128", + "--mem-fraction-static", + "0.9", "--max-running-requests", "128", *_EAGLE_SPEC_ARGS, diff --git a/test/registered/models_e2e/test_deepseek_v4_flash_fp4_b200_trtllm.py b/test/registered/models_e2e/test_deepseek_v4_flash_fp4_b200_trtllm.py new file mode 100644 index 000000000000..7dca7e0897c7 --- /dev/null +++ b/test/registered/models_e2e/test_deepseek_v4_flash_fp4_b200_trtllm.py @@ -0,0 +1,236 @@ +"""B200 per-commit CI: DeepSeek-V4-Flash FP4 with the trtllm attention backend. + +Same four recipes as test_deepseek_v4_flash_fp4_b200.py (which guards the +default FlashMLA backend), with ``--dsv4-attn-backend trtllm`` (uniform-FP8 +KV pool, trtllm-gen sparse MLA for decode and prefill). + +Registry: base-c-test-4-gpu-b200 (per-commit, 4x B200) +""" + +import unittest + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.basic_decode_correctness_kit import BasicDecodeCorrectnessMixin +from sglang.test.kits.eval_accuracy_kit import GSM8KMixin +from sglang.test.kits.spec_decoding_kit import SpecDecodingMixin +from sglang.test.test_utils import ( + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, + try_cached_model, +) + +register_cuda_ci(est_time=700, stage="base-c", runner_config="4-gpu-b200") + +MODEL = "deepseek-ai/DeepSeek-V4-Flash" +SERVER_LAUNCH_TIMEOUT = 3600 +DEEPEP_CONFIG = '{"normal_dispatch":{"num_sms":96},"normal_combine":{"num_sms":96}}' + +_DEEPEP_ENV = { + "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "1024", +} + + +class TestDSV4FlashFP4B200Trtllm( + SpecDecodingMixin, + BasicDecodeCorrectnessMixin, + GSM8KMixin, + CustomTestCase, +): + """LowLatency recipe: TP=4, FP4 (mxfp4), EAGLE spec decoding.""" + + gsm8k_accuracy_thres = 0.93 + accept_length_thres = 2.8 + bs_1_speed_thres = 220 + + @classmethod + def setUpClass(cls): + cls.model = try_cached_model(MODEL) + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=SERVER_LAUNCH_TIMEOUT, + other_args=[ + "--trust-remote-code", + "--dsv4-attn-backend", + "trtllm", + "--tp", + "4", + "--moe-runner-backend", + "flashinfer_mxfp4", + "--speculative-algorithm", + "EAGLE", + "--speculative-num-steps", + "3", + "--speculative-eagle-topk", + "1", + "--speculative-num-draft-tokens", + "4", + "--chunked-prefill-size", + "4096", + "--disable-flashinfer-autotune", + ], + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + + +class TestDSV4FlashFP4B200BalancedTrtllm( + SpecDecodingMixin, + BasicDecodeCorrectnessMixin, + GSM8KMixin, + CustomTestCase, +): + """Balanced recipe: TP=4, DP=4, DeepEP, EAGLE (1-step spec).""" + + gsm8k_accuracy_thres = 0.93 + accept_length_thres = 1.8 + bs_1_speed_thres = 100 + + @classmethod + def setUpClass(cls): + cls.model = try_cached_model(MODEL) + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=SERVER_LAUNCH_TIMEOUT, + other_args=[ + "--trust-remote-code", + "--dsv4-attn-backend", + "trtllm", + "--tp", + "4", + "--dp", + "4", + "--enable-dp-attention", + "--moe-a2a-backend", + "deepep", + "--speculative-algorithm", + "EAGLE", + "--speculative-num-steps", + "1", + "--speculative-eagle-topk", + "1", + "--speculative-num-draft-tokens", + "2", + "--deepep-config", + DEEPEP_CONFIG, + ], + env=_DEEPEP_ENV, + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + + +class TestDSV4FlashFP4NonMTPB200Trtllm( + BasicDecodeCorrectnessMixin, GSM8KMixin, CustomTestCase +): + """Non-MTP recipe: TP=4, DP=4, DeepEP, no speculative decoding.""" + + gsm8k_accuracy_thres = 0.93 + + @classmethod + def setUpClass(cls): + cls.model = try_cached_model(MODEL) + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=SERVER_LAUNCH_TIMEOUT, + other_args=[ + "--trust-remote-code", + "--dsv4-attn-backend", + "trtllm", + "--tp", + "4", + "--dp", + "4", + "--enable-dp-attention", + "--moe-a2a-backend", + "deepep", + "--deepep-config", + DEEPEP_CONFIG, + ], + env=_DEEPEP_ENV, + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + + +@unittest.skip( + "DP prefill BCG with an idle rank fabricates a dummy EXTEND whose tokens " + "are counted as real; their hidden states enter the shared EP grouped " + "GEMMs and perturb live ranks' logits (#31125). Under FlashMLA this " + "shows as nondeterministic outputs; under the trtllm backend the " + "perturbation reliably drives generations empty, so every " + "low-concurrency probe here fails. " + "Re-enable once the generic DP idle-rank fix lands (follow-up PR to " + "#30805)." +) +class TestDSV4FlashFP4BreakableCudaGraphB200Trtllm( + BasicDecodeCorrectnessMixin, GSM8KMixin, CustomTestCase +): + """BCG recipe: TP=4, DP=4, DeepEP, DP attention, mixed chunk.""" + + gsm8k_accuracy_thres = 0.93 + + @classmethod + def setUpClass(cls): + cls.model = try_cached_model(MODEL) + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=SERVER_LAUNCH_TIMEOUT, + other_args=[ + "--trust-remote-code", + "--dsv4-attn-backend", + "trtllm", + "--tp", + "4", + "--dp", + "4", + "--enable-dp-attention", + "--enable-mixed-chunk", + "--cuda-graph-backend-prefill", + "breakable", + "--moe-a2a-backend", + "deepep", + "--deepep-config", + DEEPEP_CONFIG, + "--chunked-prefill-size", + "4096", + "--piecewise-cuda-graph-max-tokens", + "1024", + "--mem-fraction-static", + "0.80", + "--cuda-graph-max-bs-decode", + "16", + "--max-running-requests", + "128", + "--watchdog-timeout", + "900", + ], + env=_DEEPEP_ENV, + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py index 34ef3c0c887a..cd2bbc6edeb8 100644 --- a/test/registered/unit/test_model_overrides.py +++ b/test/registered/unit/test_model_overrides.py @@ -86,6 +86,7 @@ def test_server_args_whitelist_is_exactly_the_migrated_fields(self): "kv_cache_dtype", "dsa_prefill_backend", "dsa_decode_backend", + "dsv4_attn_backend", "prefill_attention_backend", "decode_attention_backend", "flashinfer_allreduce_fusion_backend", From faa489dfc0e964a1efa7dca21eaff036f1115938 Mon Sep 17 00:00:00 2001 From: Akhil Goel Date: Fri, 21 Aug 2026 15:56:24 -0700 Subject: [PATCH 2/3] Add uniform-FP8 store variant to the CUDA FusedNormRopeKernel --- .../csrc/deepseek_v4/fused_norm_rope_v2.cuh | 46 +++++++++++++++---- .../kernels/ops/attention/dsv4/compress.py | 14 +++++- 2 files changed, 49 insertions(+), 11 deletions(-) diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh index a3b41157566f..20f664be1ff0 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh @@ -380,10 +380,17 @@ INDEXER_KERNEL void fused_norm_rope_indexer_fp4(const __grid_constant__ FusedNor // Each thread loads kVecSize=2 BF16, so 256 threads cover the full 512 elems. // Cache layout: 584 bytes/token = 448 fp8 nope + 64 (=32 bf16x2) rope + 8 scale. // ---------------------------------------------------------------------------- -template +template < + typename DType, + ForwardMode kMode, + int32_t kPageBits, + bool kUsePDL, + bool kBf16Store = false, + bool kUniformFp8Store = false> FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormRopeStoreParams params) { using namespace device; using enum ForwardMode; + static_assert(!(kBf16Store && kUniformFp8Store)); constexpr int64_t kHeadDim = 512; constexpr int64_t kRopeDim = 64; @@ -393,8 +400,11 @@ FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormR constexpr uint32_t kRopeWarp = kNumWarps - 1; // kBf16Store: write the whole head_dim as plain BF16 (no fp8 / no scale) into a // [num_slots, head_dim] bf16 cache (page_size==1) at row out_loc - constexpr int64_t kPageBytes = - kBf16Store ? ((kHeadDim * 2ll) << kPageBits) : host::div_ceil(584ll << kPageBits, 576) * 576; + // kUniformFp8Store: write the whole head_dim (rope tail included) as plain + // e4m3 at per-tensor scale 1.0 into the uniform 512-byte-per-token pool. + constexpr int64_t kPageBytes = kBf16Store ? ((kHeadDim * 2ll) << kPageBits) + : kUniformFp8Store ? (kHeadDim << kPageBits) + : host::div_ceil(584ll << kPageBits, 576) * 576; static_assert(kHeadDim == kBlockSize * kVecSize); static_assert(kRopeDim == kWarpThreads * kVecSize); static_assert(kHeadDim - kRopeDim == kRopeWarp * kWarpThreads * kVecSize); @@ -465,12 +475,14 @@ FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormR const int64_t page = out_loc >> kPageBits; const int64_t offset = out_loc & ((1 << kPageBits) - 1); const auto page_ptr = params.kvcache + page * kPageBytes; - const auto value_ptr = page_ptr + offset * (kBf16Store ? (kHeadDim * 2) : 576); + const auto value_ptr = + page_ptr + offset * (kBf16Store ? (kHeadDim * 2) : kUniformFp8Store ? kHeadDim : 576); PDLTriggerSecondary(); - // part 2: rope on the rope warp (BF16 store), or per-warp FP8 quant + store. - if constexpr (kBf16Store) { + // part 2: rope on the rope warp (BF16/uniform store), or per-warp FP8 + // quant + store (packed layout). + if constexpr (kBf16Store || kUniformFp8Store) { Float2 d = data; if (warp_id == kRopeWarp) { const auto x_real = data[0]; @@ -480,7 +492,15 @@ FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormR d[0] = x_real * freq_real - x_imag * freq_imag; d[1] = x_real * freq_imag + x_imag * freq_real; } - reinterpret_cast(value_ptr)[tx] = cast(fp32x2_t{d[0], d[1]}); + if constexpr (kUniformFp8Store) { + // BF16 round-trip to match the unfused path (Triton norm+rope emits + // bf16, then the pool store casts bf16 -> e4m3 at scale 1.0). + const auto x = cast(cast(d[0])); + const auto y = cast(cast(d[1])); + reinterpret_cast(value_ptr)[tx] = pack_fp8(x, y); + } else { + reinterpret_cast(value_ptr)[tx] = cast(fp32x2_t{d[0], d[1]}); + } } else if (warp_id == kRopeWarp) { // Each rope-warp lane owns exactly one (real, imag) pair within the rope // tail. Apply rotation, downcast to BF16, write to the slot's rope region. @@ -518,15 +538,21 @@ template < uint32_t kPageSize, bool kUsePDL, int32_t kPreshuffleSize = 0, - bool kBf16Store = false> + bool kBf16Store = false, + bool kUniformFp8Store = false> struct FusedNormRopeKernel { static constexpr int32_t kLogPageSize = std::countr_zero(kPageSize); static constexpr bool kIsIndexer = (kHeadDim == 128); static_assert(!(kIsIndexer && kBf16Store), "bf16 store only for flashmla head_dim=512"); + static_assert(!(kIsIndexer && kUniformFp8Store), "uniform fp8 store only for flashmla head_dim=512"); + static_assert(!(kBf16Store && kUniformFp8Store)); static constexpr int64_t kIndexerBytes = 132 * kPageSize; static constexpr int64_t kFlashMLABytes = host::div_ceil(584 * kPageSize, 576) * 576; static constexpr int64_t kBf16Bytes = kHeadDim * 2 * kPageSize; // plain bf16 cache - static constexpr int64_t kPageBytes = kBf16Store ? kBf16Bytes : (kIsIndexer ? kIndexerBytes : kFlashMLABytes); + static constexpr int64_t kUniformBytes = kHeadDim * kPageSize; // uniform e4m3, 512 B/token + static constexpr int64_t kPageBytes = kBf16Store ? kBf16Bytes + : kUniformFp8Store ? kUniformBytes + : (kIsIndexer ? kIndexerBytes : kFlashMLABytes); /// TODO: Let's fix the config for now. static_assert(kRopeDim == 64 && (kHeadDim == 128 || kHeadDim == 512)); @@ -537,7 +563,7 @@ struct FusedNormRopeKernel { if constexpr (kIsIndexer) { return fused_norm_rope_indexer; } else { - return fused_norm_rope_flashmla; + return fused_norm_rope_flashmla; } } diff --git a/python/sglang/kernels/ops/attention/dsv4/compress.py b/python/sglang/kernels/ops/attention/dsv4/compress.py index bb3540076a53..511dc67f92b9 100644 --- a/python/sglang/kernels/ops/attention/dsv4/compress.py +++ b/python/sglang/kernels/ops/attention/dsv4/compress.py @@ -49,6 +49,7 @@ def _jit_compress_norm_rope_module( rope_dim: int, page_size: int, bf16_store: bool = False, + uniform_fp8_store: bool = False, ) -> Module: args = make_cpp_args( dtype, @@ -58,6 +59,7 @@ def _jit_compress_norm_rope_module( is_arch_support_pdl(), INDEXER_K_CACHE_PRESHUFFLE_TILE if aiter_can_use_preshuffle_paged_mqa() else 0, bf16_store, + uniform_fp8_store, ) cuda_wrappers = [("forward", f"FusedNormRopeKernel<{args}>::forward")] if head_dim == 128: @@ -425,9 +427,14 @@ def compress_norm_rope_store( page_size: int, use_fp4: bool = False, bf16_store: bool = False, + uniform_fp8_store: bool = False, ) -> None: if use_fp4: assert kv.shape[-1] == 128 + if uniform_fp8_store: + # Uniform 512-byte-per-token e4m3 pool (trtllm backend): plain cast at + # per-tensor scale 1.0, rope tail included; no packed scales. + assert kv.shape[-1] == 512 and not use_fp4 and not bf16_store freq_cis = torch.view_as_real(freq_cis).flatten(-2) if _is_xpu: compress_norm_rope_store_xpu( @@ -445,7 +452,12 @@ def compress_norm_rope_store( ) else: module = _jit_compress_norm_rope_module( - kv.dtype, kv.shape[-1], freq_cis.shape[-1], page_size, bf16_store + kv.dtype, + kv.shape[-1], + freq_cis.shape[-1], + page_size, + bf16_store, + uniform_fp8_store, ) fn = module.forward_fp4 if use_fp4 else module.forward fn( From 1ee6592e48ae572e0a5cbccf861ad9c3968a988e Mon Sep 17 00:00:00 2001 From: Akhil Goel Date: Fri, 21 Aug 2026 17:14:02 -0700 Subject: [PATCH 3/3] Use the fused CUDA store for the uniform-FP8 compressor path --- .../csrc/deepseek_v4/fused_norm_rope_v2.cuh | 9 +- .../attention/dsv4/compressor_trtllm.py | 100 +++--------- .../test_dsv4_uniform_fp8_fused_store.py | 144 ++++++++++++++++++ 3 files changed, 171 insertions(+), 82 deletions(-) create mode 100644 test/registered/kernels/test_dsv4_uniform_fp8_fused_store.py diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh index 20f664be1ff0..33ea58a522d0 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh @@ -402,7 +402,7 @@ FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormR // [num_slots, head_dim] bf16 cache (page_size==1) at row out_loc // kUniformFp8Store: write the whole head_dim (rope tail included) as plain // e4m3 at per-tensor scale 1.0 into the uniform 512-byte-per-token pool. - constexpr int64_t kPageBytes = kBf16Store ? ((kHeadDim * 2ll) << kPageBits) + constexpr int64_t kPageBytes = kBf16Store ? ((kHeadDim * 2ll) << kPageBits) : kUniformFp8Store ? (kHeadDim << kPageBits) : host::div_ceil(584ll << kPageBits, 576) * 576; static_assert(kHeadDim == kBlockSize * kVecSize); @@ -475,8 +475,7 @@ FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormR const int64_t page = out_loc >> kPageBits; const int64_t offset = out_loc & ((1 << kPageBits) - 1); const auto page_ptr = params.kvcache + page * kPageBytes; - const auto value_ptr = - page_ptr + offset * (kBf16Store ? (kHeadDim * 2) : kUniformFp8Store ? kHeadDim : 576); + const auto value_ptr = page_ptr + offset * (kBf16Store ? (kHeadDim * 2) : kUniformFp8Store ? kHeadDim : 576); PDLTriggerSecondary(); @@ -549,8 +548,8 @@ struct FusedNormRopeKernel { static constexpr int64_t kIndexerBytes = 132 * kPageSize; static constexpr int64_t kFlashMLABytes = host::div_ceil(584 * kPageSize, 576) * 576; static constexpr int64_t kBf16Bytes = kHeadDim * 2 * kPageSize; // plain bf16 cache - static constexpr int64_t kUniformBytes = kHeadDim * kPageSize; // uniform e4m3, 512 B/token - static constexpr int64_t kPageBytes = kBf16Store ? kBf16Bytes + static constexpr int64_t kUniformBytes = kHeadDim * kPageSize; // uniform e4m3, 512 B/token + static constexpr int64_t kPageBytes = kBf16Store ? kBf16Bytes : kUniformFp8Store ? kUniformBytes : (kIsIndexer ? kIndexerBytes : kFlashMLABytes); diff --git a/python/sglang/srt/layers/attention/dsv4/compressor_trtllm.py b/python/sglang/srt/layers/attention/dsv4/compressor_trtllm.py index 26d2c95e3b11..a42b84b44808 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor_trtllm.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor_trtllm.py @@ -1,16 +1,12 @@ """Uniform-FP8 (trtllm backend) DSV4 compressor store path. -The fused JIT compressor epilogue writes the packed FlashMLA cache layout -(448-dim FP8 NoPE + 64-dim BF16 RoPE + block scales); the trtllm backend's -uniform 512-dim FP8 pool needs a different epilogue. This module carries a -standalone unfused pipeline for that pool -- compress, invalid-row masking, -norm + RoPE, then a plain e4m3 cast store -- so the shared -``compressor_v2.py`` (FlashMLA / HIP) stays untouched. - -The pipeline intentionally duplicates the compress/norm/RoPE steps of -``CompressorBackendMixin._forward_unified_hip`` rather than refactoring -them out of the shared file. A follow-up fused uniform-FP8 store (see PR -#32975) replaces this module wholesale. +Mirrors the packed FlashMLA pipeline: ``compress_forward`` (softmax pooling) +followed by the fused ``compress_norm_rope_store`` CUDA kernel with the +uniform-FP8 epilogue -- RMSNorm + RoPE + plain e4m3 cast (per-tensor scale +1.0) + paged scatter into the 512-byte-per-token uniform pool, one launch. +The kernel reads positions, destination slots, decode window boundaries and +prefill-row validity from the compress plan, so no Python-side masking or +position extraction is needed (same contract as the packed epilogue). """ from __future__ import annotations @@ -19,7 +15,7 @@ import torch -from sglang.kernels.ops.attention.dsv4 import compress_forward +from sglang.kernels.ops.attention.dsv4 import compress_forward, compress_norm_rope_store if TYPE_CHECKING: from sglang.srt.layers.attention.dsv4.compressor import Compressor @@ -27,34 +23,6 @@ from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool -def _mask_invalid_prefill_compress_rows( - kv_compressed: torch.Tensor, - plan_raw: torch.Tensor, - out_loc: torch.Tensor, -) -> tuple[torch.Tensor, torch.Tensor]: - """Make padded prefill-plan rows inert before the unfused store. - - ``PlanC::invalid()`` uses ``seq_len == -1`` and a ``ragged_id`` of - ``0xffff``. The fused epilogue checks ``is_invalid()`` before reading - ``out_loc``; the unfused uniform-FP8 path must provide the same guard - without introducing a dynamic-shape mask (it is also used by BCG). - """ - valid = plan_raw[:, 0] != -1 - kv_compressed = torch.where( - valid.unsqueeze(-1), kv_compressed, torch.zeros_like(kv_compressed) - ) - - ragged_ids = plan_raw[:, 1].to(torch.int32) & 0xFFFF - safe_ragged_ids = torch.where(valid, ragged_ids, torch.zeros_like(ragged_ids)) - mapped_out_loc = out_loc[safe_ragged_ids.long()] - # Slot 0 is the allocator's reserved padding sink, so duplicate invalid - # writes cannot collide with a live cache entry. - out_loc_to_store = torch.where( - valid, mapped_out_loc, torch.zeros_like(mapped_out_loc) - ) - return kv_compressed, out_loc_to_store - - def forward_compress_uniform_fp8( backend: CompressorBackendMixin, *, @@ -64,16 +32,8 @@ def forward_compress_uniform_fp8( compressor: Compressor, layer_id: int, ) -> None: - """Unfused compress + norm + RoPE + e4m3 store for the uniform-FP8 pool. - - The compression math is the same JIT kernel as the fused path; only the - epilogue differs (plain FP8 cast into the 512-dim uniform layout). - """ - from sglang.kernels.ops.attention.deepseek_v4_rope import ( - fused_norm_rope_inplace_triton, - ) + """Compress + fused norm/RoPE/e4m3 store for the uniform-FP8 pool.""" from sglang.srt.layers.attention.dsv4.compressor_v2 import ( - _extract_positions_from_plan, _use_online_compress, is_overlap_compress, ) @@ -108,33 +68,19 @@ def forward_compress_uniform_fp8( if kv_compressed.shape[0] == 0: return - plan_raw = plan[1].view(torch.int32) - if plan.is_decode: - # Zero out non-boundary tokens to prevent corrupting kvcache loc 0. - seq_lens_plan = plan_raw[:, 0].to(torch.int32) - is_boundary = (seq_lens_plan % compress_ratio == 0).unsqueeze(-1) - kv_compressed = torch.where( - is_boundary, kv_compressed, torch.zeros_like(kv_compressed) - ) - out_loc_to_store = out_loc - else: - kv_compressed, out_loc_to_store = _mask_invalid_prefill_compress_rows( - kv_compressed, - plan_raw, - out_loc, - ) - - positions = _extract_positions_from_plan(plan, compress_ratio).clamp(min=0) - fused_norm_rope_inplace_triton( + # One fused launch: RMSNorm + RoPE + e4m3 cast + paged scatter. The + # kernel skips invalid prefill plan rows and non-boundary decode tokens, + # and resolves each row's position / destination from the plan. + kv_cache = token_to_kv_pool.get_extra_key_buffer(layer_id) + page_size = token_to_kv_pool.get_extra_key_page_size(layer_id) + compress_norm_rope_store( kv_compressed, - compressor.norm.weight, - compressor.norm.variance_epsilon, - compressor.freqs_cis, - positions=positions, - ) - - token_to_kv_pool.set_extra_key_buffer_fused( - layer_id=layer_id, - loc=out_loc_to_store, - cache_k=kv_compressed, + plan, + norm_weight=compressor.norm.weight, + norm_eps=compressor.norm.variance_epsilon, + freq_cis=compressor.freqs_cis, + out_loc=out_loc, + kvcache=kv_cache.view(torch.uint8), + page_size=page_size, + uniform_fp8_store=True, ) diff --git a/test/registered/kernels/test_dsv4_uniform_fp8_fused_store.py b/test/registered/kernels/test_dsv4_uniform_fp8_fused_store.py new file mode 100644 index 000000000000..557215eba53b --- /dev/null +++ b/test/registered/kernels/test_dsv4_uniform_fp8_fused_store.py @@ -0,0 +1,144 @@ +"""Correctness of the uniform-FP8 fused store (CUDA FusedNormRopeKernel). + +The kernel fuses RMSNorm + RoPE + plain e4m3 cast (per-tensor scale 1.0) + +paged scatter into the 512-byte-per-token uniform pool, reading positions, +destinations, decode window boundaries and prefill-row validity from the +compress plan. Reference is the unfused pipeline it replaces +(fused_norm_rope_inplace_triton, then an e4m3 cast + index_put) with the +plan semantics emulated in Python. + +The CUDA block reduction sums the RMSNorm squares in a different order than +the Triton reference, so a ~1e-6 fraction of elements can land one e4m3 ulp +apart; the assertions allow exactly that and nothing more. +""" + +import pytest +import torch + +from sglang.kernels.ops.attention.deepseek_v4_rope import fused_norm_rope_inplace_triton +from sglang.kernels.ops.attention.dsv4.compress import ( + CompressorDecodePlan, + CompressorPrefillPlan, + compress_norm_rope_store, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") + +HEAD_DIM = 512 +ROPE_DIM = 64 +RATIO = 4 +PAGE_SIZE = 64 + + +def _make_freqs(max_pos: int) -> torch.Tensor: + inv_freq = 1.0 / ( + 10000.0 ** (torch.arange(0, ROPE_DIM, 2, device="cuda").float() / ROPE_DIM) + ) + t = torch.arange(max_pos, device="cuda").float() + angles = torch.outer(t, inv_freq) + return torch.polar(torch.ones_like(angles), angles) # complex64 [max_pos, 32] + + +def _reference_rows(kv, weight, eps, freqs, positions): + x = kv.clone() + fused_norm_rope_inplace_triton(x, weight, eps, freqs, positions=positions) + return x.to(torch.float8_e4m3fn) + + +def _assert_rows_match(got_fp8, ref_fp8): + got, ref = got_fp8.float(), ref_fp8.float() + mismatch = (got_fp8.view(torch.uint8) != ref_fp8.view(torch.uint8)).float().mean() + assert mismatch.item() <= 1e-4, f"{mismatch.item()=}" + # any mismatches must be within one e4m3 ulp (<= 2^-3 relative) + denom = ref.abs().clamp(min=2**-9) + assert ((got - ref).abs() / denom).max().item() <= 0.13 + + +@pytest.mark.parametrize("num_rows", [1, 7, 128, 2048]) +def test_uniform_fp8_store_decode_plan(num_rows): + torch.manual_seed(num_rows) + kv = torch.randn(num_rows, HEAD_DIM, dtype=torch.bfloat16, device="cuda") * 2 + weight = torch.randn(HEAD_DIM, dtype=torch.bfloat16, device="cuda").abs() + 0.5 + eps = 1e-6 + seq = torch.randint(1, 8192, (num_rows,), device="cuda", dtype=torch.int32) + # half the rows sit on a window boundary (stored), half not (skipped) + seq[::2] = (seq[::2] // RATIO).clamp(min=1) * RATIO + seq[1::2] = (seq[1::2] // RATIO).clamp(min=1) * RATIO + 1 + plan_i32 = torch.zeros(num_rows, 4, dtype=torch.int32, device="cuda") + plan_i32[:, 0] = seq + plan = CompressorDecodePlan(RATIO, plan_i32.view(torch.uint8)) + freqs = _make_freqs(8192 + RATIO) + out_loc = torch.randperm(num_rows * 2, device="cuda")[:num_rows].to(torch.int64) + + num_pages = (num_rows * 2 + PAGE_SIZE - 1) // PAGE_SIZE + 1 + cache = torch.zeros( + num_pages, PAGE_SIZE * HEAD_DIM, dtype=torch.uint8, device="cuda" + ) + compress_norm_rope_store( + kv, + plan, + norm_weight=weight, + norm_eps=eps, + freq_cis=freqs, + out_loc=out_loc, + kvcache=cache, + page_size=PAGE_SIZE, + uniform_fp8_store=True, + ) + + boundary = (seq % RATIO) == 0 + positions = (seq - RATIO).to(torch.int32) + ref = _reference_rows(kv, weight, eps, freqs, positions) + rows = cache.view(-1, HEAD_DIM) + got = rows[out_loc[boundary]].view(torch.float8_e4m3fn) + _assert_rows_match(got, ref[boundary]) + # non-boundary rows must not be written (their slots stay zero) + skipped = rows[out_loc[~boundary]] + assert (skipped == 0).all(), "non-boundary decode rows must be skipped" + + +@pytest.mark.parametrize("num_rows", [8, 300, 1024]) +def test_uniform_fp8_store_prefill_plan_with_invalid_rows(num_rows): + torch.manual_seed(num_rows) + kv = torch.randn(num_rows, HEAD_DIM, dtype=torch.bfloat16, device="cuda") * 2 + weight = torch.randn(HEAD_DIM, dtype=torch.bfloat16, device="cuda").abs() + 0.5 + eps = 1e-6 + seq = (torch.randint(1, 2048, (num_rows,), device="cuda") * RATIO).to(torch.int32) + ragged = torch.randperm(num_rows, device="cuda").to(torch.int32) + valid = torch.rand(num_rows, device="cuda") > 0.25 + plan_i32 = torch.zeros(num_rows, 4, dtype=torch.int32, device="cuda") + plan_i32[:, 0] = torch.where(valid, seq, torch.full_like(seq, -1)) + plan_i32[:, 1] = torch.where(valid, ragged, torch.full_like(ragged, 0xFFFF)) + plan_c = plan_i32.view(torch.uint8) + plan_w = torch.zeros(num_rows, 8, dtype=torch.uint8, device="cuda") + plan = CompressorPrefillPlan(RATIO, plan_c, plan_w) + freqs = _make_freqs(2048 * RATIO + RATIO) + out_loc = torch.randperm(num_rows * 2, device="cuda")[:num_rows].to(torch.int64) + + num_pages = (num_rows * 2 + PAGE_SIZE - 1) // PAGE_SIZE + 1 + cache = torch.zeros( + num_pages, PAGE_SIZE * HEAD_DIM, dtype=torch.uint8, device="cuda" + ) + compress_norm_rope_store( + kv, + plan, + norm_weight=weight, + norm_eps=eps, + freq_cis=freqs, + out_loc=out_loc, + kvcache=cache, + page_size=PAGE_SIZE, + uniform_fp8_store=True, + ) + + positions = (seq - RATIO).clamp(min=0).to(torch.int32) + ref = _reference_rows(kv, weight, eps, freqs, positions) + rows = cache.view(-1, HEAD_DIM) + got = rows[out_loc[ragged[valid].long()]].view(torch.float8_e4m3fn) + _assert_rows_match(got, ref[valid]) + # invalid rows are skipped entirely: every slot not mapped by a valid row + # stays zero + written = torch.zeros(rows.shape[0], dtype=torch.bool, device="cuda") + written[out_loc[ragged[valid].long()]] = True + assert (rows[~written] == 0).all(), "invalid prefill rows must be skipped"