dsv4(npu): support prefill context parallelism with interleave and zigzag (#39427)
This commit is contained in:
@@ -11,6 +11,7 @@ from sglang.srt.arg_groups.overrides import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.runtime_context import get_platform
|
from sglang.srt.runtime_context import get_platform
|
||||||
|
from sglang.srt.utils.common import is_npu
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
@@ -204,9 +205,15 @@ def validate_deepseek_v4_cp(server_args: ServerArgs) -> None:
|
|||||||
if not cfg.enable_prefill_cp:
|
if not cfg.enable_prefill_cp:
|
||||||
return
|
return
|
||||||
|
|
||||||
if cfg.cp_strategy != "interleave":
|
if cfg.cp_strategy not in ("interleave", "zigzag"):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"DeepSeekV4 only supports interleave CP strategy, got {cfg.cp_strategy}"
|
f"DeepSeekV4 only supports interleave/zigzag CP strategy, got {cfg.cp_strategy}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if cfg.cp_strategy == "zigzag" and not is_npu():
|
||||||
|
raise ValueError(
|
||||||
|
"DeepSeekV4 zigzag CP requires the NPU backend; the CUDA backend "
|
||||||
|
"reindexes with interleave order."
|
||||||
)
|
)
|
||||||
|
|
||||||
declare_resolution(
|
declare_resolution(
|
||||||
@@ -224,6 +231,7 @@ def validate_deepseek_v4_cp(server_args: ServerArgs) -> None:
|
|||||||
"validate_deepseek_v4_cp",
|
"validate_deepseek_v4_cp",
|
||||||
attn_cp_size=cfg.tp_size // cfg.dp_size,
|
attn_cp_size=cfg.tp_size // cfg.dp_size,
|
||||||
)
|
)
|
||||||
|
if not is_npu():
|
||||||
assert cfg.dp_size == 1, (
|
assert cfg.dp_size == 1, (
|
||||||
"For round-robin split mode, dp attention is not supported."
|
"For round-robin split mode, dp attention is not supported."
|
||||||
)
|
)
|
||||||
@@ -243,6 +251,7 @@ def validate_deepseek_v4_cp(server_args: ServerArgs) -> None:
|
|||||||
envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.set(False)
|
envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.set(False)
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Enable Context Parallel for DeepSeekV4, "
|
f"Enable Context Parallel for DeepSeekV4, "
|
||||||
|
f"strategy={cfg.cp_strategy}, "
|
||||||
f"dp_size={cfg.dp_size}, moe_dense_tp_size={cfg.moe_dense_tp_size}, "
|
f"dp_size={cfg.dp_size}, moe_dense_tp_size={cfg.moe_dense_tp_size}, "
|
||||||
f"attn_cp_size={cfg.attn_cp_size}, ep_size={cfg.ep_size}, tp_size={cfg.tp_size}"
|
f"attn_cp_size={cfg.attn_cp_size}, ep_size={cfg.ep_size}, tp_size={cfg.tp_size}"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -621,9 +621,7 @@ def validate_prefill_cp_platform(server_args: Any):
|
|||||||
"""Reject deprecated platform CP before resolving models or CP topology."""
|
"""Reject deprecated platform CP before resolving models or CP topology."""
|
||||||
cfg = resolving_view(server_args)
|
cfg = resolving_view(server_args)
|
||||||
platform = get_platform()
|
platform = get_platform()
|
||||||
if cfg.enable_prefill_cp and (
|
if cfg.enable_prefill_cp and (platform.is_hip or platform.is_musa):
|
||||||
platform.is_hip or platform.is_npu or platform.is_musa
|
|
||||||
):
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Prefill CP on HIP/NPU/MUSA is deprecated; CP support will be refactored soon."
|
"Prefill CP on HIP/MUSA is deprecated; CP support will be refactored soon."
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import math
|
import math
|
||||||
|
from contextlib import contextmanager
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from typing import TYPE_CHECKING, Optional
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
@@ -18,6 +19,7 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.srt.hardware_backend.npu.attention.ascend_backend import AscendAttnBackend
|
from sglang.srt.hardware_backend.npu.attention.ascend_backend import AscendAttnBackend
|
||||||
from sglang.srt.hardware_backend.npu.dsv4.dsv4_rope import Dsv4NpuRoPE, rope_cos_sin
|
from sglang.srt.hardware_backend.npu.dsv4.dsv4_rope import Dsv4NpuRoPE, rope_cos_sin
|
||||||
from sglang.srt.hardware_backend.npu.utils import is_npu_arch35
|
from sglang.srt.hardware_backend.npu.utils import is_npu_arch35
|
||||||
|
from sglang.srt.layers.cp.base import get_cp_strategy
|
||||||
from sglang.srt.model_executor.forward_batch_info import DSV4OutCacheLoc, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import DSV4OutCacheLoc, ForwardMode
|
||||||
from sglang.srt.model_executor.forward_context import get_attn_backend
|
from sglang.srt.model_executor.forward_context import get_attn_backend
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
@@ -416,6 +418,9 @@ class CompressorAscendBackendMixin:
|
|||||||
return
|
return
|
||||||
compressor(x, forward_batch)
|
compressor(x, forward_batch)
|
||||||
|
|
||||||
|
# The NPU fused compressor writes c4/indexer payloads identically.
|
||||||
|
forward_indexer_compressor = forward_core_compressor
|
||||||
|
|
||||||
def forward_compress(
|
def forward_compress(
|
||||||
self,
|
self,
|
||||||
compressor,
|
compressor,
|
||||||
@@ -623,10 +628,12 @@ class C4IndexerAscendBackendMixin:
|
|||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
q_lora: torch.Tensor,
|
q_lora: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
|
skip_compressor: bool = False,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
q = self._compute_q_npu(c4_indexer, q_lora, forward_batch)
|
q = self._compute_q_npu(c4_indexer, q_lora, forward_batch)
|
||||||
weights, _ = c4_indexer.weights_proj(x)
|
weights, _ = c4_indexer.weights_proj(x)
|
||||||
weights = weights * (c4_indexer.softmax_scale * c4_indexer.n_heads**-0.5)
|
weights = weights * (c4_indexer.softmax_scale * c4_indexer.n_heads**-0.5)
|
||||||
|
if not skip_compressor:
|
||||||
c4_indexer.compressor(x, forward_batch)
|
c4_indexer.compressor(x, forward_batch)
|
||||||
return q, weights
|
return q, weights
|
||||||
|
|
||||||
@@ -647,6 +654,7 @@ class C4IndexerAscendBackendMixin:
|
|||||||
q_lora: torch.Tensor,
|
q_lora: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
q_lora_ready,
|
q_lora_ready,
|
||||||
|
skip_compressor: bool = False,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
from sglang.srt.hardware_backend.npu.utils import (
|
from sglang.srt.hardware_backend.npu.utils import (
|
||||||
get_indexer_weight_stream,
|
get_indexer_weight_stream,
|
||||||
@@ -661,6 +669,7 @@ class C4IndexerAscendBackendMixin:
|
|||||||
stream_w.wait_stream(cur)
|
stream_w.wait_stream(cur)
|
||||||
|
|
||||||
# route-KV write on cur; ordered before the topk read by cur's program order.
|
# route-KV write on cur; ordered before the topk read by cur's program order.
|
||||||
|
if not skip_compressor:
|
||||||
c4_indexer.compressor(x, forward_batch)
|
c4_indexer.compressor(x, forward_batch)
|
||||||
|
|
||||||
# weights_proj + scale on stream_w.
|
# weights_proj + scale on stream_w.
|
||||||
@@ -794,7 +803,7 @@ class C4IndexerAscendBackendMixin:
|
|||||||
self, c4_indexer, q_lora: torch.Tensor, forward_batch: ForwardBatch
|
self, c4_indexer, q_lora: torch.Tensor, forward_batch: ForwardBatch
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|
||||||
positions = forward_batch.positions
|
positions = self._cp_local_positions(forward_batch)
|
||||||
bs = q_lora.shape[0]
|
bs = q_lora.shape[0]
|
||||||
q, _ = c4_indexer.wq_b(q_lora)
|
q, _ = c4_indexer.wq_b(q_lora)
|
||||||
q = q.view(bs, c4_indexer.n_local_heads, c4_indexer.head_dim)
|
q = q.view(bs, c4_indexer.n_local_heads, c4_indexer.head_dim)
|
||||||
@@ -877,23 +886,51 @@ class C4IndexerAscendBackendMixin:
|
|||||||
) -> None:
|
) -> None:
|
||||||
if forward_batch.forward_mode.is_idle():
|
if forward_batch.forward_mode.is_idle():
|
||||||
return
|
return
|
||||||
assert not skip_compressor, (
|
# skip_compressor=True is the CP full-metadata protocol: the compressor
|
||||||
"skip_compressor=True is not supported on the NPU indexer path"
|
# already ran via forward_indexer_compressor.
|
||||||
)
|
|
||||||
self._ensure_npu_c4_indexer(c4_indexer, x.device)
|
self._ensure_npu_c4_indexer(c4_indexer, x.device)
|
||||||
if self._can_use_indexer_multi_stream():
|
if self._can_use_indexer_multi_stream():
|
||||||
q, weights = self._forward_prepare_multi_stream(
|
q, weights = self._forward_prepare_multi_stream(
|
||||||
c4_indexer, x, q_lora, forward_batch, q_lora_ready
|
c4_indexer, x, q_lora, forward_batch, q_lora_ready, skip_compressor
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
q, weights = self._forward_prepare(c4_indexer, x, q_lora, forward_batch)
|
q, weights = self._forward_prepare(
|
||||||
|
c4_indexer, x, q_lora, forward_batch, skip_compressor
|
||||||
|
)
|
||||||
topk_idxs = self._forward_indexer(c4_indexer, x, q, weights, forward_batch)
|
topk_idxs = self._forward_indexer(c4_indexer, x, q, weights, forward_batch)
|
||||||
self.forward_metadata.c4_topk_indices = topk_idxs
|
self.forward_metadata.c4_topk_indices = topk_idxs
|
||||||
|
|
||||||
|
def _cp_local_positions(self, forward_batch: ForwardBatch) -> torch.Tensor:
|
||||||
|
"""Per-rank positions under CP-v2 (the batch keeps full-length ones)."""
|
||||||
|
local = getattr(forward_batch, "dsv4_cp_local_positions", None)
|
||||||
|
return local if local is not None else forward_batch.positions
|
||||||
|
|
||||||
|
|
||||||
class DeepseekV4AscendAttnBackend(
|
class DeepseekV4AscendAttnBackend(
|
||||||
AscendAttnBackend, C4IndexerAscendBackendMixin, CompressorAscendBackendMixin
|
AscendAttnBackend, C4IndexerAscendBackendMixin, CompressorAscendBackendMixin
|
||||||
):
|
):
|
||||||
|
_DSV4_CP_LOCAL_FIELDS = (
|
||||||
|
"actual_seq_lengths_q",
|
||||||
|
"actual_seq_lengths_q_pa",
|
||||||
|
"actual_seq_lengths_kv",
|
||||||
|
"block_tables",
|
||||||
|
"swa_page_table",
|
||||||
|
"c4_page_table",
|
||||||
|
"c128_page_table",
|
||||||
|
"kernel_metadata",
|
||||||
|
"c4_topk_indices",
|
||||||
|
"positions_cmp_padding_c4",
|
||||||
|
"positions_cmp_padding_c128",
|
||||||
|
"c4_state_page_table",
|
||||||
|
"c128_state_page_table",
|
||||||
|
"c4_loc",
|
||||||
|
"c128_loc",
|
||||||
|
"c4_state_loc",
|
||||||
|
"c128_state_loc",
|
||||||
|
"start_pos",
|
||||||
|
"seqused",
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
model_runner: ModelRunner,
|
model_runner: ModelRunner,
|
||||||
@@ -1028,6 +1065,161 @@ class DeepseekV4AscendAttnBackend(
|
|||||||
)
|
)
|
||||||
return ori_sparse_indices
|
return ori_sparse_indices
|
||||||
|
|
||||||
|
def prepare_dsv4_cp_metadata(self, forward_batch: ForwardBatch) -> None:
|
||||||
|
if getattr(forward_batch, "dsv4_cp_metadata_prepared", False):
|
||||||
|
return
|
||||||
|
if getattr(forward_batch, "attn_cp_metadata", None) is None:
|
||||||
|
return
|
||||||
|
if not forward_batch.forward_mode.is_context_parallel_extend():
|
||||||
|
return
|
||||||
|
if forward_batch.forward_mode.is_target_verify():
|
||||||
|
return
|
||||||
|
|
||||||
|
strategy = get_cp_strategy()
|
||||||
|
if strategy is None or strategy.cp_size <= 1:
|
||||||
|
return
|
||||||
|
|
||||||
|
fm = self.forward_metadata
|
||||||
|
global_positions = forward_batch.positions
|
||||||
|
if global_positions is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
device = global_positions.device
|
||||||
|
num_tokens = int(global_positions.shape[0])
|
||||||
|
local_idx = strategy.local_q_indices(num_tokens, forward_batch).to(
|
||||||
|
device=device, dtype=torch.long
|
||||||
|
)
|
||||||
|
if local_idx.numel() > 0:
|
||||||
|
# Same bound the runner uses to shard model inputs (x[:total_seq_lens]).
|
||||||
|
shard_bound = int(
|
||||||
|
getattr(forward_batch.attn_cp_metadata, "total_seq_lens", num_tokens)
|
||||||
|
)
|
||||||
|
local_idx = local_idx[local_idx < shard_bound]
|
||||||
|
local_positions = global_positions.index_select(0, local_idx)
|
||||||
|
# Sharded model inputs are padded to per_rank_actual_token; pad rows get
|
||||||
|
# position 0 so per-token metadata matches the sharded q length.
|
||||||
|
from sglang.srt.layers.cp.padding import pad_local_rows
|
||||||
|
|
||||||
|
local_positions = pad_local_rows(
|
||||||
|
local_positions, forward_batch.attn_cp_metadata, dim=0
|
||||||
|
)
|
||||||
|
|
||||||
|
extend_lens = getattr(forward_batch, "extend_seq_lens_cpu", None)
|
||||||
|
if extend_lens is None:
|
||||||
|
seq_lens_cpu = getattr(forward_batch, "seq_lens_cpu", None)
|
||||||
|
if seq_lens_cpu is not None:
|
||||||
|
extend_lens = seq_lens_cpu.tolist()
|
||||||
|
else:
|
||||||
|
extend_lens = [num_tokens]
|
||||||
|
extend_lens = [int(x) for x in extend_lens]
|
||||||
|
real_num_tokens = min(sum(extend_lens), num_tokens)
|
||||||
|
|
||||||
|
batch_ids_parts = []
|
||||||
|
for batch_id, length in enumerate(extend_lens):
|
||||||
|
if length <= 0:
|
||||||
|
continue
|
||||||
|
batch_ids_parts.append(
|
||||||
|
torch.full((length,), batch_id, dtype=torch.long, device=device)
|
||||||
|
)
|
||||||
|
if batch_ids_parts:
|
||||||
|
batch_ids = torch.cat(batch_ids_parts, dim=0)
|
||||||
|
else:
|
||||||
|
batch_ids = torch.empty(0, dtype=torch.long, device=device)
|
||||||
|
if batch_ids.shape[0] < num_tokens:
|
||||||
|
pad_len = num_tokens - batch_ids.shape[0]
|
||||||
|
batch_ids = torch.cat(
|
||||||
|
[batch_ids, torch.zeros(pad_len, dtype=torch.long, device=device)],
|
||||||
|
dim=0,
|
||||||
|
)
|
||||||
|
elif batch_ids.shape[0] > num_tokens:
|
||||||
|
batch_ids = batch_ids[:num_tokens]
|
||||||
|
|
||||||
|
local_batch_ids = batch_ids.index_select(0, local_idx)
|
||||||
|
valid_rows = local_idx < real_num_tokens
|
||||||
|
# Pad rows reuse request 0's page table (in-bounds) and stay invalid.
|
||||||
|
pad_rows = local_positions.shape[0] - local_batch_ids.shape[0]
|
||||||
|
if pad_rows > 0:
|
||||||
|
local_batch_ids = torch.cat(
|
||||||
|
[local_batch_ids, local_batch_ids.new_zeros(pad_rows)]
|
||||||
|
)
|
||||||
|
valid_rows = torch.cat([valid_rows, valid_rows.new_zeros(pad_rows)])
|
||||||
|
seqused_kv = torch.where(
|
||||||
|
valid_rows,
|
||||||
|
local_positions.to(torch.int32) + 1,
|
||||||
|
torch.ones_like(local_positions, dtype=torch.int32),
|
||||||
|
).clamp(min=1)
|
||||||
|
|
||||||
|
full_fields = {
|
||||||
|
field: getattr(fm, field, None) for field in self._DSV4_CP_LOCAL_FIELDS
|
||||||
|
}
|
||||||
|
setattr(fm, "dsv4_cp_full_fields", full_fields)
|
||||||
|
|
||||||
|
def _select_rows(table: Optional[torch.Tensor]):
|
||||||
|
if table is None:
|
||||||
|
return None
|
||||||
|
if local_batch_ids.numel() == 0:
|
||||||
|
return table.new_empty((0, *table.shape[1:]))
|
||||||
|
return table.index_select(0, local_batch_ids)
|
||||||
|
|
||||||
|
fm.block_tables = _select_rows(full_fields["block_tables"])
|
||||||
|
fm.swa_page_table = _select_rows(full_fields["swa_page_table"])
|
||||||
|
if self._dsv4_has_c4:
|
||||||
|
fm.c4_page_table = _select_rows(full_fields["c4_page_table"])
|
||||||
|
if self._dsv4_has_c128:
|
||||||
|
fm.c128_page_table = _select_rows(full_fields["c128_page_table"])
|
||||||
|
|
||||||
|
local_t = int(local_positions.shape[0])
|
||||||
|
fm.actual_seq_lengths_q = torch.arange(
|
||||||
|
1, local_t + 1, dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
fm.actual_seq_lengths_q_pa = torch.arange(
|
||||||
|
0, local_t + 1, dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
fm.actual_seq_lengths_kv = seqused_kv
|
||||||
|
fm.kernel_metadata = self._kernel_metadata_from_parts(
|
||||||
|
bs=local_t,
|
||||||
|
actual_seq_lengths_q_pa=fm.actual_seq_lengths_q_pa,
|
||||||
|
actual_seq_lengths_kv=fm.actual_seq_lengths_kv,
|
||||||
|
block_tables=fm.block_tables,
|
||||||
|
max_seqlen_q=1,
|
||||||
|
is_nextn=False,
|
||||||
|
)
|
||||||
|
if self._dsv4_has_c4:
|
||||||
|
fm.c4_topk_indices = torch.full(
|
||||||
|
(local_t, self._dsv4_index_topk),
|
||||||
|
-1,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
|
||||||
|
forward_batch.dsv4_cp_metadata_prepared = True
|
||||||
|
forward_batch.dsv4_cp_global_positions = global_positions
|
||||||
|
forward_batch.dsv4_cp_local_positions = local_positions
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def use_dsv4_cp_full_metadata(self, forward_batch: ForwardBatch):
|
||||||
|
fm = self.forward_metadata
|
||||||
|
full_fields = getattr(fm, "dsv4_cp_full_fields", None)
|
||||||
|
if not full_fields:
|
||||||
|
yield
|
||||||
|
return
|
||||||
|
|
||||||
|
local_fields = {
|
||||||
|
field: getattr(fm, field, None) for field in self._DSV4_CP_LOCAL_FIELDS
|
||||||
|
}
|
||||||
|
previous_positions = getattr(forward_batch, "positions", None)
|
||||||
|
try:
|
||||||
|
for field, value in full_fields.items():
|
||||||
|
setattr(fm, field, value)
|
||||||
|
global_positions = getattr(forward_batch, "dsv4_cp_global_positions", None)
|
||||||
|
if global_positions is not None:
|
||||||
|
forward_batch.positions = global_positions
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
for field, value in local_fields.items():
|
||||||
|
setattr(fm, field, value)
|
||||||
|
forward_batch.positions = previous_positions
|
||||||
|
|
||||||
def _init_dsv4_graph_buffers(self, *, max_bs: int, max_num_tokens: int) -> None:
|
def _init_dsv4_graph_buffers(self, *, max_bs: int, max_num_tokens: int) -> None:
|
||||||
device = self.device
|
device = self.device
|
||||||
block_tables_shape = self.graph_metadata["block_tables"].shape
|
block_tables_shape = self.graph_metadata["block_tables"].shape
|
||||||
|
|||||||
@@ -6,8 +6,8 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.communicator import ScatterMode
|
from sglang.srt.layers.communicator import ScatterMode
|
||||||
|
from sglang.srt.layers.cp.utils import cp_gather_full_sequence_states
|
||||||
from sglang.srt.layers.dp_attention import attn_tp_all_gather_into_tensor
|
from sglang.srt.layers.dp_attention import attn_tp_all_gather_into_tensor
|
||||||
from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output
|
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_executor.forward_context import (
|
from sglang.srt.model_executor.forward_context import (
|
||||||
get_attn_backend,
|
get_attn_backend,
|
||||||
@@ -206,9 +206,8 @@ class DSANPUIndexerMixin:
|
|||||||
and self.dsa_enable_prefill_cp
|
and self.dsa_enable_prefill_cp
|
||||||
and forward_batch.attn_cp_metadata is not None
|
and forward_batch.attn_cp_metadata is not None
|
||||||
):
|
):
|
||||||
k = cp_all_gather_rerange_output(
|
k = cp_gather_full_sequence_states(
|
||||||
k.contiguous().view(-1, self.head_dim),
|
k.contiguous().view(-1, self.head_dim),
|
||||||
self.cp_size,
|
|
||||||
forward_batch,
|
forward_batch,
|
||||||
torch.npu.current_stream(),
|
torch.npu.current_stream(),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_parallel,
|
get_parallel,
|
||||||
process_model_config,
|
process_model_config,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip, is_musa, is_npu
|
from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip, is_musa
|
||||||
from sglang.srt.utils.common import ceil_div
|
from sglang.srt.utils.common import ceil_div
|
||||||
|
|
||||||
|
|
||||||
@@ -115,7 +115,7 @@ def should_use_dsa_fused_topk(seed_dsa_topk_from_draft_extend: bool) -> bool:
|
|||||||
|
|
||||||
|
|
||||||
def is_dsa_enable_prefill_cp():
|
def is_dsa_enable_prefill_cp():
|
||||||
if is_hip() or is_npu() or is_musa():
|
if is_hip() or is_musa():
|
||||||
return False
|
return False
|
||||||
|
|
||||||
# Generic prefill CP derives activation from the runtime topology and model
|
# Generic prefill CP derives activation from the runtime topology and model
|
||||||
|
|||||||
@@ -145,6 +145,11 @@ class ContextParallelStrategy(ABC):
|
|||||||
) -> Any:
|
) -> Any:
|
||||||
"""Gather rank-local KV payloads back to full token order."""
|
"""Gather rank-local KV payloads back to full token order."""
|
||||||
|
|
||||||
|
def local_q_indices(self, num_tokens: int, forward_batch: ForwardBatch) -> Any:
|
||||||
|
raise NotImplementedError(
|
||||||
|
f"{self.name} strategy does not support local q indices"
|
||||||
|
)
|
||||||
|
|
||||||
def shard_per_request(
|
def shard_per_request(
|
||||||
self,
|
self,
|
||||||
extend_seqs_cpu: List[int],
|
extend_seqs_cpu: List[int],
|
||||||
|
|||||||
@@ -122,6 +122,14 @@ class InterleaveCPStrategy(ContextParallelStrategy):
|
|||||||
|
|
||||||
return input_.view(-1, cp_size, *input_.shape[1:])[:, cp_rank].contiguous()
|
return input_.view(-1, cp_size, *input_.shape[1:])[:, cp_rank].contiguous()
|
||||||
|
|
||||||
|
def local_q_indices(self, num_tokens: int, forward_batch) -> Any:
|
||||||
|
device = getattr(getattr(forward_batch, "input_ids", None), "device", None)
|
||||||
|
if device is None:
|
||||||
|
device = torch.device("cpu")
|
||||||
|
return torch.arange(
|
||||||
|
self.cp_rank, int(num_tokens), self.cp_size, device=device, dtype=torch.long
|
||||||
|
)
|
||||||
|
|
||||||
def shard_local_tokens(self, input_: Any) -> Any:
|
def shard_local_tokens(self, input_: Any) -> Any:
|
||||||
return self._interleave_shard(input_)
|
return self._interleave_shard(input_)
|
||||||
|
|
||||||
@@ -205,6 +213,18 @@ class InterleaveCPStrategy(ContextParallelStrategy):
|
|||||||
gathered = x.new_empty((self.cp_size * physical_rank_len, *x.shape[1:]))
|
gathered = x.new_empty((self.cp_size * physical_rank_len, *x.shape[1:]))
|
||||||
attn_cp_all_gather_into_tensor(gathered, padded_x.contiguous())
|
attn_cp_all_gather_into_tensor(gathered, padded_x.contiguous())
|
||||||
|
|
||||||
|
# Equal per-rank lengths: one interleave copy restores the original
|
||||||
|
# token order; cheaper than the index_select fallback below.
|
||||||
|
actual = metadata.per_rank_actual_token
|
||||||
|
if total_tokens == self.cp_size * physical_rank_len and all(
|
||||||
|
int(n) == physical_rank_len for n in actual
|
||||||
|
):
|
||||||
|
return (
|
||||||
|
gathered.view(self.cp_size, physical_rank_len, *x.shape[1:])
|
||||||
|
.transpose(0, 1)
|
||||||
|
.reshape(total_tokens, *x.shape[1:])
|
||||||
|
)
|
||||||
|
|
||||||
flat_indices = torch.arange(total_tokens, device=x.device)
|
flat_indices = torch.arange(total_tokens, device=x.device)
|
||||||
gather_indices = (
|
gather_indices = (
|
||||||
flat_indices % self.cp_size
|
flat_indices % self.cp_size
|
||||||
|
|||||||
@@ -201,6 +201,16 @@ def cp_shard_hidden_states(complete_hidden_states: Any, forward_batch):
|
|||||||
return strategy.shard_hidden_states(complete_hidden_states, forward_batch)
|
return strategy.shard_hidden_states(complete_hidden_states, forward_batch)
|
||||||
|
|
||||||
|
|
||||||
|
def cp_gather_full_sequence_states(sharded_states, forward_batch, stream=None):
|
||||||
|
"""Gather CP-sharded states into the full logical sequence order."""
|
||||||
|
assert is_cp_active(forward_batch)
|
||||||
|
strategy = get_cp_strategy()
|
||||||
|
assert strategy is not None
|
||||||
|
assert sharded_states is not None
|
||||||
|
assert getattr(forward_batch, "attn_cp_metadata", None) is not None
|
||||||
|
return strategy.gather_hidden_states(sharded_states, forward_batch, stream)
|
||||||
|
|
||||||
|
|
||||||
def cp_shard_position_ids(complete_position_ids: Any, forward_batch):
|
def cp_shard_position_ids(complete_position_ids: Any, forward_batch):
|
||||||
assert is_cp_active(forward_batch)
|
assert is_cp_active(forward_batch)
|
||||||
strategy = get_cp_strategy()
|
strategy = get_cp_strategy()
|
||||||
|
|||||||
@@ -334,6 +334,25 @@ class ZigzagCPStrategy(ContextParallelStrategy):
|
|||||||
[chunks[i] for i in forward_batch.attn_cp_metadata.cp_reverse_index], dim=0
|
[chunks[i] for i in forward_batch.attn_cp_metadata.cp_reverse_index], dim=0
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def local_q_indices(self, num_tokens: int, forward_batch) -> Any:
|
||||||
|
meta = forward_batch.attn_cp_metadata
|
||||||
|
device = getattr(getattr(forward_batch, "input_ids", None), "device", None)
|
||||||
|
if device is None and meta.cu_seqlens_q_prev_tensor is not None:
|
||||||
|
device = meta.cu_seqlens_q_prev_tensor.device
|
||||||
|
if device is None:
|
||||||
|
device = torch.device("cpu")
|
||||||
|
|
||||||
|
offsets = [0] + list(accumulate(meta.split_list))
|
||||||
|
pieces = []
|
||||||
|
for chunk_idx in meta.zigzag_index:
|
||||||
|
start = offsets[chunk_idx]
|
||||||
|
end = offsets[chunk_idx + 1]
|
||||||
|
if end > start:
|
||||||
|
pieces.append(torch.arange(start, end, device=device, dtype=torch.long))
|
||||||
|
if not pieces:
|
||||||
|
return torch.empty(0, device=device, dtype=torch.long)
|
||||||
|
return torch.cat(pieces, dim=0)
|
||||||
|
|
||||||
def get_supported_attention_backend(self):
|
def get_supported_attention_backend(self):
|
||||||
return [
|
return [
|
||||||
CPAttentionBackendKind.FLASH_ATTENTION,
|
CPAttentionBackendKind.FLASH_ATTENTION,
|
||||||
|
|||||||
@@ -67,6 +67,7 @@ from sglang.srt.layers.communicator_dsa_cp import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers.cp.cp_decode_attn_tp import get_cp_decode_attn_tp_ctx
|
from sglang.srt.layers.cp.cp_decode_attn_tp import get_cp_decode_attn_tp_ctx
|
||||||
from sglang.srt.layers.cp.utils import (
|
from sglang.srt.layers.cp.utils import (
|
||||||
|
cp_gather_full_sequence_states,
|
||||||
cp_materialize_global_token_order,
|
cp_materialize_global_token_order,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
@@ -1633,9 +1634,16 @@ class MQALayer(MqaAttentionBase):
|
|||||||
sin4,
|
sin4,
|
||||||
qk_nope_dim=self.qk_nope_head_dim,
|
qk_nope_dim=self.qk_nope_head_dim,
|
||||||
)
|
)
|
||||||
|
kv_for_cache = kv
|
||||||
|
if use_cp:
|
||||||
|
kv_for_cache = cp_gather_full_sequence_states(
|
||||||
|
kv.contiguous(),
|
||||||
|
forward_batch,
|
||||||
|
torch.cuda.current_stream(),
|
||||||
|
)
|
||||||
attn_backend.store_cache(
|
attn_backend.store_cache(
|
||||||
layer_id=self.layer_id,
|
layer_id=self.layer_id,
|
||||||
swa_k=kv,
|
swa_k=kv_for_cache,
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
)
|
)
|
||||||
kv = None
|
kv = None
|
||||||
@@ -1670,7 +1678,24 @@ class MQALayer(MqaAttentionBase):
|
|||||||
|
|
||||||
del qkv_a
|
del qkv_a
|
||||||
|
|
||||||
|
use_npu_cp_full_metadata = use_cp and _is_npu
|
||||||
if self.indexer is not None:
|
if self.indexer is not None:
|
||||||
|
if use_npu_cp_full_metadata:
|
||||||
|
with attn_backend.use_dsv4_cp_full_metadata(forward_batch):
|
||||||
|
attn_backend.forward_indexer_compressor(
|
||||||
|
x,
|
||||||
|
forward_batch,
|
||||||
|
self.indexer.layer_id,
|
||||||
|
self.indexer.compressor,
|
||||||
|
)
|
||||||
|
self.indexer(
|
||||||
|
x=x,
|
||||||
|
q_lora=q_lora,
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
attn_backend=attn_backend,
|
||||||
|
skip_compressor=True,
|
||||||
|
)
|
||||||
|
else:
|
||||||
self.indexer(
|
self.indexer(
|
||||||
x=x,
|
x=x,
|
||||||
q_lora=q_lora,
|
q_lora=q_lora,
|
||||||
@@ -1678,6 +1703,15 @@ class MQALayer(MqaAttentionBase):
|
|||||||
attn_backend=attn_backend,
|
attn_backend=attn_backend,
|
||||||
)
|
)
|
||||||
if self.compressor is not None:
|
if self.compressor is not None:
|
||||||
|
if use_npu_cp_full_metadata:
|
||||||
|
with attn_backend.use_dsv4_cp_full_metadata(forward_batch):
|
||||||
|
attn_backend.forward_core_compressor(
|
||||||
|
x,
|
||||||
|
forward_batch,
|
||||||
|
self.layer_id,
|
||||||
|
self.compressor,
|
||||||
|
)
|
||||||
|
else:
|
||||||
attn_backend.forward_core_compressor(
|
attn_backend.forward_core_compressor(
|
||||||
x,
|
x,
|
||||||
forward_batch,
|
forward_batch,
|
||||||
@@ -3209,6 +3243,7 @@ class DeepseekV4Model(nn.Module):
|
|||||||
)
|
)
|
||||||
return hidden_states
|
return hidden_states
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
input_ids: torch.Tensor,
|
input_ids: torch.Tensor,
|
||||||
@@ -3250,9 +3285,16 @@ class DeepseekV4Model(nn.Module):
|
|||||||
|
|
||||||
capture_dspark = self.dspark_layers_to_capture is not None
|
capture_dspark = self.dspark_layers_to_capture is not None
|
||||||
dspark_aux_hidden_states: List[torch.Tensor] = []
|
dspark_aux_hidden_states: List[torch.Tensor] = []
|
||||||
# DSpark aux capture needs the per-layer eager loop (TBO's overlapped
|
|
||||||
# execution cannot expose per-layer completed hidden states), so skip
|
attn_backend = get_attn_backend()
|
||||||
# TBO when capturing -- a perf-only downgrade, not a correctness one.
|
if _is_npu and forward_batch.attn_cp_metadata is not None:
|
||||||
|
attn_backend.prepare_dsv4_cp_metadata(forward_batch)
|
||||||
|
local_positions = getattr(forward_batch, "dsv4_cp_local_positions", None)
|
||||||
|
if (
|
||||||
|
local_positions is not None
|
||||||
|
and positions.shape[0] == local_positions.shape[0]
|
||||||
|
):
|
||||||
|
forward_batch.positions = positions
|
||||||
|
|
||||||
# Reset Compressor's per-step freqs_cis cache from any previous step.
|
# Reset Compressor's per-step freqs_cis cache from any previous step.
|
||||||
for _attr in ("freqs_cis_c4", "freqs_cis_c128"):
|
for _attr in ("freqs_cis_c4", "freqs_cis_c128"):
|
||||||
@@ -3449,7 +3491,6 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
0 if is_shared_experts_fusion_disabled() else self.config.n_shared_experts
|
0 if is_shared_experts_fusion_disabled() else self.config.n_shared_experts
|
||||||
)
|
)
|
||||||
|
|
||||||
@torch.no_grad()
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
input_ids: torch.Tensor,
|
input_ids: torch.Tensor,
|
||||||
@@ -3458,7 +3499,6 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
input_embeds: Optional[torch.Tensor] = None,
|
input_embeds: Optional[torch.Tensor] = None,
|
||||||
pp_proxy_tensors: Optional[PPProxyTensors] = None,
|
pp_proxy_tensors: Optional[PPProxyTensors] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|
||||||
with get_attn_tp_context().maybe_input_scattered(forward_batch):
|
with get_attn_tp_context().maybe_input_scattered(forward_batch):
|
||||||
hidden_states = self.model.forward(
|
hidden_states = self.model.forward(
|
||||||
input_ids, positions, forward_batch, input_embeds, pp_proxy_tensors
|
input_ids, positions, forward_batch, input_embeds, pp_proxy_tensors
|
||||||
|
|||||||
@@ -8,6 +8,12 @@ from transformers import PretrainedConfig
|
|||||||
|
|
||||||
from sglang.srt.distributed import get_pp_group
|
from sglang.srt.distributed import get_pp_group
|
||||||
from sglang.srt.hardware_backend.npu.dsv4.dsv4_rope import prime_rope_cos_sin
|
from sglang.srt.hardware_backend.npu.dsv4.dsv4_rope import prime_rope_cos_sin
|
||||||
|
from sglang.srt.layers.attention.dsa.utils import (
|
||||||
|
dsa_use_prefill_cp,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.cp.utils import (
|
||||||
|
is_cp_active,
|
||||||
|
)
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
dp_gather_replicate,
|
dp_gather_replicate,
|
||||||
get_global_dp_buffer_len,
|
get_global_dp_buffer_len,
|
||||||
@@ -24,6 +30,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
VocabParallelEmbedding,
|
VocabParallelEmbedding,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
from sglang.srt.model_executor.forward_context import get_attn_backend
|
||||||
from sglang.srt.models.deepseek_v4 import (
|
from sglang.srt.models.deepseek_v4 import (
|
||||||
DeepseekV4DecoderLayer,
|
DeepseekV4DecoderLayer,
|
||||||
DeepseekV4ForCausalLM,
|
DeepseekV4ForCausalLM,
|
||||||
@@ -160,6 +167,20 @@ class DeepseekV4ModelNextN(nn.Module):
|
|||||||
else:
|
else:
|
||||||
input_ids_global = getattr(forward_batch, "input_ids_global", input_ids)
|
input_ids_global = getattr(forward_batch, "input_ids_global", input_ids)
|
||||||
|
|
||||||
|
use_prefill_cp = dsa_use_prefill_cp(forward_batch)
|
||||||
|
if use_prefill_cp and is_cp_active(forward_batch):
|
||||||
|
attn_backend = get_attn_backend()
|
||||||
|
if hasattr(attn_backend, "prepare_dsv4_cp_metadata"):
|
||||||
|
attn_backend.prepare_dsv4_cp_metadata(forward_batch)
|
||||||
|
local_positions = getattr(
|
||||||
|
forward_batch, "dsv4_cp_local_positions", None
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
local_positions is not None
|
||||||
|
and positions.shape[0] == local_positions.shape[0]
|
||||||
|
):
|
||||||
|
forward_batch.positions = positions
|
||||||
|
|
||||||
if _is_npu:
|
if _is_npu:
|
||||||
# Same per-forward rope prime as DeepseekV4Model.forward: the
|
# Same per-forward rope prime as DeepseekV4Model.forward: the
|
||||||
# decoder layer reads the memoized gather instead of re-gathering.
|
# decoder layer reads the memoized gather instead of re-gathering.
|
||||||
@@ -220,7 +241,6 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM):
|
|||||||
positions: torch.Tensor,
|
positions: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|
||||||
hidden_states, pre_hc_head = self.model(input_ids, positions, forward_batch)
|
hidden_states, pre_hc_head = self.model(input_ids, positions, forward_batch)
|
||||||
return self.logits_processor(
|
return self.logits_processor(
|
||||||
input_ids,
|
input_ids,
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ register_cpu_ci(est_time=2, suite="base-a-test-cpu")
|
|||||||
|
|
||||||
class TestPlatformPrefillCPDeprecation(CustomTestCase):
|
class TestPlatformPrefillCPDeprecation(CustomTestCase):
|
||||||
def test_platform_cp_rejected_before_model_lookup(self):
|
def test_platform_cp_rejected_before_model_lookup(self):
|
||||||
for platform in ("is_hip", "is_npu", "is_musa"):
|
for platform in ("is_hip", "is_musa"):
|
||||||
facts = dict(is_hip=False, is_npu=False, is_musa=False)
|
facts = dict(is_hip=False, is_npu=False, is_musa=False)
|
||||||
facts[platform] = True
|
facts[platform] = True
|
||||||
for strategy in (None, "zigzag", "interleave"):
|
for strategy in (None, "zigzag", "interleave"):
|
||||||
@@ -31,7 +31,7 @@ class TestPlatformPrefillCPDeprecation(CustomTestCase):
|
|||||||
validate_prefill_cp_platform(args)
|
validate_prefill_cp_platform(args)
|
||||||
|
|
||||||
def test_context_parallel_handler_rejects_before_model_lookup(self):
|
def test_context_parallel_handler_rejects_before_model_lookup(self):
|
||||||
for platform in ("is_hip", "is_npu", "is_musa"):
|
for platform in ("is_hip", "is_musa"):
|
||||||
facts = dict(is_hip=False, is_npu=False, is_musa=False)
|
facts = dict(is_hip=False, is_npu=False, is_musa=False)
|
||||||
facts[platform] = True
|
facts[platform] = True
|
||||||
with self.subTest(platform=platform), override_platform(**facts):
|
with self.subTest(platform=platform), override_platform(**facts):
|
||||||
@@ -44,7 +44,7 @@ class TestPlatformPrefillCPDeprecation(CustomTestCase):
|
|||||||
handle_context_parallelism(args)
|
handle_context_parallelism(args)
|
||||||
|
|
||||||
def test_resolution_rejects_even_dummy_models(self):
|
def test_resolution_rejects_even_dummy_models(self):
|
||||||
for platform in ("is_hip", "is_npu", "is_musa"):
|
for platform in ("is_hip", "is_musa"):
|
||||||
facts = dict(is_hip=False, is_npu=False, is_musa=False)
|
facts = dict(is_hip=False, is_npu=False, is_musa=False)
|
||||||
facts[platform] = True
|
facts[platform] = True
|
||||||
for model_path in ("dummy", "none", "missing-model-must-not-be-loaded"):
|
for model_path in ("dummy", "none", "missing-model-must-not-be-loaded"):
|
||||||
@@ -59,7 +59,7 @@ class TestPlatformPrefillCPDeprecation(CustomTestCase):
|
|||||||
args.resolve_once()
|
args.resolve_once()
|
||||||
|
|
||||||
def test_non_cp_and_decode_cp_are_not_rejected(self):
|
def test_non_cp_and_decode_cp_are_not_rejected(self):
|
||||||
for platform in ("is_hip", "is_npu", "is_musa"):
|
for platform in ("is_hip", "is_musa"):
|
||||||
facts = dict(is_hip=False, is_npu=False, is_musa=False)
|
facts = dict(is_hip=False, is_npu=False, is_musa=False)
|
||||||
facts[platform] = True
|
facts[platform] = True
|
||||||
for dcp_size in (1, 2):
|
for dcp_size in (1, 2):
|
||||||
|
|||||||
Reference in New Issue
Block a user