[NPU] DeepSeek-V4 adapt sgl-kernel-npu ops (compressor/sparse-attn/sparse-attn-metadata) (#35676)
Co-authored-by: vstone-w <374330057@qq.com> Co-authored-by: unclezhou486 <154310456@qq.com> Co-authored-by: 摆渡人 <2044145178@qq.com>
This commit is contained in:
co-authored by
vstone-w
unclezhou486
摆渡人
parent
61b67316d8
commit
e2b50930b9
@@ -89,6 +89,9 @@ class ForwardMetadata:
|
|||||||
seq_lens: Optional[torch.Tensor] = None
|
seq_lens: Optional[torch.Tensor] = None
|
||||||
actual_seq_lengths_q: Optional[torch.Tensor] = None
|
actual_seq_lengths_q: Optional[torch.Tensor] = None
|
||||||
actual_seq_lengths_q_pa: Optional[torch.Tensor] = None
|
actual_seq_lengths_q_pa: Optional[torch.Tensor] = None
|
||||||
|
# CPU mirror of actual_seq_lengths_q_pa for the host metadata op
|
||||||
|
# (torch.ops.npu.sparse_attn_sharedkv_metadata_host reads CPU int32 inputs).
|
||||||
|
actual_seq_lengths_q_pa_cpu: Optional[torch.Tensor] = None
|
||||||
actual_seq_lengths_kv: Optional[torch.Tensor] = None
|
actual_seq_lengths_kv: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# swa attention mask for graph mode decode
|
# swa attention mask for graph mode decode
|
||||||
|
|||||||
@@ -398,7 +398,7 @@ class CompressorAscendBackendMixin:
|
|||||||
allow_build=False,
|
allow_build=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
cmp_kv = torch.ops.custom.compressor(
|
cmp_kv = torch.ops.npu.compressor(
|
||||||
x,
|
x,
|
||||||
compressor._fused_wkv_w,
|
compressor._fused_wkv_w,
|
||||||
compressor._fused_wgate_w,
|
compressor._fused_wgate_w,
|
||||||
@@ -989,6 +989,14 @@ class DeepseekV4AscendAttnBackend(
|
|||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=device,
|
device=device,
|
||||||
)
|
)
|
||||||
|
# q_pa is constant per graph shape (never rewritten at replay), so the
|
||||||
|
# CPU mirror for the host metadata op is built once here.
|
||||||
|
metadata.actual_seq_lengths_q_pa_cpu = torch.arange(
|
||||||
|
0,
|
||||||
|
bs * tokens_per_req + tokens_per_req,
|
||||||
|
tokens_per_req,
|
||||||
|
dtype=torch.int32,
|
||||||
|
)
|
||||||
|
|
||||||
# init >=1 so the captured kernel records valid attention work; replay overwrites in-place
|
# init >=1 so the captured kernel records valid attention work; replay overwrites in-place
|
||||||
metadata.actual_seq_lengths_kv = torch.ones(
|
metadata.actual_seq_lengths_kv = torch.ones(
|
||||||
@@ -1189,6 +1197,10 @@ class DeepseekV4AscendAttnBackend(
|
|||||||
fm.seq_lens_cpu_int,
|
fm.seq_lens_cpu_int,
|
||||||
ctx.live_seq_lens_cpu.int(),
|
ctx.live_seq_lens_cpu.int(),
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
# CPU mirror of the kv buffer written below, from its CPU source — the
|
||||||
|
# host metadata op reads this instead of a D2H sync.
|
||||||
|
fm.seq_lens_cpu_int = ctx.final_seq_lens_cpu[: ctx.bs].int().clamp(min=1)
|
||||||
fm.actual_seq_lengths_kv.copy_(attn_seq_lens.clamp(min=1))
|
fm.actual_seq_lengths_kv.copy_(attn_seq_lens.clamp(min=1))
|
||||||
|
|
||||||
def _refresh_graph_compress_page_tables_direct(self, ctx) -> None:
|
def _refresh_graph_compress_page_tables_direct(self, ctx) -> None:
|
||||||
@@ -1432,6 +1444,13 @@ class DeepseekV4AscendAttnBackend(
|
|||||||
[torch.zeros(1, dtype=torch.int32, device=device), actual_q],
|
[torch.zeros(1, dtype=torch.int32, device=device), actual_q],
|
||||||
dim=0,
|
dim=0,
|
||||||
)
|
)
|
||||||
|
fm.actual_seq_lengths_q_pa_cpu = torch.cat(
|
||||||
|
[
|
||||||
|
torch.zeros(1, dtype=torch.int32),
|
||||||
|
torch.cumsum(seq_lens_cpu, dim=0).int(),
|
||||||
|
],
|
||||||
|
dim=0,
|
||||||
|
)
|
||||||
elif forward_batch.forward_mode.is_decode():
|
elif forward_batch.forward_mode.is_decode():
|
||||||
B = forward_batch.batch_size
|
B = forward_batch.batch_size
|
||||||
fm.actual_seq_lengths_q = torch.arange(
|
fm.actual_seq_lengths_q = torch.arange(
|
||||||
@@ -1440,6 +1459,7 @@ class DeepseekV4AscendAttnBackend(
|
|||||||
fm.actual_seq_lengths_q_pa = torch.arange(
|
fm.actual_seq_lengths_q_pa = torch.arange(
|
||||||
0, B + 1, dtype=torch.int32, device=device
|
0, B + 1, dtype=torch.int32, device=device
|
||||||
)
|
)
|
||||||
|
fm.actual_seq_lengths_q_pa_cpu = torch.arange(0, B + 1, dtype=torch.int32)
|
||||||
elif (
|
elif (
|
||||||
forward_batch.forward_mode.is_target_verify()
|
forward_batch.forward_mode.is_target_verify()
|
||||||
or forward_batch.forward_mode.is_draft_extend_v2()
|
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
@@ -1458,9 +1478,17 @@ class DeepseekV4AscendAttnBackend(
|
|||||||
[torch.zeros(1, dtype=torch.int32, device=device), actual_q],
|
[torch.zeros(1, dtype=torch.int32, device=device), actual_q],
|
||||||
dim=0,
|
dim=0,
|
||||||
)
|
)
|
||||||
|
fm.actual_seq_lengths_q_pa_cpu = torch.cat(
|
||||||
|
[
|
||||||
|
torch.zeros(1, dtype=torch.int32),
|
||||||
|
torch.arange(n_draft, B * n_draft + 1, n_draft, dtype=torch.int32),
|
||||||
|
],
|
||||||
|
dim=0,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
fm.actual_seq_lengths_q = None
|
fm.actual_seq_lengths_q = None
|
||||||
fm.actual_seq_lengths_q_pa = None
|
fm.actual_seq_lengths_q_pa = None
|
||||||
|
fm.actual_seq_lengths_q_pa_cpu = None
|
||||||
|
|
||||||
fm.swa_page_table = (
|
fm.swa_page_table = (
|
||||||
fm.block_tables_swa if fm.block_tables_swa is not None else fm.block_tables
|
fm.block_tables_swa if fm.block_tables_swa is not None else fm.block_tables
|
||||||
@@ -1514,8 +1542,6 @@ class DeepseekV4AscendAttnBackend(
|
|||||||
) -> dict:
|
) -> dict:
|
||||||
fm = self.forward_metadata
|
fm = self.forward_metadata
|
||||||
common = {
|
common = {
|
||||||
"cu_seqlens_q": actual_seq_lengths_q_pa,
|
|
||||||
"seqused_kv": actual_seq_lengths_kv,
|
|
||||||
"cmp_ratio": 1,
|
"cmp_ratio": 1,
|
||||||
"ori_mask_mode": 4,
|
"ori_mask_mode": 4,
|
||||||
"cmp_mask_mode": 3,
|
"cmp_mask_mode": 3,
|
||||||
@@ -1534,28 +1560,27 @@ class DeepseekV4AscendAttnBackend(
|
|||||||
"has_ori_kv": True,
|
"has_ori_kv": True,
|
||||||
"has_cmp_kv": False,
|
"has_cmp_kv": False,
|
||||||
}
|
}
|
||||||
|
# The host metadata op reads CPU int32 mirrors — never a D2H sync of the
|
||||||
|
# device tensors (that would drain the stream and stall overlapped prep).
|
||||||
c1a_kwargs = base_kwargs | common
|
c1a_kwargs = base_kwargs | common
|
||||||
if self._is_dspark_draft_worker:
|
if self._is_dspark_draft_worker:
|
||||||
seq_lens_cpu = getattr(fm, "seq_lens_cpu_int", None)
|
cu_q_cpu = fm.actual_seq_lengths_q_pa_cpu
|
||||||
max_seqlen_kv = (
|
if cu_q_cpu is not None and cu_q_cpu.numel() > bs + 1:
|
||||||
int(seq_lens_cpu[:bs].max().item())
|
cu_q_cpu = cu_q_cpu[: bs + 1]
|
||||||
if seq_lens_cpu is not None and bs > 0
|
host_inputs = {"seqused_kv": fm.seq_lens_cpu_int[:bs].int()}
|
||||||
else int(actual_seq_lengths_kv[:bs].max().item())
|
if cu_q_cpu is not None:
|
||||||
)
|
host_inputs["cu_seqlens_q"] = cu_q_cpu
|
||||||
c1a_kwargs.update(
|
c1a_kwargs = c1a_kwargs | host_inputs
|
||||||
cu_seqlens_ori_kv=actual_seq_lengths_q_pa,
|
metadata_op = torch.ops.npu.sparse_attn_sharedkv_metadata_host
|
||||||
max_seqlen_q=max_seqlen_q,
|
|
||||||
max_seqlen_kv=max_seqlen_kv,
|
|
||||||
)
|
|
||||||
c1a_metadata = torch.ops._C_ascend.npu_sparse_attn_sharedkv_metadata(
|
|
||||||
device=str(actual_seq_lengths_kv.device),
|
|
||||||
**c1a_kwargs,
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
c1a_metadata = torch.ops.custom.npu_sparse_attn_sharedkv_metadata(
|
# The device-side op requires tensor args for backend dispatch; pass
|
||||||
**c1a_kwargs,
|
# the device mirrors just like the pre-refactor call did.
|
||||||
)
|
c1a_kwargs = c1a_kwargs | {
|
||||||
kernel_metadata = {"c1a_metadata": c1a_metadata}
|
"cu_seqlens_q": actual_seq_lengths_q_pa,
|
||||||
|
"seqused_kv": actual_seq_lengths_kv,
|
||||||
|
}
|
||||||
|
metadata_op = torch.ops.custom.npu_sparse_attn_sharedkv_metadata
|
||||||
|
kernel_metadata = {"c1a_metadata": metadata_op(**c1a_kwargs)}
|
||||||
|
|
||||||
if self._dsv4_has_c4:
|
if self._dsv4_has_c4:
|
||||||
c4a_overrides = {
|
c4a_overrides = {
|
||||||
@@ -1665,10 +1690,7 @@ class DeepseekV4AscendAttnBackend(
|
|||||||
if ori_sparse_indices is not None:
|
if ori_sparse_indices is not None:
|
||||||
attn_kwargs["ori_sparse_indices"] = ori_sparse_indices
|
attn_kwargs["ori_sparse_indices"] = ori_sparse_indices
|
||||||
q_arg = attn_kwargs.pop("q")
|
q_arg = attn_kwargs.pop("q")
|
||||||
if self._is_dspark_draft_worker:
|
out, _ = torch.ops.npu.sparse_attn_sharedkv(q_arg, **attn_kwargs)
|
||||||
out, _ = torch.ops._C_ascend.npu_sparse_attn_sharedkv(q_arg, **attn_kwargs)
|
|
||||||
else:
|
|
||||||
out, _ = torch.ops.custom.npu_sparse_attn_sharedkv(q_arg, **attn_kwargs)
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def _forward_compressed(
|
def _forward_compressed(
|
||||||
@@ -1739,7 +1761,7 @@ class DeepseekV4AscendAttnBackend(
|
|||||||
else:
|
else:
|
||||||
attn_kwargs["cmp_sparse_indices"] = None
|
attn_kwargs["cmp_sparse_indices"] = None
|
||||||
q_arg = attn_kwargs.pop("q")
|
q_arg = attn_kwargs.pop("q")
|
||||||
out, _ = torch.ops.custom.npu_sparse_attn_sharedkv(q_arg, **attn_kwargs)
|
out, _ = torch.ops.npu.sparse_attn_sharedkv(q_arg, **attn_kwargs)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def get_swa_out_cache_loc(self, forward_batch: ForwardBatch) -> torch.Tensor:
|
def get_swa_out_cache_loc(self, forward_batch: ForwardBatch) -> torch.Tensor:
|
||||||
|
|||||||
@@ -97,7 +97,7 @@ class NPUCompressStatePool(CompressStatePool):
|
|||||||
128,
|
128,
|
||||||
), f"NPUCompressStatePool only supports ratio in (4, 128); got {ratio}"
|
), f"NPUCompressStatePool only supports ratio in (4, 128); got {ratio}"
|
||||||
assert dtype == torch.float32, (
|
assert dtype == torch.float32, (
|
||||||
"Atlas A3 custom.compressor requires FP32 state_cache, "
|
"Atlas A3 npu.compressor requires FP32 state_cache, "
|
||||||
f"but NPUCompressStatePool got {dtype}."
|
f"but NPUCompressStatePool got {dtype}."
|
||||||
)
|
)
|
||||||
assert ring_size > 0, f"ring_size must be positive, got {ring_size}"
|
assert ring_size > 0, f"ring_size must be positive, got {ring_size}"
|
||||||
@@ -404,7 +404,7 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
|
|||||||
def get_state_cache(self, layer_id: int, from_indexer: bool) -> torch.Tensor:
|
def get_state_cache(self, layer_id: int, from_indexer: bool) -> torch.Tensor:
|
||||||
"""FP32 ``[block_num, ring_size, 2*coff*D]`` view of this layer's
|
"""FP32 ``[block_num, ring_size, 2*coff*D]`` view of this layer's
|
||||||
kv+score buffer — the fused compressor op
|
kv+score buffer — the fused compressor op
|
||||||
(``torch.ops.custom.compressor``)'s ``state_cache`` argument."""
|
(``torch.ops.npu.compressor``)'s ``state_cache`` argument."""
|
||||||
return self._get_state_pool(layer_id, from_indexer).state_cache_3d
|
return self._get_state_pool(layer_id, from_indexer).state_cache_3d
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|||||||
@@ -25,7 +25,36 @@ class OpLibSpec:
|
|||||||
|
|
||||||
|
|
||||||
class TorchOpLoader:
|
class TorchOpLoader:
|
||||||
"""Load a standalone .so into ``torch.ops`` and validate its operators."""
|
"""
|
||||||
|
Loader for PyTorch custom operators from shared libraries.
|
||||||
|
|
||||||
|
This class handles the registration and initialization of custom PyTorch
|
||||||
|
operators (Ops) from dynamically linked shared object (.so) files. It supports
|
||||||
|
environment-based library path discovery, dependency pre-loading, and
|
||||||
|
operator existence validation.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
1. Create an OpLibSpec with operator metadata
|
||||||
|
2. Instantiate TorchOpLoader with the spec
|
||||||
|
3. Call initialize() to load and register the operators
|
||||||
|
|
||||||
|
Example:
|
||||||
|
>>> spec = OpLibSpec(
|
||||||
|
... name="My custom ops",
|
||||||
|
... so_env="MY_LIB_SO_PATH",
|
||||||
|
... namespace="_C_my_lib",
|
||||||
|
... required_ops=("op1", "op2"),
|
||||||
|
... pre_load_imports=("torch", "other_dep"),
|
||||||
|
... )
|
||||||
|
>>> loader = TorchOpLoader(spec)
|
||||||
|
>>> lib_path = loader.initialize()
|
||||||
|
>>> # Ops are now registered under namespace: _C_my_lib.op1()
|
||||||
|
|
||||||
|
The loader will raise appropriate exceptions if:
|
||||||
|
- The shared library cannot be found (via SO_PATH env var or default paths)
|
||||||
|
- Pre-load imports fail
|
||||||
|
- Required operators are missing after loading
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(self, spec: OpLibSpec) -> None:
|
def __init__(self, spec: OpLibSpec) -> None:
|
||||||
self._spec = spec
|
self._spec = spec
|
||||||
@@ -105,18 +134,3 @@ class TorchOpLoader:
|
|||||||
self._loaded_library = library_path
|
self._loaded_library = library_path
|
||||||
logger.info("Registered %s operators from %s", self._spec.name, library_path)
|
logger.info("Registered %s operators from %s", self._spec.name, library_path)
|
||||||
return library_path
|
return library_path
|
||||||
|
|
||||||
|
|
||||||
def initialize_dspark_sparse_attn_ops() -> Optional[Path]:
|
|
||||||
"""Register the DSpark sparse-attention ops before backend execution."""
|
|
||||||
spec = OpLibSpec(
|
|
||||||
name="DSpark sparse-attention",
|
|
||||||
so_env="SGLANG_DSPARK_EXTRA_OPS_SO",
|
|
||||||
namespace="_C_ascend",
|
|
||||||
required_ops=(
|
|
||||||
"npu_sparse_attn_sharedkv_metadata",
|
|
||||||
"npu_sparse_attn_sharedkv",
|
|
||||||
),
|
|
||||||
pre_load_imports=("torch_npu",),
|
|
||||||
)
|
|
||||||
return TorchOpLoader(spec).initialize()
|
|
||||||
|
|||||||
@@ -114,16 +114,15 @@ def get_last_loc(
|
|||||||
"torch_native",
|
"torch_native",
|
||||||
) and decode_backend not in ("ascend", "torch_native")
|
) and decode_backend not in ("ascend", "torch_native")
|
||||||
|
|
||||||
if _is_hip and uses_triton_dispatch:
|
if (_is_hip or _is_npu) and uses_triton_dispatch:
|
||||||
# HIP-only: the legacy get_last_loc_triton kernel emits a
|
# HIP and NPU DSV4: the legacy get_last_loc_triton kernel emits a
|
||||||
# mixed-width int32->int64 store that Triton mis-compiles on HIP,
|
# mixed-width int32->int64 store that Triton mis-compiles on HIP,
|
||||||
# producing out-of-range last_loc values under EAGLE +
|
# producing out-of-range last_loc values under EAGLE +
|
||||||
# page_size>1 (e.g. with aiter unified attention or the triton
|
# page_size>1 (e.g. with aiter unified attention or the triton
|
||||||
# attention backend). The bug is in the Triton HIP codegen, not
|
# attention backend), and can fault in the equivalent NPU DSV4
|
||||||
# in any particular attention backend, so route every HIP path
|
# allocation path. Route those paths through the variant whose
|
||||||
# that would otherwise use get_last_loc_triton through the
|
# in-kernel result remains int32 and is promoted only after launch.
|
||||||
# int32-safe variant. Non-HIP hardware keeps the original
|
# Other hardware/backends keep the original dispatcher below.
|
||||||
# dispatcher below.
|
|
||||||
return get_last_loc_triton_safe(
|
return get_last_loc_triton_safe(
|
||||||
req_to_token, req_pool_indices_tensor, prefix_lens_tensor
|
req_to_token, req_pool_indices_tensor, prefix_lens_tensor
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -349,12 +349,6 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
def init_attention_backends(self):
|
def init_attention_backends(self):
|
||||||
with self._draft_context():
|
with self._draft_context():
|
||||||
if _is_npu:
|
|
||||||
from sglang.srt.hardware_backend.npu.extra_ops_loader import (
|
|
||||||
initialize_dspark_sparse_attn_ops,
|
|
||||||
)
|
|
||||||
|
|
||||||
initialize_dspark_sparse_attn_ops()
|
|
||||||
self._draft_worker.init_attention_backends()
|
self._draft_worker.init_attention_backends()
|
||||||
self._need_mamba_verify_commit = mambaish_config(
|
self._need_mamba_verify_commit = mambaish_config(
|
||||||
self.model_runner.model_config
|
self.model_runner.model_config
|
||||||
|
|||||||
@@ -153,6 +153,7 @@ class TestForwardMetadata(unittest.TestCase):
|
|||||||
"seq_lens",
|
"seq_lens",
|
||||||
"actual_seq_lengths_q",
|
"actual_seq_lengths_q",
|
||||||
"actual_seq_lengths_q_pa",
|
"actual_seq_lengths_q_pa",
|
||||||
|
"actual_seq_lengths_q_pa_cpu",
|
||||||
"actual_seq_lengths_kv",
|
"actual_seq_lengths_kv",
|
||||||
"swa_mask",
|
"swa_mask",
|
||||||
"prefix_lens",
|
"prefix_lens",
|
||||||
|
|||||||
Reference in New Issue
Block a user