[Inkling] Add minimal DFLASH support (#31840)
This commit is contained in:
@@ -438,7 +438,14 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None:
|
|||||||
"""
|
"""
|
||||||
from sglang.srt.utils import is_hip
|
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.
|
# Use triton on ROCm (no FlashInfer), flashinfer on CUDA.
|
||||||
fallback_backend = "triton" if is_hip() else "flashinfer"
|
fallback_backend = "triton" if is_hip() else "flashinfer"
|
||||||
|
|
||||||
@@ -453,10 +460,34 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None:
|
|||||||
if draft_backend is None:
|
if draft_backend is None:
|
||||||
draft_backend = fallback_backend
|
draft_backend = fallback_backend
|
||||||
elif draft_backend == "trtllm_mha":
|
elif draft_backend == "trtllm_mha":
|
||||||
|
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_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(
|
logger.warning(
|
||||||
"DFLASH draft worker does not support 'trtllm_mha' because the "
|
"DFLASH only enables 'trtllm_mha' when all layers use sliding "
|
||||||
"draft path requires per-layer DFlash attention. Falling back to "
|
"attention or the draft is explicitly causal; got "
|
||||||
"'%s'.",
|
"layer_types=%r, is_causal=%r. "
|
||||||
|
"Falling back to '%s'.",
|
||||||
|
layer_types,
|
||||||
|
getattr(draft_text_config, "is_causal", None),
|
||||||
fallback_backend,
|
fallback_backend,
|
||||||
)
|
)
|
||||||
draft_backend = fallback_backend
|
draft_backend = fallback_backend
|
||||||
|
|||||||
@@ -41,7 +41,7 @@ import copy
|
|||||||
import inspect
|
import inspect
|
||||||
import logging
|
import logging
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from typing import TYPE_CHECKING, Dict, Optional, Union
|
from typing import TYPE_CHECKING, Any, Dict, Optional, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import tqdm
|
import tqdm
|
||||||
@@ -120,6 +120,25 @@ logger = logging.getLogger(__name__)
|
|||||||
_MAX_PREFILL_CUDA_GRAPH_PADDING_FACTOR = 2
|
_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:
|
def prefill_failure_msg(backend_name: str) -> str:
|
||||||
"""Render PREFILL_CUDA_GRAPH_CAPTURE_FAILED_MSG with a backend-specific
|
"""Render PREFILL_CUDA_GRAPH_CAPTURE_FAILED_MSG with a backend-specific
|
||||||
numbered suggestion list. The runner is only constructed for BREAKABLE
|
numbered suggestion list. The runner is only constructed for BREAKABLE
|
||||||
@@ -1189,7 +1208,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
1, static_n
|
1, static_n
|
||||||
)[: ie.shape[0]].copy_(ie)
|
)[: ie.shape[0]].copy_(ie)
|
||||||
hs = self.backend.replay(shape_key, static_forward_batch, **kwargs)
|
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
|
original_layer_forward = self.layer_model.forward
|
||||||
self.layer_model.forward = replay_layer_forward
|
self.layer_model.forward = replay_layer_forward
|
||||||
|
|||||||
@@ -57,7 +57,13 @@ def _get_dflash_layer_attention_params(
|
|||||||
|
|
||||||
layer_type = layer_types[layer_id]
|
layer_type = layer_types[layer_id]
|
||||||
if layer_type == "full_attention":
|
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":
|
if layer_type == "sliding_attention":
|
||||||
sliding_window_size = get_dflash_attention_sliding_window_size(config)
|
sliding_window_size = get_dflash_attention_sliding_window_size(config)
|
||||||
assert sliding_window_size is not None
|
assert sliding_window_size is not None
|
||||||
|
|||||||
@@ -677,6 +677,32 @@ class InklingCausalLLM(nn.Module):
|
|||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
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):
|
def get_input_embeddings(self):
|
||||||
# Fold embed_norm into the embedding so general_mm_embed_routine norms the text
|
# 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_ar_sconv = True
|
||||||
fuse_attn_ar = True
|
fuse_attn_ar = True
|
||||||
prev_mlp_partial = False
|
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:
|
for layer in self.layers:
|
||||||
hidden_states, residual = layer(
|
hidden_states, residual = layer(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
@@ -792,6 +824,18 @@ class InklingCausalLLM(nn.Module):
|
|||||||
)
|
)
|
||||||
prev_mlp_sconv = layer.mlp_sconv
|
prev_mlp_sconv = layer.mlp_sconv
|
||||||
prev_mlp_partial = fuse_ar_sconv and layer.mlp_ar_fusable
|
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
|
# 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.
|
# 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():
|
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=self.norm,
|
||||||
norm_residual=residual,
|
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
|
# Fused extend tail: {AR + scattered sconv}, then the final
|
||||||
# norm unfused on the gathered [T, H].
|
# norm unfused on the gathered [T, H].
|
||||||
hidden_states = ar_scattered_sconv_fused(
|
hidden_states = ar_scattered_sconv_fused(
|
||||||
@@ -818,7 +866,11 @@ class InklingCausalLLM(nn.Module):
|
|||||||
get_tensor_model_parallel_group(),
|
get_tensor_model_parallel_group(),
|
||||||
)
|
)
|
||||||
hidden_states, _ = self.norm(hidden_states, residual)
|
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:
|
if prev_mlp_partial:
|
||||||
fm = forward_batch.forward_mode
|
fm = forward_batch.forward_mode
|
||||||
if fm.is_decode() or fm.is_target_verify():
|
if fm.is_decode() or fm.is_target_verify():
|
||||||
@@ -832,7 +884,11 @@ class InklingCausalLLM(nn.Module):
|
|||||||
forward_batch,
|
forward_batch,
|
||||||
get_tensor_model_parallel_group(),
|
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}
|
# Fused extend tail: {AR + full-width sconv + cache update}
|
||||||
# (non-scattered), then the final norm unfused.
|
# (non-scattered), then the final norm unfused.
|
||||||
hidden_states = ar_fullwidth_sconv_fused(
|
hidden_states = ar_fullwidth_sconv_fused(
|
||||||
@@ -842,7 +898,11 @@ class InklingCausalLLM(nn.Module):
|
|||||||
get_tensor_model_parallel_group(),
|
get_tensor_model_parallel_group(),
|
||||||
)
|
)
|
||||||
hidden_states, _ = self.norm(hidden_states, residual)
|
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
|
# 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
|
# context (installed only by the prefill BCG runner) to read the live
|
||||||
# forward_batch at replay; else run inline with the passed forward_batch.
|
# 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.layers[-1].attn_tp_group
|
||||||
)
|
)
|
||||||
hidden_states, _ = self.norm(hidden_states, residual)
|
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):
|
class InklingAudio(nn.Module):
|
||||||
@@ -1066,6 +1130,14 @@ class InklingForConditionalGeneration(nn.Module):
|
|||||||
def get_embed_and_head(self):
|
def get_embed_and_head(self):
|
||||||
return self.llm.embed_tokens.weight, self.llm.lm_head.weight
|
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:
|
def get_num_kv_cache_layers(self) -> int:
|
||||||
return self.text_config.num_hidden_layers
|
return self.text_config.num_hidden_layers
|
||||||
|
|
||||||
@@ -1101,6 +1173,9 @@ class InklingForConditionalGeneration(nn.Module):
|
|||||||
data_embedding_funcs=data_embedding_funcs,
|
data_embedding_funcs=data_embedding_funcs,
|
||||||
positions=positions,
|
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
|
mup_width_multiplier = self.config.text_config.logits_mup_width_multiplier
|
||||||
hidden_states_for_logits = (
|
hidden_states_for_logits = (
|
||||||
hidden_states
|
hidden_states
|
||||||
@@ -1115,7 +1190,12 @@ class InklingForConditionalGeneration(nn.Module):
|
|||||||
hidden_states_for_logits,
|
hidden_states_for_logits,
|
||||||
self.llm.lm_head,
|
self.llm.lm_head,
|
||||||
forward_batch,
|
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(
|
def update_conv_state_after_mtp_verify(
|
||||||
|
|||||||
@@ -324,12 +324,16 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
def init_attention_backends(self):
|
def init_attention_backends(self):
|
||||||
self._draft_worker.init_attention_backends()
|
self._draft_worker.init_attention_backends()
|
||||||
|
target_model = self.model_runner.model
|
||||||
self._need_mamba_verify_commit = mambaish_config(
|
self._need_mamba_verify_commit = mambaish_config(
|
||||||
self.model_runner.model_config
|
self.model_runner.model_config
|
||||||
) is not None and hasattr(
|
) is not None and (
|
||||||
|
hasattr(
|
||||||
self.model_runner.attn_backend,
|
self.model_runner.attn_backend,
|
||||||
"update_mamba_state_after_mtp_verify",
|
"update_mamba_state_after_mtp_verify",
|
||||||
)
|
)
|
||||||
|
or hasattr(target_model, "update_conv_state_after_mtp_verify")
|
||||||
|
)
|
||||||
|
|
||||||
def init_cuda_graphs(self):
|
def init_cuda_graphs(self):
|
||||||
capture_decode_cuda_graph = (
|
capture_decode_cuda_graph = (
|
||||||
@@ -1278,11 +1282,24 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
torch.full_like(to_track_ith, -1, dtype=torch.int64),
|
torch.full_like(to_track_ith, -1, dtype=torch.int64),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
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(
|
attn_backend.update_mamba_state_after_mtp_verify(
|
||||||
last_correct_step_indices=last_correct_step_indices,
|
last_correct_step_indices=last_correct_step_indices,
|
||||||
mamba_track_indices=batch.mamba_track_indices,
|
mamba_track_indices=batch.mamba_track_indices,
|
||||||
mamba_steps_to_track=mamba_steps_to_track,
|
mamba_steps_to_track=mamba_steps_to_track,
|
||||||
model=self.target_worker.model_runner.model,
|
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:
|
def _ensure_accept_bonus_buffers(self, bs: int) -> None:
|
||||||
|
|||||||
@@ -22,7 +22,14 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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):
|
class DraftWorkerBundle(msgspec.Struct, frozen=True):
|
||||||
|
|||||||
Reference in New Issue
Block a user