diff --git a/python/sglang/srt/arg_groups/speculative_hook.py b/python/sglang/srt/arg_groups/speculative_hook.py index cbda12929..0c24ae32e 100644 --- a/python/sglang/srt/arg_groups/speculative_hook.py +++ b/python/sglang/srt/arg_groups/speculative_hook.py @@ -438,7 +438,14 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None: """ from sglang.srt.utils import is_hip - supported_draft_backends = ("flashinfer", "fa3", "fa4", "triton", "ascend") + supported_draft_backends = ( + "flashinfer", + "fa3", + "fa4", + "triton", + "trtllm_mha", + "ascend", + ) # Use triton on ROCm (no FlashInfer), flashinfer on CUDA. fallback_backend = "triton" if is_hip() else "flashinfer" @@ -453,13 +460,37 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None: if draft_backend is None: draft_backend = fallback_backend elif draft_backend == "trtllm_mha": - logger.warning( - "DFLASH draft worker does not support 'trtllm_mha' because the " - "draft path requires per-layer DFlash attention. Falling back to " - "'%s'.", - fallback_backend, + from sglang.srt.speculative.dflash_utils import get_dflash_layer_types + from sglang.srt.utils.hf_transformers_utils import get_config + + draft_hf_config = get_config( + server_args.speculative_draft_model_path, + trust_remote_code=server_args.trust_remote_code, + revision=server_args.speculative_draft_model_revision, + model_override_args=json.loads(server_args.json_model_override_args), ) - draft_backend = fallback_backend + draft_text_config = ( + getattr(draft_hf_config, "text_config", None) or draft_hf_config + ) + layer_types = get_dflash_layer_types(draft_hf_config) + num_layers = getattr(draft_text_config, "num_hidden_layers", None) + all_sliding = ( + layer_types + and len(layer_types) == num_layers + and set(layer_types) == {"sliding_attention"} + ) + all_causal = getattr(draft_text_config, "is_causal", False) is True + if not (all_sliding or all_causal): + logger.warning( + "DFLASH only enables 'trtllm_mha' when all layers use sliding " + "attention or the draft is explicitly causal; got " + "layer_types=%r, is_causal=%r. " + "Falling back to '%s'.", + layer_types, + getattr(draft_text_config, "is_causal", None), + fallback_backend, + ) + draft_backend = fallback_backend elif draft_backend not in supported_draft_backends: logger.warning( "DFLASH draft worker only supports attention_backend in %s for now, " diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index d4ecd4564..3edd24ed3 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -41,7 +41,7 @@ import copy import inspect import logging from contextlib import contextmanager -from typing import TYPE_CHECKING, Dict, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, Optional, Union import torch import tqdm @@ -120,6 +120,25 @@ logger = logging.getLogger(__name__) _MAX_PREFILL_CUDA_GRAPH_PADDING_FACTOR = 2 +def _slice_output_rows(output: Any, num_tokens: int) -> Any: + """Slice every tensor leaf in a transformer-body output by token rows. + + Full prefill graphs replay at a padded token bucket. Most models return a + single hidden-state tensor, while auxiliary-hidden-state models (for + example DFLASH targets) return nested tuple/list structures. Slicing the + outer sequence would change its structure instead of removing padded rows. + """ + if output is None: + return None + if torch.is_tensor(output) or isinstance(output, PPProxyTensors): + return output[:num_tokens] + if isinstance(output, tuple): + return tuple(_slice_output_rows(item, num_tokens) for item in output) + if isinstance(output, list): + return [_slice_output_rows(item, num_tokens) for item in output] + raise TypeError(f"Unsupported full prefill CUDA graph output: {type(output)}") + + def prefill_failure_msg(backend_name: str) -> str: """Render PREFILL_CUDA_GRAPH_CAPTURE_FAILED_MSG with a backend-specific numbered suggestion list. The runner is only constructed for BREAKABLE @@ -1189,7 +1208,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): 1, static_n )[: ie.shape[0]].copy_(ie) hs = self.backend.replay(shape_key, static_forward_batch, **kwargs) - return hs[:raw_num_tokens] if full_path else hs + return _slice_output_rows(hs, raw_num_tokens) if full_path else hs original_layer_forward = self.layer_model.forward self.layer_model.forward = replay_layer_forward diff --git a/python/sglang/srt/models/dflash.py b/python/sglang/srt/models/dflash.py index ffef278fb..69c01c7ad 100644 --- a/python/sglang/srt/models/dflash.py +++ b/python/sglang/srt/models/dflash.py @@ -57,7 +57,13 @@ def _get_dflash_layer_attention_params( layer_type = layer_types[layer_id] if layer_type == "full_attention": - return -1, AttentionType.ENCODER_ONLY + text_config = getattr(config, "text_config", None) or config + attention_type = ( + AttentionType.DECODER + if getattr(text_config, "is_causal", False) + else AttentionType.ENCODER_ONLY + ) + return -1, attention_type if layer_type == "sliding_attention": sliding_window_size = get_dflash_attention_sliding_window_size(config) assert sliding_window_size is not None diff --git a/python/sglang/srt/models/inkling.py b/python/sglang/srt/models/inkling.py index 6d71fb7b1..9cba08763 100644 --- a/python/sglang/srt/models/inkling.py +++ b/python/sglang/srt/models/inkling.py @@ -677,6 +677,32 @@ class InklingCausalLLM(nn.Module): prefix=add_prefix("lm_head", prefix), ) self.logits_processor = LogitsProcessor(config) + self._dflash_layers_to_capture: set[int] = set() + + def set_dflash_layers_to_capture(self, layer_ids: list[int]) -> None: + """Capture post-layer hidden states consumed by a DFLASH drafter.""" + if layer_ids is None: + raise ValueError("DFLASH requires explicit target layer IDs.") + if len(layer_ids) != len(set(layer_ids)): + raise ValueError(f"DFLASH target layer IDs must be unique: {layer_ids}") + if layer_ids != sorted(layer_ids): + raise ValueError(f"DFLASH target layer IDs must be sorted: {layer_ids}") + invalid = [idx for idx in layer_ids if idx < 0 or idx >= len(self.layers)] + if invalid: + raise ValueError( + f"DFLASH target layer IDs out of range [0, {len(self.layers)}): {invalid}" + ) + + self._dflash_layers_to_capture = set(layer_ids) + # Inkling can defer an MoE all-reduce into the next layer. A tapped + # layer must instead materialize a complete hidden state at its tap. + for layer in self.layers: + if not hasattr(layer, "_mlp_ar_fusable_without_dflash"): + layer._mlp_ar_fusable_without_dflash = layer.mlp_ar_fusable + layer.mlp_ar_fusable = ( + layer._mlp_ar_fusable_without_dflash + and layer.layer_id not in self._dflash_layers_to_capture + ) def get_input_embeddings(self): # Fold embed_norm into the embedding so general_mm_embed_routine norms the text @@ -778,6 +804,12 @@ class InklingCausalLLM(nn.Module): fuse_ar_sconv = True fuse_attn_ar = True prev_mlp_partial = False + aux_hidden_states: Optional[list[torch.Tensor]] = ( + [] + if self._dflash_layers_to_capture + and not forward_batch.forward_mode.is_idle() + else None + ) for layer in self.layers: hidden_states, residual = layer( hidden_states, @@ -792,6 +824,18 @@ class InklingCausalLLM(nn.Module): ) prev_mlp_sconv = layer.mlp_sconv prev_mlp_partial = fuse_ar_sconv and layer.mlp_ar_fusable + if ( + aux_hidden_states is not None + and layer.layer_id in self._dflash_layers_to_capture + ): + # The trained taps are post-layer and precede the deferred + # mlp_sconv belonging to this layer. + tap_hidden = hidden_states + if layer.scattered_sconv: + tap_hidden = all_gather_hidden(tap_hidden, layer.attn_tp_group) + aux_hidden_states.append( + tap_hidden if residual is None else tap_hidden + residual + ) # The final layer's mlp_sconv was deferred; run it now — as an eager break # under BCG (so it re-reads live per-seq metadata at replay), else inline. if prev_mlp_sconv is not None and not forward_batch.forward_mode.is_idle(): @@ -808,7 +852,11 @@ class InklingCausalLLM(nn.Module): norm=self.norm, norm_residual=residual, ) - return hidden_states + return ( + (hidden_states, aux_hidden_states) + if self._dflash_layers_to_capture + else hidden_states + ) # Fused extend tail: {AR + scattered sconv}, then the final # norm unfused on the gathered [T, H]. hidden_states = ar_scattered_sconv_fused( @@ -818,7 +866,11 @@ class InklingCausalLLM(nn.Module): get_tensor_model_parallel_group(), ) hidden_states, _ = self.norm(hidden_states, residual) - return hidden_states + return ( + (hidden_states, aux_hidden_states) + if self._dflash_layers_to_capture + else hidden_states + ) if prev_mlp_partial: fm = forward_batch.forward_mode if fm.is_decode() or fm.is_target_verify(): @@ -832,7 +884,11 @@ class InklingCausalLLM(nn.Module): forward_batch, get_tensor_model_parallel_group(), ) - return hidden_states + return ( + (hidden_states, aux_hidden_states) + if self._dflash_layers_to_capture + else hidden_states + ) # Fused extend tail: {AR + full-width sconv + cache update} # (non-scattered), then the final norm unfused. hidden_states = ar_fullwidth_sconv_fused( @@ -842,7 +898,11 @@ class InklingCausalLLM(nn.Module): get_tensor_model_parallel_group(), ) hidden_states, _ = self.norm(hidden_states, residual) - return hidden_states + return ( + (hidden_states, aux_hidden_states) + if self._dflash_layers_to_capture + else hidden_states + ) # Same gate as the per-layer group: the eager break needs the tc_piecewise # context (installed only by the prefill BCG runner) to read the live # forward_batch at replay; else run inline with the passed forward_batch. @@ -870,7 +930,11 @@ class InklingCausalLLM(nn.Module): hidden_states, self.layers[-1].attn_tp_group ) hidden_states, _ = self.norm(hidden_states, residual) - return hidden_states + return ( + (hidden_states, aux_hidden_states) + if self._dflash_layers_to_capture + else hidden_states + ) class InklingAudio(nn.Module): @@ -1066,6 +1130,14 @@ class InklingForConditionalGeneration(nn.Module): def get_embed_and_head(self): return self.llm.embed_tokens.weight, self.llm.lm_head.weight + @property + def lm_head(self) -> nn.Module: + """Expose the target head through the common speculative API.""" + return self.llm.lm_head + + def set_dflash_layers_to_capture(self, layer_ids: list[int]) -> None: + self.llm.set_dflash_layers_to_capture(layer_ids) + def get_num_kv_cache_layers(self) -> int: return self.text_config.num_hidden_layers @@ -1101,6 +1173,9 @@ class InklingForConditionalGeneration(nn.Module): data_embedding_funcs=data_embedding_funcs, positions=positions, ) + aux_hidden_states = None + if self.llm._dflash_layers_to_capture: + hidden_states, aux_hidden_states = hidden_states mup_width_multiplier = self.config.text_config.logits_mup_width_multiplier hidden_states_for_logits = ( hidden_states @@ -1115,7 +1190,12 @@ class InklingForConditionalGeneration(nn.Module): hidden_states_for_logits, self.llm.lm_head, forward_batch, - hidden_states_before_norm=hidden_states, + aux_hidden_states=aux_hidden_states, + # DFLASH needs the concatenated tap states. LogitsProcessor gives + # hidden_states_before_norm precedence, so omit it in this mode. + hidden_states_before_norm=( + None if aux_hidden_states is not None else hidden_states + ), ) def update_conv_state_after_mtp_verify( diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index 0be95e06d..0bfd3b249 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -324,11 +324,15 @@ class DFlashWorkerV2(BaseSpecWorker): def init_attention_backends(self): self._draft_worker.init_attention_backends() + target_model = self.model_runner.model self._need_mamba_verify_commit = mambaish_config( self.model_runner.model_config - ) is not None and hasattr( - self.model_runner.attn_backend, - "update_mamba_state_after_mtp_verify", + ) is not None and ( + hasattr( + self.model_runner.attn_backend, + "update_mamba_state_after_mtp_verify", + ) + or hasattr(target_model, "update_conv_state_after_mtp_verify") ) def init_cuda_graphs(self): @@ -1278,12 +1282,25 @@ class DFlashWorkerV2(BaseSpecWorker): torch.full_like(to_track_ith, -1, dtype=torch.int64), ) - attn_backend.update_mamba_state_after_mtp_verify( - last_correct_step_indices=last_correct_step_indices, - mamba_track_indices=batch.mamba_track_indices, - mamba_steps_to_track=mamba_steps_to_track, - model=self.target_worker.model_runner.model, - ) + model_runner = self.target_worker.model_runner + if hasattr(attn_backend, "update_mamba_state_after_mtp_verify"): + attn_backend.update_mamba_state_after_mtp_verify( + last_correct_step_indices=last_correct_step_indices, + mamba_track_indices=batch.mamba_track_indices, + mamba_steps_to_track=mamba_steps_to_track, + model=model_runner.model, + ) + elif hasattr(model_runner.model, "update_conv_state_after_mtp_verify"): + # Inkling's short convolutions access the mamba pool directly, so + # their accepted verify state is committed by the model rather + # than an attention-backend wrapper. + model_runner.model.update_conv_state_after_mtp_verify( + req_to_token_pool=model_runner.req_to_token_pool, + req_pool_indices=batch.req_pool_indices[: commit_lens.shape[0]], + last_correct_step_indices=last_correct_step_indices, + mamba_track_indices=batch.mamba_track_indices, + mamba_steps_to_track=mamba_steps_to_track, + ) def _ensure_accept_bonus_buffers(self, bs: int) -> None: if self._accept_bonus_buffer_cap >= int(bs): diff --git a/python/sglang/srt/speculative/draft_worker_common.py b/python/sglang/srt/speculative/draft_worker_common.py index 901aa88cf..a42d74bf8 100644 --- a/python/sglang/srt/speculative/draft_worker_common.py +++ b/python/sglang/srt/speculative/draft_worker_common.py @@ -22,7 +22,14 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) -_SUPPORTED_DRAFT_BACKENDS = ("flashinfer", "fa3", "fa4", "triton", "ascend") +_SUPPORTED_DRAFT_BACKENDS = ( + "flashinfer", + "fa3", + "fa4", + "triton", + "trtllm_mha", + "ascend", +) class DraftWorkerBundle(msgspec.Struct, frozen=True):