run pass cp+megamoe+bcg
(cherry picked from commit 85b8a150ef69c35cbd8c3fe940e41b414b903757)
This commit is contained in:
@@ -750,7 +750,14 @@ class DSV4AttnMetadata:
|
|||||||
if src_val is None and dst_val is None:
|
if src_val is None and dst_val is None:
|
||||||
continue
|
continue
|
||||||
assert dst_val is not None, f"{field_name=} {src_val=} {dst_val=}"
|
assert dst_val is not None, f"{field_name=} {src_val=} {dst_val=}"
|
||||||
dst_val.copy_(src_val)
|
shape_mismatch = dst_val.shape != src_val.shape
|
||||||
|
assert not shape_mismatch or field_name in self._CP_GLOBAL_FIELDS, (
|
||||||
|
f"Only CP-global replay metadata may use a shorter live prefix, "
|
||||||
|
f"got {field_name=} {src_val.shape=} {dst_val.shape=}"
|
||||||
|
)
|
||||||
|
_copy_tensor_allowing_storage_alias(
|
||||||
|
dst_val, src_val, pad_value=0 if shape_mismatch else None
|
||||||
|
)
|
||||||
|
|
||||||
# These fields are safe to replace because captured kernels only need
|
# These fields are safe to replace because captured kernels only need
|
||||||
# the current per-replay objects, or the field is produced inside the
|
# the current per-replay objects, or the field is produced inside the
|
||||||
@@ -1002,6 +1009,27 @@ def _prefill_graph_max_seq_len() -> Optional[int]:
|
|||||||
return get_exec().graph.cuda_graph_config.prefill.max_seq_len
|
return get_exec().graph.cuda_graph_config.prefill.max_seq_len
|
||||||
|
|
||||||
|
|
||||||
|
def _copy_tensor_allowing_storage_alias(
|
||||||
|
dst: torch.Tensor, src: torch.Tensor, *, pad_value: Optional[int] = None
|
||||||
|
) -> None:
|
||||||
|
"""Copy replay metadata while preserving capture-stable destination addresses."""
|
||||||
|
if dst is src:
|
||||||
|
return
|
||||||
|
if dst.untyped_storage().data_ptr() == src.untyped_storage().data_ptr():
|
||||||
|
src = src.clone()
|
||||||
|
if dst.shape == src.shape:
|
||||||
|
dst.copy_(src)
|
||||||
|
return
|
||||||
|
assert (
|
||||||
|
pad_value is not None
|
||||||
|
and dst.ndim == src.ndim
|
||||||
|
and dst.shape[0] >= src.shape[0]
|
||||||
|
and dst.shape[1:] == src.shape[1:]
|
||||||
|
), f"Cannot copy replay metadata from {src.shape=} to {dst.shape=}"
|
||||||
|
dst.fill_(pad_value)
|
||||||
|
dst[: src.shape[0]].copy_(src)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class DSV4Metadata:
|
class DSV4Metadata:
|
||||||
core_attn_metadata: DSV4AttnMetadata
|
core_attn_metadata: DSV4AttnMetadata
|
||||||
|
|||||||
@@ -23,17 +23,22 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.arg_groups.overrides import (
|
from sglang.srt.arg_groups.overrides import (
|
||||||
attention_backends_of,
|
attention_backends_of,
|
||||||
|
model_config_of,
|
||||||
resolved_view,
|
resolved_view,
|
||||||
resolving_view,
|
resolving_view,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.configs.model_config import is_deepseek_v4
|
||||||
from sglang.srt.layers.cp.base import get_cp_strategy
|
from sglang.srt.layers.cp.base import get_cp_strategy
|
||||||
|
from sglang.srt.layers.cp.interleave import InterleaveCPStrategy
|
||||||
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
|
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
|
||||||
from sglang.srt.layers.cp.utils import (
|
from sglang.srt.layers.cp.utils import (
|
||||||
cp_gather_after_forward,
|
cp_gather_after_forward,
|
||||||
|
cp_shard_hidden_states,
|
||||||
cp_split_before_forward,
|
cp_split_before_forward,
|
||||||
prepare_cp_forward,
|
prepare_cp_forward,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy
|
from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy
|
||||||
|
from sglang.srt.layers.logits_processor import LogitsMetadata
|
||||||
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -50,12 +55,18 @@ def supports_prefill_cp_bcg(server_args: ServerArgs) -> bool:
|
|||||||
cfg = resolving_view(server_args)
|
cfg = resolving_view(server_args)
|
||||||
resolved = resolved_view(server_args)
|
resolved = resolved_view(server_args)
|
||||||
prefill_attention_backend, _ = attention_backends_of(resolved_view(server_args))
|
prefill_attention_backend, _ = attention_backends_of(resolved_view(server_args))
|
||||||
|
supports_layout = (
|
||||||
|
cfg.cp_strategy == "zigzag" and prefill_attention_backend == "trtllm_mha"
|
||||||
|
) or (
|
||||||
|
cfg.cp_strategy == "interleave"
|
||||||
|
and prefill_attention_backend == "dsv4"
|
||||||
|
and is_deepseek_v4(model_config_of(server_args).hf_config)
|
||||||
|
)
|
||||||
return (
|
return (
|
||||||
cfg.enable_prefill_cp
|
cfg.enable_prefill_cp
|
||||||
and cfg.pp_size == 1
|
and cfg.pp_size == 1
|
||||||
and resolved.attn_cp_size == cfg.tp_size
|
and resolved.attn_cp_size == cfg.tp_size
|
||||||
and cfg.cp_strategy == "zigzag"
|
and supports_layout
|
||||||
and prefill_attention_backend == "trtllm_mha"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -67,8 +78,12 @@ def enable_cp_bcg_capture(server_args: ServerArgs) -> bool:
|
|||||||
def filter_prefill_cp_bcg_capture_num_tokens(
|
def filter_prefill_cp_bcg_capture_num_tokens(
|
||||||
capture_num_tokens: list[int], server_args: ServerArgs
|
capture_num_tokens: list[int], server_args: ServerArgs
|
||||||
) -> list[int]:
|
) -> list[int]:
|
||||||
"""Keep only token buckets where the zigzag CP strategy can run."""
|
"""Keep only token buckets where the configured CP strategy can run."""
|
||||||
min_num_tokens = resolved_view(server_args).attn_cp_size * 2
|
cfg = resolving_view(server_args)
|
||||||
|
cp_segments_per_token_block = 2 if cfg.cp_strategy == "zigzag" else 1
|
||||||
|
min_num_tokens = (
|
||||||
|
resolved_view(server_args).attn_cp_size * cp_segments_per_token_block
|
||||||
|
)
|
||||||
filtered = [size for size in capture_num_tokens if size >= min_num_tokens]
|
filtered = [size for size in capture_num_tokens if size >= min_num_tokens]
|
||||||
if not filtered:
|
if not filtered:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -96,6 +111,8 @@ class PrefillCPBCGInput:
|
|||||||
|
|
||||||
input_embeds: torch.Tensor
|
input_embeds: torch.Tensor
|
||||||
positions: torch.Tensor
|
positions: torch.Tensor
|
||||||
|
input_ids: Optional[torch.Tensor] = None
|
||||||
|
num_token_non_padded: Optional[torch.Tensor] = None
|
||||||
bucket_local_tokens: Dict[int, int] = field(default_factory=dict)
|
bucket_local_tokens: Dict[int, int] = field(default_factory=dict)
|
||||||
live_local_tokens: int = 0
|
live_local_tokens: int = 0
|
||||||
|
|
||||||
@@ -114,12 +131,22 @@ class PrefillCPBCGInput:
|
|||||||
(runner.max_num_tokens,),
|
(runner.max_num_tokens,),
|
||||||
dtype=torch.int64,
|
dtype=torch.int64,
|
||||||
),
|
),
|
||||||
|
input_ids=torch.zeros((runner.max_num_tokens,), dtype=torch.int64),
|
||||||
|
num_token_non_padded=torch.zeros((), dtype=torch.int32),
|
||||||
)
|
)
|
||||||
|
|
||||||
def required_local_tokens(self, extend_seq_lens: Any) -> Optional[int]:
|
def required_local_tokens(self, extend_seq_lens: Any) -> Optional[int]:
|
||||||
"""Return the aligned CP-local rows required by a live zigzag layout."""
|
"""Return the aligned CP-local rows required by the active layout."""
|
||||||
strategy = get_cp_strategy()
|
strategy = get_cp_strategy()
|
||||||
if not isinstance(strategy, ZigzagCPStrategy) or extend_seq_lens is None:
|
if extend_seq_lens is None:
|
||||||
|
return None
|
||||||
|
if isinstance(strategy, InterleaveCPStrategy):
|
||||||
|
logical_tokens = (
|
||||||
|
sum(int(length) for length in extend_seq_lens) + strategy.cp_size - 1
|
||||||
|
) // strategy.cp_size
|
||||||
|
align_size = get_cp_padding_align_size()
|
||||||
|
return (logical_tokens + align_size - 1) // align_size * align_size
|
||||||
|
if not isinstance(strategy, ZigzagCPStrategy):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
cp_segment_num = strategy.cp_size * 2
|
cp_segment_num = strategy.cp_size * 2
|
||||||
@@ -219,6 +246,7 @@ class PrefillCPBCGInput:
|
|||||||
raw_tokens = int(forward_batch.extend_num_tokens)
|
raw_tokens = int(forward_batch.extend_num_tokens)
|
||||||
global_input_ids = forward_batch.input_ids[:raw_tokens]
|
global_input_ids = forward_batch.input_ids[:raw_tokens]
|
||||||
global_positions = forward_batch.positions[:raw_tokens]
|
global_positions = forward_batch.positions[:raw_tokens]
|
||||||
|
local_input_ids = cp_shard_hidden_states(global_input_ids, forward_batch)
|
||||||
global_input_embeds = runner.model_runner.model.get_input_embeddings()(
|
global_input_embeds = runner.model_runner.model.get_input_embeddings()(
|
||||||
global_input_ids
|
global_input_ids
|
||||||
)
|
)
|
||||||
@@ -249,12 +277,31 @@ class PrefillCPBCGInput:
|
|||||||
|
|
||||||
input_embeds = self.input_embeds[:captured_local_tokens]
|
input_embeds = self.input_embeds[:captured_local_tokens]
|
||||||
positions = self.positions[:captured_local_tokens]
|
positions = self.positions[:captured_local_tokens]
|
||||||
|
assert self.input_ids is not None
|
||||||
|
input_ids = self.input_ids[:captured_local_tokens]
|
||||||
input_embeds.zero_()
|
input_embeds.zero_()
|
||||||
positions.zero_()
|
positions.zero_()
|
||||||
|
input_ids.zero_()
|
||||||
input_embeds[:live_local_tokens].copy_(local_input_embeds)
|
input_embeds[:live_local_tokens].copy_(local_input_embeds)
|
||||||
positions[:live_local_tokens].copy_(local_positions)
|
positions[:live_local_tokens].copy_(local_positions)
|
||||||
|
input_ids[:live_local_tokens].copy_(local_input_ids)
|
||||||
forward_batch.input_embeds = input_embeds
|
forward_batch.input_embeds = input_embeds
|
||||||
forward_batch.positions = positions
|
forward_batch._cp_positions = positions
|
||||||
|
# Keep the global input_ids field intact: the runner uses its length to
|
||||||
|
# select the global capture bucket. The DSV4 body consumes this fixed,
|
||||||
|
# rank-local view for hash routing and MegaMoE.
|
||||||
|
forward_batch._cp_input_ids = input_ids
|
||||||
|
forward_batch.input_ids_global = input_ids
|
||||||
|
if forward_batch.num_token_non_padded is not None:
|
||||||
|
assert self.num_token_non_padded is not None
|
||||||
|
metadata = forward_batch.attn_cp_metadata
|
||||||
|
logical_tokens = (
|
||||||
|
metadata.per_rank_logical_token or metadata.per_rank_actual_token
|
||||||
|
)
|
||||||
|
strategy = get_cp_strategy()
|
||||||
|
assert strategy is not None
|
||||||
|
self.num_token_non_padded.fill_(logical_tokens[strategy.cp_rank])
|
||||||
|
forward_batch.num_token_non_padded = self.num_token_non_padded
|
||||||
self.live_local_tokens = live_local_tokens
|
self.live_local_tokens = live_local_tokens
|
||||||
|
|
||||||
|
|
||||||
@@ -307,10 +354,50 @@ def execute_prefill_cp_bcg(
|
|||||||
static_forward_batch,
|
static_forward_batch,
|
||||||
torch.cuda.current_stream(),
|
torch.cuda.current_stream(),
|
||||||
)
|
)
|
||||||
return model.logits_processor(
|
if aux_hidden_states is not None:
|
||||||
forward_batch.input_ids,
|
if torch.is_tensor(aux_hidden_states):
|
||||||
|
aux_hidden_states = cp_gather_after_forward(
|
||||||
|
aux_hidden_states, static_forward_batch, torch.cuda.current_stream()
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
aux_hidden_states = [
|
||||||
|
cp_gather_after_forward(
|
||||||
|
aux, static_forward_batch, torch.cuda.current_stream()
|
||||||
|
)
|
||||||
|
for aux in aux_hidden_states
|
||||||
|
]
|
||||||
|
hidden_states_before_norm = None
|
||||||
|
if isinstance(hidden_states, tuple):
|
||||||
|
assert len(hidden_states) == 2
|
||||||
|
hidden_states, hidden_states_before_norm = hidden_states
|
||||||
|
|
||||||
|
input_ids = forward_batch.input_ids
|
||||||
|
logits_metadata = forward_batch
|
||||||
|
tail = None
|
||||||
|
language_model = getattr(model, "model", None)
|
||||||
|
if (
|
||||||
|
capture_aux_hidden_states
|
||||||
|
and getattr(language_model, "late_layer_start", None) is not None
|
||||||
|
and forward_batch.forward_mode.is_extend_without_speculative()
|
||||||
|
):
|
||||||
|
tail_metadata = runner.model_runner.attn_backend.tail_forward_metadata
|
||||||
|
tail = tail_metadata.late_layer_tail
|
||||||
|
input_ids = tail.rows(input_ids)
|
||||||
|
logits_metadata = LogitsMetadata.from_forward_batch(forward_batch)
|
||||||
|
logits_metadata.extend_seq_lens = tail.extend_seq_lens
|
||||||
|
logits_metadata.extend_seq_lens_cpu = tail.extend_seq_lens_cpu
|
||||||
|
logits_metadata.extend_logprob_start_lens_cpu = tail.extend_seq_lens_cpu
|
||||||
|
|
||||||
|
output = model.logits_processor(
|
||||||
|
input_ids,
|
||||||
hidden_states,
|
hidden_states,
|
||||||
model.lm_head,
|
model.lm_head,
|
||||||
forward_batch,
|
logits_metadata,
|
||||||
aux_hidden_states,
|
aux_hidden_states,
|
||||||
|
hidden_states_before_norm=(
|
||||||
|
None if aux_hidden_states is not None else hidden_states_before_norm
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
if tail is not None:
|
||||||
|
output.hidden_states_token_indices = tail.token_indices
|
||||||
|
return output
|
||||||
|
|||||||
@@ -701,6 +701,9 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
|
|
||||||
def _get_layer_model_positions(self, forward_batch: ForwardBatch) -> torch.Tensor:
|
def _get_layer_model_positions(self, forward_batch: ForwardBatch) -> torch.Tensor:
|
||||||
"""Mirror outer multimodal wrappers when BCG captures layer_model directly."""
|
"""Mirror outer multimodal wrappers when BCG captures layer_model directly."""
|
||||||
|
cp_positions = getattr(forward_batch, "_cp_positions", None)
|
||||||
|
if cp_positions is not None:
|
||||||
|
return cp_positions
|
||||||
if forward_batch.mrope_positions is None:
|
if forward_batch.mrope_positions is None:
|
||||||
return forward_batch.positions
|
return forward_batch.positions
|
||||||
|
|
||||||
@@ -782,7 +785,9 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
if self._uses_eager_prefill_tail():
|
if self._uses_eager_prefill_tail():
|
||||||
# BCG / Full: capture the transformer body only.
|
# BCG / Full: capture the transformer body only.
|
||||||
positions = self._get_layer_model_positions(forward_batch)
|
positions = self._get_layer_model_positions(forward_batch)
|
||||||
input_ids = forward_batch.input_ids
|
input_ids = getattr(
|
||||||
|
forward_batch, "_cp_input_ids", forward_batch.input_ids
|
||||||
|
)
|
||||||
kwargs = _build_layer_model_forward_kwargs(
|
kwargs = _build_layer_model_forward_kwargs(
|
||||||
self.layer_model, forward_batch, pp_proxy_tensors
|
self.layer_model, forward_batch, pp_proxy_tensors
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2039,7 +2039,10 @@ class MQALayer(MqaAttentionBase):
|
|||||||
if (
|
if (
|
||||||
forward_batch.forward_mode.is_extend()
|
forward_batch.forward_mode.is_extend()
|
||||||
and is_in_breakable_cuda_graph()
|
and is_in_breakable_cuda_graph()
|
||||||
and not getattr(attn_backend, "low_ratio_prefill_graph", False)
|
and (
|
||||||
|
dsa_use_prefill_cp(forward_batch)
|
||||||
|
or not getattr(attn_backend, "low_ratio_prefill_graph", False)
|
||||||
|
)
|
||||||
):
|
):
|
||||||
bcg_deepseek_v4_low_ratio_sources(self, x, q_lora, positions)
|
bcg_deepseek_v4_low_ratio_sources(self, x, q_lora, positions)
|
||||||
else:
|
else:
|
||||||
@@ -4411,11 +4414,18 @@ class DeepseekV4Model(nn.Module):
|
|||||||
)
|
)
|
||||||
if self.engram_hasher is not None:
|
if self.engram_hasher is not None:
|
||||||
if cp_extend:
|
if cp_extend:
|
||||||
# n-gram hashing needs each token's predecessors: hash the whole prompt
|
# N-gram hashing needs each token's predecessors, so hash the
|
||||||
|
# whole prompt before selecting this CP rank's interleaved rows.
|
||||||
|
# The hasher builds request-to-token indices dynamically; keep
|
||||||
|
# that work at an eager break during breakable graph capture.
|
||||||
total = int(forward_batch.attn_cp_metadata.total_seq_lens)
|
total = int(forward_batch.attn_cp_metadata.total_seq_lens)
|
||||||
hash_ids = self.engram_hasher(
|
global_input_ids = forward_batch.input_ids[:total]
|
||||||
forward_batch.input_ids[:total], forward_batch
|
if is_in_breakable_cuda_graph():
|
||||||
)
|
hash_ids = bcg_deepseek_v4_engram_hash_ids(
|
||||||
|
self.engram_hasher, global_input_ids
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
hash_ids = self.engram_hasher(global_input_ids, forward_batch)
|
||||||
parallel = get_parallel()
|
parallel = get_parallel()
|
||||||
hash_ids = hash_ids[parallel.attn_cp_rank :: parallel.attn_cp_size]
|
hash_ids = hash_ids[parallel.attn_cp_rank :: parallel.attn_cp_size]
|
||||||
pad_rows = hidden_states.shape[0] - hash_ids.shape[0]
|
pad_rows = hidden_states.shape[0] - hash_ids.shape[0]
|
||||||
|
|||||||
Reference in New Issue
Block a user