diff --git a/python/sglang/jit_kernel/hadamard.py b/python/sglang/jit_kernel/hadamard.py index 25930ce942d3..6e845474903e 100644 --- a/python/sglang/jit_kernel/hadamard.py +++ b/python/sglang/jit_kernel/hadamard.py @@ -5,6 +5,7 @@ import torch from sglang.jit_kernel.utils import KERNEL_PATH, cache_once, load_jit, make_cpp_args +from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: from tvm_ffi.module import Module @@ -56,6 +57,14 @@ def _hadamard_transform_impl( return out.reshape(shapes_og) +def _hadamard_transform_fake_impl( + x: torch.Tensor, + scale: float = 1.0, +) -> torch.Tensor: + return torch.empty_like(x) + + +@register_custom_op(fake_impl=_hadamard_transform_fake_impl) def hadamard_transform(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor: module = _jit_hadamard_module(x.dtype) return _hadamard_transform_impl(x, scale, 8, module.hadamard_transform) diff --git a/python/sglang/srt/compilation/piecewise_context_manager.py b/python/sglang/srt/compilation/piecewise_context_manager.py index 20a08a9972b9..bbe0d040d705 100644 --- a/python/sglang/srt/compilation/piecewise_context_manager.py +++ b/python/sglang/srt/compilation/piecewise_context_manager.py @@ -71,6 +71,7 @@ def __init__(self): self.quant_config = None self.moe_layers = None self.moe_fusions = None + self.dsa_indexers = None def set_forward_batch(self, forward_batch: ForwardBatch): self.forward_batch = forward_batch @@ -87,6 +88,9 @@ def set_moe_layers(self, layers: List[Any]): def set_moe_fusions(self, fusions: List[Any]): self.moe_fusions = fusions + def set_dsa_indexers(self, indexers: List[Any]): + self.dsa_indexers = indexers + _forward_context: Optional[ForwardContext] = None @@ -104,6 +108,7 @@ def set_forward_context( quant_config: Any, moe_layers: List[Any], moe_fusions: List[Any], + dsa_indexers: Optional[List[Any]] = None, ): global _forward_context _forward_context = ForwardContext() @@ -112,6 +117,8 @@ def set_forward_context( _forward_context.set_quant_config(quant_config) _forward_context.set_moe_layers(moe_layers) _forward_context.set_moe_fusions(moe_fusions) + if dsa_indexers is not None: + _forward_context.set_dsa_indexers(dsa_indexers) try: yield finally: diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 6f89a69025d1..e3551392a753 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -329,7 +329,6 @@ def __init__( self.use_ngram_embedding = getattr(self.hf_config, "use_ngram_embedding", False) self.is_piecewise_cuda_graph_disabled_model = ( is_piecewise_cuda_graph_disabled_model(self.hf_config.architectures) - or is_deepseek_dsa(self.hf_text_config) ) self.dtype = _get_and_verify_dtype(self.hf_text_config, dtype) @@ -1556,11 +1555,9 @@ def is_generation_model(model_architectures: List[str], is_embedding: bool = Fal ] piecewise_cuda_graph_disabled_model_archs = [ - "DeepseekV32ForCausalLM", "DeepseekV4ForCausalLM", "DeepseekV4ForCausalLMNextN", "Qwen3NextForCausalLM", - "GlmMoeDsaForCausalLM", "BailingMoeV2_5ForCausalLM", "LLaDAModelLM", ] diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index 7548c811ff23..d18e4dd695e6 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -12,6 +12,10 @@ can_use_dsa_fused_store, fused_store_index_k_cache, ) +from sglang.srt.compilation.piecewise_context_manager import ( + get_forward_context, + is_in_piecewise_cuda_graph, +) from sglang.srt.environ import envs from sglang.srt.layers.attention.dsa.utils import ( aiter_can_use_preshuffle_paged_mqa, @@ -95,6 +99,80 @@ DUAL_STREAM_TOKEN_THRESHOLD = 1024 if _is_cuda else 0 +if _is_cuda: + from sglang.srt.compilation.compilation_config import register_split_op + from sglang.srt.utils.custom_op import register_custom_op + + @register_custom_op(mutates_args=["topk_result"]) + @register_split_op() + def k_cache_and_topk_result( + layer_id: int, + key: torch.Tensor, + q_fp8: torch.Tensor, + weights: torch.Tensor, + topk_result: torch.Tensor, + ) -> None: + assert ( + _is_cuda + ), "Internal error: piecewise CUDA graph is only supported on CUDA" + from sglang.srt.layers.attention.dsa.triton_kernel import act_quant + + forward_batch = get_forward_context().forward_batch + indexer = get_forward_context().dsa_indexers[layer_id] + metadata = get_attn_backend().get_indexer_metadata(layer_id, forward_batch) + + # slice off padding from piecewise CUDA graph + extend_num_tokens = forward_batch.extend_num_tokens + + indexer._store_index_k_cache( + forward_batch=forward_batch, + layer_id=layer_id, + key=key[:extend_num_tokens], + act_quant=act_quant, + out_cache_loc=forward_batch.out_cache_loc[:extend_num_tokens], + ) + indexer._get_topk_ragged( + False, + forward_batch, + layer_id, + q_fp8[:extend_num_tokens], + weights, + metadata, + topk_result, + ) + + def _logits_head_gate_pcg_fake_impl( + x: torch.Tensor, + weight: torch.Tensor, + n_heads_inv_sqrt: float, + softmax_scale: float, + q_scale: torch.Tensor, + ) -> torch.Tensor: + return torch.empty( + (x.shape[0], weight.shape[0], q_scale.shape[-1]), + dtype=torch.float32, + device=x.device, + ) + + @register_custom_op(fake_impl=_logits_head_gate_pcg_fake_impl) + def logits_head_gate_pcg( + x: torch.Tensor, + weight: torch.Tensor, + n_heads_inv_sqrt: float, + softmax_scale: float, + q_scale: torch.Tensor, + ) -> torch.Tensor: + from sglang.srt.layers.deep_gemm_wrapper import entrypoint as deep_gemm_wrapper + + out = torch.empty( + (x.shape[0], weight.shape[0]), dtype=torch.float32, device=x.device + ) + deep_gemm_wrapper.gemm_nt_bf16bf16f32(x, weight, out) + weights = out * n_heads_inv_sqrt + weights = weights.unsqueeze(-1) * q_scale * softmax_scale + return weights + + class BaseIndexerMetadata(ABC): @abstractmethod def get_seqlens_int32(self) -> torch.Tensor: @@ -441,7 +519,8 @@ def _get_k_bf16( def _update_rope_guarded(dst: torch.Tensor, src: torch.Tensor) -> None: # On AMD with in-place RoPE kernels, self-aliasing can occur; # skip write-back when src/dst tensors point to a single memory. - if src.data_ptr() == dst.data_ptr(): + # data_ptr() is not comparable inside torch.compile, so skip the guard there. + if not torch.compiler.is_compiling() and src.data_ptr() == dst.data_ptr(): return dst.copy_(src) @@ -627,6 +706,7 @@ def _get_topk_ragged( q_fp8: torch.Tensor, weights: torch.Tensor, metadata: BaseIndexerMetadata, + topk_result: Optional[torch.Tensor] = None, ) -> torch.Tensor: if TYPE_CHECKING: assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool) @@ -669,9 +749,10 @@ def _get_topk_ragged( device_index = device.index assert device_index is not None, "q_fp8 must be on an indexed CUDA device" - topk_result = torch.full( - (token_nums, self.index_topk), -1, device=device, dtype=torch.int32 - ) + if topk_result is None: + topk_result = torch.full( + (token_nums, self.index_topk), -1, device=device, dtype=torch.int32 + ) if batch_size == 0: return topk_result @@ -855,6 +936,9 @@ def _get_topk_ragged_with_cp( actual_seq_q: int, cp_index: List[Tuple[int, int, int]] = None, ) -> torch.Tensor: + assert ( + not is_in_piecewise_cuda_graph() + ), "DSA context parallel (_get_topk_ragged_with_cp) not supported under piecewise CUDA graph" if TYPE_CHECKING: assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool) @@ -1001,6 +1085,9 @@ def forward_indexer( topk: int, layer_id: int, ) -> Optional[torch.Tensor]: + assert ( + not is_in_piecewise_cuda_graph() + ), "DSA forward_indexer (non-CUDA loop path) not supported under piecewise CUDA graph" if not _is_npu: from sglang.srt.layers.attention.dsa.tilelang_kernel import fp8_index @@ -1083,21 +1170,26 @@ def _store_index_k_cache( key: torch.Tensor, *, act_quant=None, # fallback only + out_cache_loc: Optional[torch.Tensor] = None, ) -> None: """ Store DSA indexer K cache for current step. Preferred: fused_store_index_k_cache(key, cache, out_cache_loc, page_size) Fallback : act_quant(key) + token_to_kv_pool.set_index_k_scale_buffer(...) + + out_cache_loc will default to forward_batch.out_cache_loc if not provided. """ - # Fast path: JIT fused store (CUDA, page_size=64, non-fnuz) + if out_cache_loc is None: + out_cache_loc = forward_batch.out_cache_loc + if ( _is_cuda and (not _is_fp8_fnuz) and can_use_dsa_fused_store( key.dtype, - forward_batch.out_cache_loc.dtype, + out_cache_loc.dtype, get_token_to_kv_pool().page_size, ) ): @@ -1108,7 +1200,7 @@ def _store_index_k_cache( fused_store_index_k_cache( key, buf, - forward_batch.out_cache_loc, + out_cache_loc, get_token_to_kv_pool().page_size, ) return @@ -1141,13 +1233,12 @@ def _store_index_k_cache( assert act_quant is not None k_fp8, k_scale = act_quant(key, self.block_size, self.scale_fmt) - out_loc = forward_batch.out_cache_loc - if not out_loc.is_contiguous(): - out_loc = out_loc.contiguous() + if not out_cache_loc.is_contiguous(): + out_cache_loc = out_cache_loc.contiguous() get_token_to_kv_pool().set_index_k_scale_buffer( layer_id=layer_id, - loc=out_loc, + loc=out_cache_loc, index_k=k_fp8, index_k_scale=k_scale, ) @@ -1186,7 +1277,15 @@ def forward_cuda( # a tuple like (x_fp8, x_scale[, y]). Use `x_meta` for shape/device queries. x_meta = x[0] if isinstance(x, tuple) else x - metadata = get_attn_backend().get_indexer_metadata(layer_id, forward_batch) + # In piecewise CUDA graph mode, metadata is fetched inside custom ops via get_forward_context() to + # prevent Dynamo from guarding on forward_metadata identity (which changes each + # replay when init_forward_metadata creates a new ForwardMetadata object). + if not is_in_piecewise_cuda_graph(): + metadata = get_attn_backend().get_indexer_metadata(layer_id, forward_batch) + if metadata is None: + return None + else: + metadata = None enable_dual_stream = ( self.alt_stream is not None @@ -1195,14 +1294,13 @@ def forward_cuda( and q_lora.shape[0] <= DUAL_STREAM_TOKEN_THRESHOLD ) - # skip DSA if attention backend choose to skip this batch - if metadata is None: - return None - # Determine if should skip topk based on sequence length # We can only skip the logits computation if cuda graph is not involved skip_logits_computation = False - if forward_batch.forward_mode.is_extend_without_speculative(): + if ( + not is_in_piecewise_cuda_graph() + and forward_batch.forward_mode.is_extend_without_speculative() + ): if forward_batch.seq_lens_cpu is not None: max_kv_len = forward_batch.seq_lens_cpu.max().item() skip_logits_computation = max_kv_len <= self.index_topk @@ -1258,7 +1356,7 @@ def forward_cuda( act_quant=act_quant, ) current_stream.wait_stream(self.alt_stream) - else: + elif not is_in_piecewise_cuda_graph(): q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt) self._store_index_k_cache( forward_batch=forward_batch, @@ -1266,6 +1364,10 @@ def forward_cuda( key=key, act_quant=act_quant, ) + else: + # piecewise CUDA graph need to split graph on store_k_cache and mqa_logits, + # so delay store_k_cache after weights proj. + q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt) # aiter (ROCm gfx95): the 3-tuple (fp8, scale, bf16) from # fused_rms_fp8_group_quant is passed directly to _get_logits_head_gate, @@ -1308,25 +1410,37 @@ def forward_cuda( else: x_for_gate = x - weights = self._get_logits_head_gate(x_for_gate, q_scale) + if is_in_piecewise_cuda_graph(): + weights = logits_head_gate_pcg( + x_for_gate, + self.weights_proj.weight, + self.n_heads**-0.5, + self.softmax_scale, + q_scale, + ) + else: + weights = self._get_logits_head_gate(x_for_gate, q_scale) if _is_cuda or _is_hip: - assert forward_batch.seq_lens_cpu is not None - if len(forward_batch.seq_lens_cpu) == 0: - # this seems b/c max-pad, no worries? - # if x.shape[0] != 0: - # print( - # "HACK: seq_lens empty but x not empty, hackily return all-invalid topk_result" - # ) - return maybe_capture_indexer_topk( - layer_id, - torch.full( - (x_meta.shape[0], self.index_topk), - -1, - dtype=torch.int, - device=x_meta.device, - ), - ) + # In piecewise CUDA graph, any access to seq_lens_cpu creates a Dynamo shape guard. + # Piecewise CUDA graph never has empty batches. + if not is_in_piecewise_cuda_graph(): + assert forward_batch.seq_lens_cpu is not None + if len(forward_batch.seq_lens_cpu) == 0: + # this seems b/c max-pad, no worries? + # if x.shape[0] != 0: + # print( + # "HACK: seq_lens empty but x not empty, hackily return all-invalid topk_result" + # ) + return maybe_capture_indexer_topk( + layer_id, + torch.full( + (x_meta.shape[0], self.index_topk), + -1, + dtype=torch.int, + device=x_meta.device, + ), + ) if ( forward_batch.forward_mode.is_decode_or_idle() @@ -1379,6 +1493,24 @@ def forward_cuda( layer_id, torch.cat([topk_result_prev, topk_result_next], dim=0), ) + elif is_in_piecewise_cuda_graph(): + assert ( + not enable_dual_stream + ), "Internal error: piecewise CUDA graph should not be enabled with dual stream" + + topk_result = torch.full( + (q_fp8.shape[0], self.index_topk), + -1, + device=q_fp8.device, + dtype=torch.int32, + ) + k_cache_and_topk_result( + layer_id=layer_id, + key=key, + q_fp8=q_fp8, + weights=weights, + topk_result=topk_result, + ) else: topk_result = self._get_topk_ragged( enable_dual_stream, diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 9e5c8f1347e5..c9f1f4028868 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -1,5 +1,6 @@ from __future__ import annotations +import logging from dataclasses import dataclass from enum import IntEnum, auto from typing import TYPE_CHECKING, Dict, List, Literal, Optional, Tuple, TypeAlias @@ -7,6 +8,8 @@ import torch from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa + +logger = logging.getLogger(__name__) from sglang.srt.environ import envs from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.dsa.dequant_k_cache import dequantize_k_cache_paged @@ -2180,8 +2183,8 @@ def _forward_trtllm( backend="trtllm-gen", skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(), ) - # Output: [batch, q_len=1, heads, v_dim] -> [batch, heads, v_dim] - return out.squeeze(1) + + return out def _pad_topk_indices( self, topk_indices: torch.Tensor, num_tokens: int @@ -2212,10 +2215,18 @@ def set_dsa_prefill_impl(self, forward_batch: Optional[ForwardBatch] = None): """ Decide all attention prefill dispatch strategies for this batch. """ + from sglang.srt.compilation.piecewise_context_manager import ( + is_in_piecewise_cuda_graph, + ) from sglang.srt.utils import get_device_sm, is_blackwell # Decide MHA vs MLA - if forward_batch and forward_batch.forward_mode.is_extend_without_speculative(): + if is_in_piecewise_cuda_graph(): + # Can't branch on seq_lens_cpu in PCG, force mha off to guarantee correctness. + self.use_mha = False + elif ( + forward_batch and forward_batch.forward_mode.is_extend_without_speculative() + ): # Check if sequence meets criteria for MHA_ONE_SHOT assert forward_batch.seq_lens_cpu is not None max_kv_len = forward_batch.seq_lens_cpu.max().item() diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index e9c9e7be8c71..90e42e893336 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -53,7 +53,26 @@ if _is_cuda or _is_xpu or _is_musa: if _is_flashinfer_available: try: - from flashinfer.norm import layernorm + import flashinfer.norm + + from sglang.srt.utils.custom_op import register_custom_op + + def _layernorm_fake_impl( + input: torch.Tensor, + gamma: torch.Tensor, + beta: torch.Tensor, + eps: float = 1e-6, + ) -> torch.Tensor: + return torch.empty_like(input) + + @register_custom_op(fake_impl=_layernorm_fake_impl) + def layernorm( + input: torch.Tensor, + gamma: torch.Tensor, + beta: torch.Tensor, + eps: float = 1e-6, + ) -> torch.Tensor: + return flashinfer.norm.layernorm(input, gamma, beta, eps) _flashinfer_layernorm_available = True except (ImportError, AttributeError): diff --git a/python/sglang/srt/layers/radix_attention.py b/python/sglang/srt/layers/radix_attention.py index 1e8784f1d53b..c3409c8c27e9 100644 --- a/python/sglang/srt/layers/radix_attention.py +++ b/python/sglang/srt/layers/radix_attention.py @@ -160,6 +160,12 @@ def unified_attention_with_output( q_rope: Optional[torch.Tensor] = None, k_rope: Optional[torch.Tensor] = None, sinks: Optional[torch.Tensor] = None, + # MLA / TRT-LLM / NSA paths pass these through RadixAttention.forward(**kwargs); + # they must appear in the schema when --enforce-piecewise-cuda-graph is on. + cos_sin_cache: Optional[torch.Tensor] = None, + is_neox: Optional[bool] = None, + llama_4_scaling: Optional[torch.Tensor] = None, + topk_indices: Optional[torch.Tensor] = None, ) -> None: context = get_forward_context() forward_batch = context.forward_batch @@ -178,6 +184,14 @@ def unified_attention_with_output( kwargs["k_rope"] = k_rope[:real_num_tokens] if sinks is not None: kwargs["sinks"] = sinks + if cos_sin_cache is not None: + kwargs["cos_sin_cache"] = cos_sin_cache + if is_neox is not None: + kwargs["is_neox"] = is_neox + if llama_4_scaling is not None: + kwargs["llama_4_scaling"] = llama_4_scaling + if topk_indices is not None: + kwargs["topk_indices"] = topk_indices[:real_num_tokens] original_out_cache_loc = forward_batch.out_cache_loc # Keep the original ForwardBatch object and only narrow cache locations for diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 52a483802dae..b204d15da936 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -2851,6 +2851,7 @@ def init_piecewise_cuda_graphs(self, force_for_draft_worker: bool = False): self.attention_layers = [] self.moe_layers = [] self.moe_fusions = [] + self.dsa_indexers = [] for layer in layer_model.layers: attn_layer = None if hasattr(layer, "self_attn"): @@ -2903,6 +2904,11 @@ def init_piecewise_cuda_graphs(self, force_for_draft_worker: bool = False): moe_fusion = layer.mixer self.moe_layers.append(moe_block) self.moe_fusions.append(moe_fusion) + # NSA indexers (None for layers without NSA) + dsa_indexer = None + if hasattr(layer, "self_attn") and hasattr(layer.self_attn, "indexer"): + dsa_indexer = layer.self_attn.indexer + self.dsa_indexers.append(dsa_indexer) if len(self.attention_layers) < self.model_config.num_hidden_layers: # TODO(yuwei): support Non-Standard GQA diff --git a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py index 8653bce86c74..877cebf2de10 100644 --- a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py @@ -298,6 +298,7 @@ def __init__(self, model_runner: ModelRunner): self.attention_layers = self.model_runner.attention_layers self.moe_layers = self.model_runner.moe_layers self.moe_fusions = self.model_runner.moe_fusions + self.dsa_indexers = getattr(self.model_runner, "dsa_indexers", None) if get_global_graph_memory_pool() is None: set_global_graph_memory_pool(self.device_module.graph_pool_handle()) @@ -432,6 +433,7 @@ def warmup_compile(self, num_tokens: int): self.quant_config, self.moe_layers, self.moe_fusions, + dsa_indexers=self.dsa_indexers, ): _ = self.model_runner.model.forward( forward_batch.input_ids, @@ -622,6 +624,7 @@ def run_once(): self.quant_config, self.moe_layers, self.moe_fusions, + dsa_indexers=self.dsa_indexers, ): self.model_runner.model.forward( forward_batch.input_ids, @@ -795,6 +798,7 @@ def replay( self.quant_config, self.moe_layers, self.moe_fusions, + dsa_indexers=self.dsa_indexers, ): # Due to the dispatch kernel for MLA model, we init the metadata with original forward_batch self.model_runner.attn_backend.init_forward_metadata(forward_batch) diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 44fd33cb6b04..bbecb680adbf 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -1980,6 +1980,7 @@ def forward( if ( isinstance(self.mlp, DeepseekV2MoE) and not self.mlp.experts.moe_runner_config.inplace + and not torch.compiler.is_compiling() ): from sglang.srt.layers.moe.moe_runner.base import moe_output_buffer_ctx diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 1d1b8d29959d..961b1b787a5e 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1347,6 +1347,10 @@ def _handle_piecewise_cuda_graph(self): # 18. CUDA Graph debug mode if self.debug_cuda_graph: self.disable_piecewise_cuda_graph = True + # 19. DSA prefill context parallelism (attn_cp_size is set later in + # _handle_model_specific_adjustments, so check the flag directly here) + if self.enable_dsa_prefill_context_parallel: + self.disable_piecewise_cuda_graph = True def _handle_multi_item_scoring(self): """Setup and validate multi-item scoring constraints. @@ -3183,22 +3187,6 @@ def _handle_moe_kernel_config(self): self.ep_size == 1 ), "FP8/MXFP8 Cutlass MoE is only supported with ep_size == 1" - # TODO(yuwei): Fix piecewise cuda graph support for bypassed topk MoE backends. - # Exception: GptOssForCausalLM wraps the entire MoE block in its own - # custom op (moe_impl), so bypassed topk is handled inside the op body. - if ( - not self.enforce_piecewise_cuda_graph - and self.moe_runner_backend in ("flashinfer_trtllm", "flashinfer_mxfp4") - and self.get_model_config().hf_config.architectures[0] - != "GptOssForCausalLM" - ): - self.disable_piecewise_cuda_graph = True - logger.info( - f"Piecewise cuda graph is disabled for MoE runner backend " - f"'{self.moe_runner_backend}' (bypassed topk is incompatible " - f"with torch.compile)." - ) - def _handle_a2a_moe(self): if self.enable_deepep_waterfill and self.moe_a2a_backend != "deepep": logger.warning( diff --git a/test/registered/piecewise_cuda_graph/test_pcg_glm5_fp4.py b/test/registered/piecewise_cuda_graph/test_pcg_glm5_fp4.py new file mode 100644 index 000000000000..6cce520debf1 --- /dev/null +++ b/test/registered/piecewise_cuda_graph/test_pcg_glm5_fp4.py @@ -0,0 +1,71 @@ +import unittest +from types import SimpleNamespace + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.run_eval import run_eval +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=900, stage="base-c", runner_config="4-gpu-b200") + +GLM5_FP4_MODEL = "nvidia/GLM-5-NVFP4" + + +class TestPCGGlm5Fp4(CustomTestCase): + """PCG prefill on GLM-5-NVFP4 (DSA model, TP=4, B200). + + GLM-5 uses GlmMoeDsaForCausalLM (DSA attention). This test verifies that + piecewise CUDA graph works correctly after the DSA indexer was updated to + cache k_fp8/k_scale for PCG-compatible prefill. + """ + + @classmethod + def setUpClass(cls): + cls.model = GLM5_FP4_MODEL + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--tp-size", + "4", + "--trust-remote-code", + "--reasoning-parser", + "glm45", + "--tool-call-parser", + "glm47", + "--quantization", + "modelopt_fp4", + "--disable-flashinfer-autotune", + "--enforce-piecewise-cuda-graph", + "--model-loader-extra-config", + '{"enable_multithread_load": true, "num_threads": 64}', + ], + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_gsm8k(self): + args = SimpleNamespace( + base_url=self.base_url, + model=self.model, + eval_name="gsm8k", + num_examples=200, + num_threads=200, + max_tokens=4096, + ) + metrics = run_eval(args) + print(f"{metrics=}") + self.assertGreater(metrics["score"], 0.92) + + +if __name__ == "__main__": + unittest.main()