[CP] Consolidate decode-context-parallel (DCP) helpers into layers/dcp/ (#29365)
Co-authored-by: Hao Phan <htphan@nvidia.com> Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Hao Phan
Claude Opus 4.8
parent
7820dc60a7
commit
0203c60fdf
@@ -23,13 +23,13 @@ from sglang.srt.layers.attention.flashinfer_backend import (
|
|||||||
create_flashinfer_kv_indices_triton,
|
create_flashinfer_kv_indices_triton,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.utils import assert_buffer_fits
|
from sglang.srt.layers.attention.utils import assert_buffer_fits
|
||||||
from sglang.srt.layers.utils.dcp_utils import (
|
from sglang.srt.layers.dcp import (
|
||||||
DecodeContextParallelMetadata,
|
DecodeContextParallelMetadata,
|
||||||
dcp_enabled,
|
dcp_enabled,
|
||||||
get_attention_dcp_world_size,
|
get_attention_dcp_world_size,
|
||||||
plan_dcp_decode_metadata,
|
|
||||||
update_local_kv_lens_for_dcp,
|
update_local_kv_lens_for_dcp,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.dcp.planner import plan_dcp_decode_metadata
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||||
is_in_tc_piecewise_cuda_graph,
|
is_in_tc_piecewise_cuda_graph,
|
||||||
|
|||||||
@@ -17,12 +17,12 @@ from sglang.srt.layers.attention.utils import (
|
|||||||
create_flashmla_kv_indices_triton,
|
create_flashmla_kv_indices_triton,
|
||||||
get_num_kv_index_blocks_flashmla,
|
get_num_kv_index_blocks_flashmla,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import scaled_fp8_quant
|
from sglang.srt.layers.dcp import (
|
||||||
from sglang.srt.layers.utils.dcp_utils import (
|
|
||||||
dcp_enabled,
|
dcp_enabled,
|
||||||
get_attention_dcp_rank,
|
get_attention_dcp_rank,
|
||||||
get_attention_dcp_world_size,
|
get_attention_dcp_world_size,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.quantization.fp8_kernel import scaled_fp8_quant
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
|
|||||||
@@ -17,8 +17,8 @@ from sglang.srt.layers.attention.triton_ops.kv_indices import (
|
|||||||
create_flashinfer_kv_indices_triton,
|
create_flashinfer_kv_indices_triton,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.triton_ops.metadata import get_num_kv_splits_triton
|
from sglang.srt.layers.attention.triton_ops.metadata import get_num_kv_splits_triton
|
||||||
from sglang.srt.layers.attention.utils import (
|
from sglang.srt.layers.dcp import (
|
||||||
cp_lse_ag_out_rs,
|
cp_lse_ag_out_rs_mha,
|
||||||
create_triton_kv_indices_for_dcp_triton,
|
create_triton_kv_indices_for_dcp_triton,
|
||||||
get_dcp_lens,
|
get_dcp_lens,
|
||||||
)
|
)
|
||||||
@@ -1489,7 +1489,7 @@ class TritonAttnBackend(AttentionBackend):
|
|||||||
skip_extend=True,
|
skip_extend=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
prefix_out, prefix_lse = cp_lse_ag_out_rs(
|
prefix_out, prefix_lse = cp_lse_ag_out_rs_mha(
|
||||||
prefix_out, prefix_lse, group, return_lse=True
|
prefix_out, prefix_lse, group, return_lse=True
|
||||||
)
|
)
|
||||||
final_lse = torch.logaddexp(prefix_lse, current_lse)
|
final_lse = torch.logaddexp(prefix_lse, current_lse)
|
||||||
@@ -1748,7 +1748,7 @@ class TritonAttnBackend(AttentionBackend):
|
|||||||
],
|
],
|
||||||
dim=-1,
|
dim=-1,
|
||||||
)
|
)
|
||||||
o = cp_lse_ag_out_rs(o_for_decode, local_lse, group)
|
o = cp_lse_ag_out_rs_mha(o_for_decode, local_lse, group)
|
||||||
return o.reshape(-1, layer.tp_q_head_num * layer.v_head_dim).to(q.dtype)
|
return o.reshape(-1, layer.tp_q_head_num * layer.v_head_dim).to(q.dtype)
|
||||||
|
|
||||||
self.decode_attention_fwd(
|
self.decode_attention_fwd(
|
||||||
|
|||||||
@@ -3,7 +3,6 @@ import triton
|
|||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
from sglang.jit_kernel.utils import is_arch_support_pdl
|
from sglang.jit_kernel.utils import is_arch_support_pdl
|
||||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
|
||||||
from sglang.srt.layers.attention.triton_ops.cache_ops import (
|
from sglang.srt.layers.attention.triton_ops.cache_ops import (
|
||||||
concat_and_cast_mha_k_kernel as concat_and_cast_mha_k_kernel,
|
concat_and_cast_mha_k_kernel as concat_and_cast_mha_k_kernel,
|
||||||
)
|
)
|
||||||
@@ -179,104 +178,6 @@ def concat_mla_absorb_q_general(q_nope, q_rope):
|
|||||||
return torch.cat([q_nope, q_rope], dim=-1)
|
return torch.cat([q_nope, q_rope], dim=-1)
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Decode Context Parallel (DCP) helpers.
|
|
||||||
#
|
|
||||||
# Not part of upstream main (PR #26000 centralized the other Triton utility
|
|
||||||
# kernels into triton_ops/*). These three live here because they are DCP-only:
|
|
||||||
# - create_triton_kv_indices_for_dcp_triton: per-rank local KV indices
|
|
||||||
# - get_dcp_lens: per-rank visible KV length
|
|
||||||
# - cp_lse_ag_out_rs: merge DCP partial attention via natural-log LSE
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
@triton.jit
|
|
||||||
def create_triton_kv_indices_for_dcp_triton(
|
|
||||||
req_to_token_ptr, # [max_batch, max_context_len]
|
|
||||||
req_pool_indices_ptr,
|
|
||||||
dcp_kernel_lens_ptr,
|
|
||||||
kv_indptr,
|
|
||||||
kv_start_idx,
|
|
||||||
kv_indices_ptr,
|
|
||||||
req_to_token_ptr_stride: tl.constexpr,
|
|
||||||
dcp_size: tl.constexpr,
|
|
||||||
dcp_rank: tl.constexpr,
|
|
||||||
):
|
|
||||||
BLOCK_SIZE: tl.constexpr = 512
|
|
||||||
pid = tl.program_id(axis=0)
|
|
||||||
req_pool_index = tl.load(req_pool_indices_ptr + pid)
|
|
||||||
kv_indices_offset = tl.load(kv_indptr + pid)
|
|
||||||
|
|
||||||
kv_start = 0
|
|
||||||
if kv_start_idx:
|
|
||||||
kv_start = tl.load(kv_start_idx + pid).to(tl.int32)
|
|
||||||
|
|
||||||
# First absolute token position in this range owned by dcp_rank.
|
|
||||||
# Triton follows C-style remainder for negative values, so avoid
|
|
||||||
# computing the offset as a negative remainder when kv_start > dcp_rank.
|
|
||||||
kv_start_mod = kv_start % dcp_size
|
|
||||||
first = kv_start + ((dcp_rank + dcp_size - kv_start_mod) % dcp_size)
|
|
||||||
local_len = tl.load(dcp_kernel_lens_ptr + pid).to(tl.int32)
|
|
||||||
|
|
||||||
num_loop = tl.cdiv(local_len, BLOCK_SIZE)
|
|
||||||
for i in range(num_loop):
|
|
||||||
offset = tl.arange(0, BLOCK_SIZE).to(tl.int64) + i * BLOCK_SIZE
|
|
||||||
mask = offset < local_len
|
|
||||||
abs_pos = first + offset * dcp_size
|
|
||||||
data = tl.load(
|
|
||||||
req_to_token_ptr + req_pool_index * req_to_token_ptr_stride + abs_pos,
|
|
||||||
mask=mask,
|
|
||||||
)
|
|
||||||
tl.store(
|
|
||||||
kv_indices_ptr + kv_indices_offset + offset, data // dcp_size, mask=mask
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def get_dcp_lens(
|
|
||||||
lens: torch.Tensor,
|
|
||||||
dcp_size: int,
|
|
||||||
dcp_rank: int,
|
|
||||||
start: torch.Tensor | None = None,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
if dcp_size == 1:
|
|
||||||
return lens
|
|
||||||
if start is None:
|
|
||||||
return lens // dcp_size + (dcp_rank < lens % dcp_size)
|
|
||||||
|
|
||||||
first = start + torch.remainder(dcp_rank - start, dcp_size)
|
|
||||||
remaining = start + lens - first
|
|
||||||
return torch.clamp((remaining + dcp_size - 1) // dcp_size, min=0)
|
|
||||||
|
|
||||||
|
|
||||||
def cp_lse_ag_out_rs(
|
|
||||||
cp_attn_out: torch.Tensor,
|
|
||||||
cp_attn_lse: torch.Tensor,
|
|
||||||
cp_group: GroupCoordinator,
|
|
||||||
return_lse: bool = False,
|
|
||||||
):
|
|
||||||
"""Merge DCP partial attention outputs using natural-log LSE."""
|
|
||||||
if cp_group.world_size == 1:
|
|
||||||
return (cp_attn_out, cp_attn_lse) if return_lse else cp_attn_out
|
|
||||||
|
|
||||||
cp_attn_lse = cp_attn_lse.contiguous()
|
|
||||||
lses = cp_group.all_gather(cp_attn_lse, dim=0).view(
|
|
||||||
(cp_group.world_size,) + cp_attn_lse.shape
|
|
||||||
)
|
|
||||||
global_lse = torch.logsumexp(lses, dim=0)
|
|
||||||
scale = torch.exp(cp_attn_lse - global_lse).unsqueeze(-1)
|
|
||||||
scale = torch.nan_to_num(scale, nan=0.0, posinf=0.0, neginf=0.0)
|
|
||||||
|
|
||||||
out = torch.nan_to_num(cp_attn_out, nan=0.0, posinf=0.0, neginf=0.0) * scale
|
|
||||||
out = cp_group.all_reduce(out)
|
|
||||||
|
|
||||||
cp_num_heads = global_lse.shape[1] // cp_group.world_size
|
|
||||||
cp_rank = cp_group.rank_in_group
|
|
||||||
head_start = cp_num_heads * cp_rank
|
|
||||||
head_end = cp_num_heads * (cp_rank + 1)
|
|
||||||
out = out[:, head_start:head_end, :].contiguous()
|
|
||||||
if return_lse:
|
|
||||||
return out, global_lse[:, head_start:head_end].contiguous()
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def reshape_and_cache_flash(
|
def reshape_and_cache_flash(
|
||||||
key_ptr,
|
key_ptr,
|
||||||
|
|||||||
@@ -0,0 +1,72 @@
|
|||||||
|
# Copyright 2023-2026 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
"""Decode Context Parallel (DCP) — consolidated home for the primitives that
|
||||||
|
were previously split between layers/attention/utils.py (PR #25090, Triton/MHA)
|
||||||
|
and layers/utils/dcp_utils.py (PR #14194, FlashInfer-MLA).
|
||||||
|
|
||||||
|
The two ``cp_lse_ag_out_rs`` variants are kept distinct (``_mha`` torch/all-reduce,
|
||||||
|
``_mla`` Triton/reduce-scatter) because their bodies are backend-forced.
|
||||||
|
|
||||||
|
Only the symbols imported by code OUTSIDE this subpackage are re-exported here.
|
||||||
|
Package-internal helpers (the @triton.jit kernels, ``CPTritonContext``,
|
||||||
|
``correct_attn_out``, ``create_dcp_kv_indices``, ``update_kv_lens_and_indices``,
|
||||||
|
``_all_gather_dcp_kv_cache``) stay private to their submodules — import them from
|
||||||
|
``sglang.srt.layers.dcp.{kernels,comm}`` if ever needed internally."""
|
||||||
|
|
||||||
|
from sglang.srt.layers.dcp.comm import (
|
||||||
|
all_gather_kv_cache_for_dcp,
|
||||||
|
all_gather_kv_cache_for_mha_chunk_extend,
|
||||||
|
all_gather_kv_cache_for_mha_extend,
|
||||||
|
all_gather_kv_cache_for_mla_extend,
|
||||||
|
all_gather_q_for_mla_decode,
|
||||||
|
cp_lse_ag_out_rs_mha,
|
||||||
|
cp_lse_ag_out_rs_mla,
|
||||||
|
dcp_enabled,
|
||||||
|
get_attention_dcp_rank,
|
||||||
|
get_attention_dcp_world_size,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.dcp.kernels import create_triton_kv_indices_for_dcp_triton
|
||||||
|
from sglang.srt.layers.dcp.layout import (
|
||||||
|
filter_dcp_local_kv_indices,
|
||||||
|
get_dcp_lens,
|
||||||
|
update_local_kv_lens_for_dcp,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata
|
||||||
|
|
||||||
|
# NOTE: planner.py is intentionally NOT imported here. It depends on server_args
|
||||||
|
# (get_global_server_args), whereas this package-init executes at module-load time
|
||||||
|
# for every eager importer of the DCP primitives — triton_backend,
|
||||||
|
# mem_cache.memory_pool, mem_cache.triton_ops.mla_buffer, mem_cache.kv_cache_builder,
|
||||||
|
# the FlashInfer-MLA / FlashMLA backends, and the deepseek forward methods. Keeping
|
||||||
|
# the init server_args-free avoids a load-time import edge into server_args. Import
|
||||||
|
# planner functions from sglang.srt.layers.dcp.planner directly.
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"DecodeContextParallelMetadata",
|
||||||
|
"all_gather_kv_cache_for_dcp",
|
||||||
|
"all_gather_kv_cache_for_mha_chunk_extend",
|
||||||
|
"all_gather_kv_cache_for_mha_extend",
|
||||||
|
"all_gather_kv_cache_for_mla_extend",
|
||||||
|
"all_gather_q_for_mla_decode",
|
||||||
|
"cp_lse_ag_out_rs_mha",
|
||||||
|
"cp_lse_ag_out_rs_mla",
|
||||||
|
"create_triton_kv_indices_for_dcp_triton",
|
||||||
|
"dcp_enabled",
|
||||||
|
"filter_dcp_local_kv_indices",
|
||||||
|
"get_attention_dcp_rank",
|
||||||
|
"get_attention_dcp_world_size",
|
||||||
|
"get_dcp_lens",
|
||||||
|
"update_local_kv_lens_for_dcp",
|
||||||
|
]
|
||||||
@@ -0,0 +1,339 @@
|
|||||||
|
# Copyright 2023-2026 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
"""Group accessors, LSE-merge and all-gather collectives for decode CP (DCP).
|
||||||
|
|
||||||
|
The two LSE-merge variants kept separate (bodies are backend-forced, see
|
||||||
|
PR #25090 vs #14194):
|
||||||
|
- cp_lse_ag_out_rs_mha: torch / natural-log logsumexp / all-reduce + head slice
|
||||||
|
- cp_lse_ag_out_rs_mla: Triton (log2/exp2) correction / reduce-scatter
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
|
use_symmetric_memory,
|
||||||
|
)
|
||||||
|
from sglang.srt.distributed.parallel_state import (
|
||||||
|
GroupCoordinator,
|
||||||
|
get_dcp_group,
|
||||||
|
get_dcp_group_no_assert,
|
||||||
|
get_dcp_rank,
|
||||||
|
get_dcp_world_size,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.dcp.kernels import CPTritonContext, correct_attn_out
|
||||||
|
from sglang.srt.utils import is_cuda
|
||||||
|
|
||||||
|
|
||||||
|
def dcp_enabled() -> bool:
|
||||||
|
"""
|
||||||
|
only checks whether dcp enabled for cuda platform
|
||||||
|
"""
|
||||||
|
if get_dcp_group_no_assert() is None:
|
||||||
|
return False
|
||||||
|
if not is_cuda():
|
||||||
|
return False
|
||||||
|
return get_dcp_world_size() > 1
|
||||||
|
|
||||||
|
|
||||||
|
def get_attention_dcp_world_size() -> int:
|
||||||
|
if not dcp_enabled():
|
||||||
|
return 1
|
||||||
|
return get_dcp_world_size()
|
||||||
|
|
||||||
|
|
||||||
|
def get_attention_dcp_rank() -> int:
|
||||||
|
if not dcp_enabled():
|
||||||
|
return 0
|
||||||
|
return get_dcp_rank()
|
||||||
|
|
||||||
|
|
||||||
|
def _ag_lse(cp_attn_lse: torch.Tensor, cp_group: GroupCoordinator) -> torch.Tensor:
|
||||||
|
"""All-gather each rank's LSE into a ``[world_size, *lse.shape]`` stack.
|
||||||
|
|
||||||
|
Shared prologue of both ``cp_lse_ag_out_rs_{mha,mla}``. Callers do their own
|
||||||
|
pre-processing (``contiguous()`` for MHA, fp32 cast for MLA) before calling.
|
||||||
|
"""
|
||||||
|
return cp_group.all_gather(cp_attn_lse, dim=0).view(
|
||||||
|
(cp_group.world_size,) + cp_attn_lse.shape
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def cp_lse_ag_out_rs_mha(
|
||||||
|
cp_attn_out: torch.Tensor,
|
||||||
|
cp_attn_lse: torch.Tensor,
|
||||||
|
cp_group: GroupCoordinator,
|
||||||
|
return_lse: bool = False,
|
||||||
|
):
|
||||||
|
"""Merge DCP partial attention outputs using natural-log LSE (PR #25090)."""
|
||||||
|
if cp_group.world_size == 1:
|
||||||
|
return (cp_attn_out, cp_attn_lse) if return_lse else cp_attn_out
|
||||||
|
|
||||||
|
cp_attn_lse = cp_attn_lse.contiguous()
|
||||||
|
lses = _ag_lse(cp_attn_lse, cp_group)
|
||||||
|
global_lse = torch.logsumexp(lses, dim=0)
|
||||||
|
scale = torch.exp(cp_attn_lse - global_lse).unsqueeze(-1)
|
||||||
|
scale = torch.nan_to_num(scale, nan=0.0, posinf=0.0, neginf=0.0)
|
||||||
|
|
||||||
|
out = torch.nan_to_num(cp_attn_out, nan=0.0, posinf=0.0, neginf=0.0) * scale
|
||||||
|
out = cp_group.all_reduce(out)
|
||||||
|
|
||||||
|
cp_num_heads = global_lse.shape[1] // cp_group.world_size
|
||||||
|
cp_rank = cp_group.rank_in_group
|
||||||
|
head_start = cp_num_heads * cp_rank
|
||||||
|
head_end = cp_num_heads * (cp_rank + 1)
|
||||||
|
out = out[:, head_start:head_end, :].contiguous()
|
||||||
|
if return_lse:
|
||||||
|
return out, global_lse[:, head_start:head_end].contiguous()
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def cp_lse_ag_out_rs_mla(
|
||||||
|
cp_attn_out: torch.Tensor,
|
||||||
|
cp_attn_lse: torch.Tensor,
|
||||||
|
cp_group: GroupCoordinator,
|
||||||
|
ctx: Optional[CPTritonContext] = None,
|
||||||
|
):
|
||||||
|
"""Merge DCP partial attention outputs via Triton correction (PR #14194).
|
||||||
|
|
||||||
|
cp_attn_out: [ B, H, D ]
|
||||||
|
cp_attn_lse: [ B, H ]
|
||||||
|
"""
|
||||||
|
if cp_group.world_size == 1:
|
||||||
|
return cp_attn_out
|
||||||
|
|
||||||
|
if ctx is None:
|
||||||
|
ctx = CPTritonContext()
|
||||||
|
|
||||||
|
with use_symmetric_memory(cp_group):
|
||||||
|
# cp_attn_out is [B,H,D], we want to transpose it to [H,B,D] for the kernel, and then transpose back after correction.
|
||||||
|
new_output = cp_attn_out.new_empty(
|
||||||
|
cp_attn_out.transpose(0, 1).shape, dtype=torch.float32
|
||||||
|
)
|
||||||
|
cp_attn_lse = cp_attn_lse.to(torch.float32)
|
||||||
|
lses = _ag_lse(cp_attn_lse, cp_group)
|
||||||
|
out, _ = correct_attn_out(
|
||||||
|
cp_attn_out, lses, cp_group.rank_in_group, ctx, new_output
|
||||||
|
)
|
||||||
|
out = cp_group.reduce_scatter_along_dim(out, dim=0)
|
||||||
|
return out.to(cp_attn_out.dtype)
|
||||||
|
|
||||||
|
|
||||||
|
def _all_gather_dcp_kv_cache(kv_a: torch.Tensor):
|
||||||
|
dcp_world_size = get_dcp_world_size()
|
||||||
|
# not use symmetric_memory unless torch mem_pool updated, see https://github.com/pytorch/pytorch/issues/178138
|
||||||
|
gathered_kv_a = kv_a.new_empty(
|
||||||
|
(kv_a.shape[0] * dcp_world_size, *kv_a.shape[1:]),
|
||||||
|
)
|
||||||
|
get_dcp_group().all_gather_into_tensor(gathered_kv_a, kv_a)
|
||||||
|
gathered_kv_a = (
|
||||||
|
gathered_kv_a.reshape((dcp_world_size,) + kv_a.shape)
|
||||||
|
.transpose(0, 1)
|
||||||
|
.reshape(-1, *kv_a.shape[1:])
|
||||||
|
)
|
||||||
|
return gathered_kv_a
|
||||||
|
|
||||||
|
|
||||||
|
def all_gather_kv_cache_for_mha_chunk_extend(
|
||||||
|
kv_a: torch.Tensor,
|
||||||
|
k_pe: torch.Tensor,
|
||||||
|
prefix_kv_lens_cpu: torch.Tensor,
|
||||||
|
prefix_starts_cpu: torch.Tensor = None,
|
||||||
|
):
|
||||||
|
if dcp_enabled():
|
||||||
|
kv_a = kv_a.unsqueeze(1)
|
||||||
|
gathered_kv = all_gather_kv_cache_for_dcp(
|
||||||
|
kv_a,
|
||||||
|
k_pe,
|
||||||
|
prefix_kv_lens_cpu,
|
||||||
|
prefix_starts_cpu,
|
||||||
|
)
|
||||||
|
kv_a, k_pe = gathered_kv.split([kv_a.shape[-1], k_pe.shape[-1]], dim=-1)
|
||||||
|
kv_a = kv_a.squeeze(1)
|
||||||
|
return kv_a.contiguous(), k_pe.contiguous()
|
||||||
|
|
||||||
|
|
||||||
|
def all_gather_kv_cache_for_mha_extend(
|
||||||
|
token_to_kv_pool,
|
||||||
|
attn_mqa,
|
||||||
|
dcp_local_prefix_kv_indices,
|
||||||
|
seq_lens,
|
||||||
|
extend_prefix_lens,
|
||||||
|
extend_prefix_lens_cpu: list[int],
|
||||||
|
extend_seq_lens,
|
||||||
|
kv_a: torch.Tensor,
|
||||||
|
k_pe: torch.Tensor,
|
||||||
|
):
|
||||||
|
prefix_kv_a, prefix_k_pe = token_to_kv_pool.get_mla_kv_buffer(
|
||||||
|
attn_mqa, dcp_local_prefix_kv_indices
|
||||||
|
)
|
||||||
|
extend_prefix_lens_cpu = torch.tensor(extend_prefix_lens_cpu)
|
||||||
|
gathered_kv_cache = all_gather_kv_cache_for_dcp(
|
||||||
|
prefix_kv_a,
|
||||||
|
prefix_k_pe,
|
||||||
|
extend_prefix_lens_cpu,
|
||||||
|
)
|
||||||
|
prefix_kv_a, prefix_k_pe = gathered_kv_cache.split(
|
||||||
|
[kv_a.shape[-1], k_pe.shape[-1]], dim=-1
|
||||||
|
)
|
||||||
|
prefix_kv_a = prefix_kv_a.squeeze(1)
|
||||||
|
# re-organize kv with query orders
|
||||||
|
prefix_lens_cu = torch.zeros(
|
||||||
|
len(seq_lens) + 1,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=kv_a.device,
|
||||||
|
)
|
||||||
|
extend_lens_cu = torch.zeros_like(prefix_lens_cu)
|
||||||
|
prefix_lens_cu[1:] = torch.cumsum(extend_prefix_lens, dim=0)
|
||||||
|
extend_lens_cu[1:] = torch.cumsum(extend_seq_lens, dim=0)
|
||||||
|
kv_a_tuple = ()
|
||||||
|
k_pe_tuple = ()
|
||||||
|
for i in range(len(seq_lens)):
|
||||||
|
kv_a_tuple += (
|
||||||
|
prefix_kv_a[prefix_lens_cu[i] : prefix_lens_cu[i + 1]],
|
||||||
|
kv_a[extend_lens_cu[i] : extend_lens_cu[i + 1]],
|
||||||
|
)
|
||||||
|
k_pe_tuple += (
|
||||||
|
prefix_k_pe[prefix_lens_cu[i] : prefix_lens_cu[i + 1]],
|
||||||
|
k_pe[extend_lens_cu[i] : extend_lens_cu[i + 1]],
|
||||||
|
)
|
||||||
|
kv_a = torch.cat(kv_a_tuple, dim=0)
|
||||||
|
k_pe = torch.cat(k_pe_tuple, dim=0)
|
||||||
|
return kv_a.contiguous(), k_pe.contiguous()
|
||||||
|
|
||||||
|
|
||||||
|
def all_gather_q_for_mla_decode(
|
||||||
|
q_nope_out: torch.Tensor,
|
||||||
|
q_pe: torch.Tensor,
|
||||||
|
):
|
||||||
|
with use_symmetric_memory(get_dcp_group()):
|
||||||
|
# transpose q_pe and q_nope_out from [B, H, L] to [H, B, L]
|
||||||
|
combined = torch.cat([q_pe.transpose(0, 1), q_nope_out.transpose(0, 1)], dim=-1)
|
||||||
|
gathered = get_dcp_group().all_gather(combined, dim=0)
|
||||||
|
d_pe = q_pe.size(-1)
|
||||||
|
d_nope = q_nope_out.size(-1)
|
||||||
|
q_pe, q_nope_out = gathered.split([d_pe, d_nope], dim=-1)
|
||||||
|
q_pe = q_pe.transpose(0, 1)
|
||||||
|
q_nope_out = q_nope_out.transpose(0, 1)
|
||||||
|
return q_nope_out, q_pe
|
||||||
|
|
||||||
|
|
||||||
|
def all_gather_kv_cache_for_mla_extend(
|
||||||
|
token_to_kv_pool,
|
||||||
|
attn_mqa,
|
||||||
|
extend_prefix_lens_cpu: list[int],
|
||||||
|
dcp_local_prefix_kv_indices,
|
||||||
|
dcp_extend_prefix_lens_sum,
|
||||||
|
dcp_kv_buffer,
|
||||||
|
kv_lora_rank,
|
||||||
|
k_nope,
|
||||||
|
k_pe,
|
||||||
|
):
|
||||||
|
cache_k_nope, cache_k_rope = token_to_kv_pool.get_mla_kv_buffer(
|
||||||
|
attn_mqa,
|
||||||
|
dcp_local_prefix_kv_indices,
|
||||||
|
)
|
||||||
|
extend_prefix_lens_cpu = torch.tensor(extend_prefix_lens_cpu)
|
||||||
|
# all gather kv cache into forward_batch.attn_dcp_metadata.dcp_kv_buffer
|
||||||
|
gathered_kv = all_gather_kv_cache_for_dcp(
|
||||||
|
cache_k_nope,
|
||||||
|
cache_k_rope,
|
||||||
|
extend_prefix_lens_cpu,
|
||||||
|
prefix_starts_cpu=torch.zeros_like(extend_prefix_lens_cpu),
|
||||||
|
)
|
||||||
|
dcp_kv_buffer[:dcp_extend_prefix_lens_sum] = gathered_kv
|
||||||
|
|
||||||
|
# copy local kv cache into forward_batch.attn_dcp_metadata.dcp_kv_buffer
|
||||||
|
dcp_kv_buffer[
|
||||||
|
dcp_extend_prefix_lens_sum:,
|
||||||
|
...,
|
||||||
|
:kv_lora_rank,
|
||||||
|
] = k_nope
|
||||||
|
dcp_kv_buffer[
|
||||||
|
dcp_extend_prefix_lens_sum:,
|
||||||
|
...,
|
||||||
|
kv_lora_rank:,
|
||||||
|
] = k_pe
|
||||||
|
|
||||||
|
|
||||||
|
# all gather kv cache and re-org to query orders
|
||||||
|
def all_gather_kv_cache_for_dcp(
|
||||||
|
prefix_kv_a: torch.Tensor,
|
||||||
|
prefix_k_pe: torch.Tensor,
|
||||||
|
prefix_kv_lens_cpu: torch.Tensor,
|
||||||
|
prefix_starts_cpu: torch.Tensor = None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
prefix_kv_a and prefix_k_pe should have same shape, expect for last dim
|
||||||
|
"""
|
||||||
|
if not dcp_enabled():
|
||||||
|
return torch.cat([prefix_kv_a, prefix_k_pe], dim=-1)
|
||||||
|
# 1. compute max kv_lens for each seq
|
||||||
|
dcp_world_size = get_dcp_world_size()
|
||||||
|
dcp_rank = get_dcp_rank()
|
||||||
|
|
||||||
|
if prefix_starts_cpu is None:
|
||||||
|
prefix_starts_cpu = torch.zeros_like(prefix_kv_lens_cpu)
|
||||||
|
|
||||||
|
left_pads = prefix_starts_cpu % dcp_world_size > dcp_rank
|
||||||
|
left_pads = left_pads.to(torch.int32)
|
||||||
|
right_pads = (
|
||||||
|
prefix_starts_cpu + prefix_kv_lens_cpu - 1
|
||||||
|
) % dcp_world_size < dcp_rank
|
||||||
|
right_pads = right_pads.to(torch.int32)
|
||||||
|
padded_lens = (
|
||||||
|
prefix_kv_lens_cpu + (prefix_starts_cpu % dcp_world_size) + dcp_world_size - 1
|
||||||
|
) // dcp_world_size
|
||||||
|
|
||||||
|
local_kv_lens = padded_lens - left_pads - right_pads
|
||||||
|
local_kv_lens_cu = torch.zeros(
|
||||||
|
len(prefix_kv_lens_cpu) + 1,
|
||||||
|
dtype=torch.int32,
|
||||||
|
)
|
||||||
|
local_kv_lens_cu[1:] = torch.cumsum(local_kv_lens, dim=0)
|
||||||
|
|
||||||
|
padded_kv_cache_arr = []
|
||||||
|
prefix_kv_cache = torch.cat([prefix_kv_a, prefix_k_pe], dim=-1)
|
||||||
|
for req_idx in range(len(prefix_kv_lens_cpu)):
|
||||||
|
padded_tensor = prefix_kv_cache.new_empty(
|
||||||
|
(padded_lens[req_idx].item(),) + prefix_kv_cache.size()[1:]
|
||||||
|
)
|
||||||
|
padded_tensor[
|
||||||
|
left_pads[req_idx] : left_pads[req_idx] + local_kv_lens[req_idx]
|
||||||
|
] = prefix_kv_cache[local_kv_lens_cu[req_idx] : local_kv_lens_cu[req_idx + 1]]
|
||||||
|
padded_kv_cache_arr.append(padded_tensor)
|
||||||
|
|
||||||
|
padded_kv_cache = torch.cat(padded_kv_cache_arr, dim=0)
|
||||||
|
|
||||||
|
gatherd_kv_cache = _all_gather_dcp_kv_cache(padded_kv_cache)
|
||||||
|
|
||||||
|
# 2. re-org kv cache to query orders
|
||||||
|
padded_lens_cu = torch.zeros(
|
||||||
|
len(prefix_kv_lens_cpu) + 1,
|
||||||
|
dtype=torch.int32,
|
||||||
|
)
|
||||||
|
padded_lens_cu[1:] = torch.cumsum(padded_lens, dim=0)
|
||||||
|
kv_cache_tuple = ()
|
||||||
|
for req_idx in range(len(prefix_kv_lens_cpu)):
|
||||||
|
kv_cache_tuple += (
|
||||||
|
gatherd_kv_cache[
|
||||||
|
padded_lens_cu[req_idx] * dcp_world_size
|
||||||
|
+ (prefix_starts_cpu[req_idx] % dcp_world_size) :
|
||||||
|
][: prefix_kv_lens_cpu[req_idx]],
|
||||||
|
)
|
||||||
|
gatherd_kv_cache = torch.cat(kv_cache_tuple, dim=0)
|
||||||
|
|
||||||
|
return gatherd_kv_cache
|
||||||
@@ -0,0 +1,333 @@
|
|||||||
|
# Copyright 2023-2026 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
"""Triton kernels for decode context parallel (DCP).
|
||||||
|
|
||||||
|
Consolidated from the two merged DCP implementations:
|
||||||
|
- create_triton_kv_indices_for_dcp_triton (PR #25090, Triton/MHA path)
|
||||||
|
- create_dcp_kv_indices / update_kv_lens_and_indices (PR #14194, MLA path)
|
||||||
|
- _correct_attn_cp_out_kernel / correct_attn_out / CPTritonContext (PR #14194)
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# KV-index build (PR #25090, Triton/MHA): per-rank local KV indices.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
@triton.jit
|
||||||
|
def create_triton_kv_indices_for_dcp_triton(
|
||||||
|
req_to_token_ptr, # [max_batch, max_context_len]
|
||||||
|
req_pool_indices_ptr,
|
||||||
|
dcp_kernel_lens_ptr,
|
||||||
|
kv_indptr,
|
||||||
|
kv_start_idx,
|
||||||
|
kv_indices_ptr,
|
||||||
|
req_to_token_ptr_stride: tl.constexpr,
|
||||||
|
dcp_size: tl.constexpr,
|
||||||
|
dcp_rank: tl.constexpr,
|
||||||
|
):
|
||||||
|
BLOCK_SIZE: tl.constexpr = 512
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
req_pool_index = tl.load(req_pool_indices_ptr + pid)
|
||||||
|
kv_indices_offset = tl.load(kv_indptr + pid)
|
||||||
|
|
||||||
|
kv_start = 0
|
||||||
|
if kv_start_idx:
|
||||||
|
kv_start = tl.load(kv_start_idx + pid).to(tl.int32)
|
||||||
|
|
||||||
|
# First absolute token position in this range owned by dcp_rank.
|
||||||
|
# Triton follows C-style remainder for negative values, so avoid
|
||||||
|
# computing the offset as a negative remainder when kv_start > dcp_rank.
|
||||||
|
kv_start_mod = kv_start % dcp_size
|
||||||
|
first = kv_start + ((dcp_rank + dcp_size - kv_start_mod) % dcp_size)
|
||||||
|
local_len = tl.load(dcp_kernel_lens_ptr + pid).to(tl.int32)
|
||||||
|
|
||||||
|
num_loop = tl.cdiv(local_len, BLOCK_SIZE)
|
||||||
|
for i in range(num_loop):
|
||||||
|
offset = tl.arange(0, BLOCK_SIZE).to(tl.int64) + i * BLOCK_SIZE
|
||||||
|
mask = offset < local_len
|
||||||
|
abs_pos = first + offset * dcp_size
|
||||||
|
data = tl.load(
|
||||||
|
req_to_token_ptr + req_pool_index * req_to_token_ptr_stride + abs_pos,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
tl.store(
|
||||||
|
kv_indices_ptr + kv_indices_offset + offset, data // dcp_size, mask=mask
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# KV-index build (PR #14194, MLA): global prefix+extend layout for the
|
||||||
|
# all-gathered dcp_kv_buffer, plus the per-rank shard/compact kernel.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
@triton.jit
|
||||||
|
def create_dcp_kv_indices(
|
||||||
|
kv_indptr,
|
||||||
|
extend_lens_ptr,
|
||||||
|
extend_cu_lens_ptr,
|
||||||
|
extend_prefix_lens_ptr,
|
||||||
|
extend_cu_prefix_lens_ptr,
|
||||||
|
kv_indices_ptr,
|
||||||
|
extend_prefix_lens_sum,
|
||||||
|
dcp_world_size: tl.constexpr,
|
||||||
|
):
|
||||||
|
BLOCK_SIZE: tl.constexpr = 512
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
prefix_len = tl.load(extend_prefix_lens_ptr + pid)
|
||||||
|
prefix_start = tl.load(extend_cu_prefix_lens_ptr + pid)
|
||||||
|
kv_ind_start = tl.load(kv_indptr + pid)
|
||||||
|
num_loop = tl.cdiv(prefix_len, BLOCK_SIZE)
|
||||||
|
for i in range(num_loop):
|
||||||
|
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
||||||
|
mask = offset < prefix_len
|
||||||
|
data = prefix_start + offset
|
||||||
|
tl.store(kv_indices_ptr + kv_ind_start + offset, data, mask=mask)
|
||||||
|
extend_len = tl.load(extend_lens_ptr + pid)
|
||||||
|
extend_start = tl.load(extend_cu_lens_ptr + pid)
|
||||||
|
num_loop = tl.cdiv(extend_len, BLOCK_SIZE)
|
||||||
|
for i in range(num_loop):
|
||||||
|
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
||||||
|
mask = offset < extend_len
|
||||||
|
data = extend_prefix_lens_sum + extend_start + offset
|
||||||
|
tl.store(
|
||||||
|
kv_indices_ptr + kv_ind_start + prefix_len + offset,
|
||||||
|
data,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def update_kv_lens_and_indices(
|
||||||
|
kv_lens: torch.Tensor,
|
||||||
|
kv_lens_cumsum: torch.Tensor,
|
||||||
|
kv_indices: torch.Tensor,
|
||||||
|
local_kv_lens: torch.Tensor,
|
||||||
|
local_kv_lens_cumsum: torch.Tensor,
|
||||||
|
local_kv_indices: torch.Tensor,
|
||||||
|
dcp_rank: tl.constexpr,
|
||||||
|
dcp_world_size: tl.constexpr,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
):
|
||||||
|
bs_idx = tl.program_id(0)
|
||||||
|
block_idx = tl.program_id(1)
|
||||||
|
|
||||||
|
local_kv_len = tl.load(local_kv_lens + bs_idx)
|
||||||
|
local_kv_indices_start = tl.load(local_kv_lens_cumsum + bs_idx)
|
||||||
|
kv_indices_start = tl.load(kv_lens_cumsum + bs_idx)
|
||||||
|
|
||||||
|
block_start = block_idx * BLOCK_SIZE
|
||||||
|
offsets = block_start + tl.arange(0, BLOCK_SIZE)
|
||||||
|
|
||||||
|
mask = offsets < local_kv_len
|
||||||
|
|
||||||
|
kv_indice_offsets = offsets * dcp_world_size + dcp_rank + kv_indices_start
|
||||||
|
local_kv_indices_offsets = local_kv_indices_start + offsets
|
||||||
|
|
||||||
|
kv_values = tl.load(kv_indices + kv_indice_offsets, mask=mask)
|
||||||
|
tl.store(
|
||||||
|
local_kv_indices + local_kv_indices_offsets,
|
||||||
|
kv_values // dcp_world_size,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Partial-attention LSE correction (PR #14194, MLA path).
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
@triton.jit
|
||||||
|
def _correct_attn_cp_out_kernel(
|
||||||
|
outputs_ptr,
|
||||||
|
new_output_ptr,
|
||||||
|
lses_ptr,
|
||||||
|
vlse_ptr,
|
||||||
|
outputs_stride_B,
|
||||||
|
outputs_stride_H,
|
||||||
|
outputs_stride_D,
|
||||||
|
lses_stride_N,
|
||||||
|
lses_stride_B,
|
||||||
|
lses_stride_H,
|
||||||
|
new_outputs_stride_H,
|
||||||
|
new_outputs_stride_B,
|
||||||
|
new_outputs_stride_D,
|
||||||
|
lse_idx,
|
||||||
|
HEAD_DIM: tl.constexpr,
|
||||||
|
N_ROUNDED: tl.constexpr,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Apply the all-gathered lses to correct each local rank's attention
|
||||||
|
output. we still need perform a cross-rank reduction to obtain the
|
||||||
|
final attention output.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
outputs_ptr (triton.PointerType):
|
||||||
|
Pointer to input tensor of shape [ B, H, D ]
|
||||||
|
lses_ptr (triton.PointerType):
|
||||||
|
Pointer to input tensor of shape [ N, B, H ]
|
||||||
|
new_output_ptr (triton.PointerType):
|
||||||
|
Pointer to output tensor of shape [ H, B, D ]
|
||||||
|
vlse_ptr (triton.PointerType):
|
||||||
|
Pointer to output tensor of shape [ B, H ]
|
||||||
|
"""
|
||||||
|
batch_idx = tl.program_id(axis=0).to(tl.int64)
|
||||||
|
head_idx = tl.program_id(axis=1).to(tl.int64)
|
||||||
|
|
||||||
|
# Use int32 for offsets where possible to reduce register pressure
|
||||||
|
b_i32 = batch_idx.to(tl.int32)
|
||||||
|
h_i32 = head_idx.to(tl.int32)
|
||||||
|
|
||||||
|
# Vectorized load of LSE values: shape = [N]
|
||||||
|
num_n_offsets = tl.arange(0, N_ROUNDED)
|
||||||
|
lse_offsets = (
|
||||||
|
num_n_offsets * lses_stride_N + b_i32 * lses_stride_B + h_i32 * lses_stride_H
|
||||||
|
)
|
||||||
|
|
||||||
|
# Compute final LSE using online softmax algorithm (more numerically stable)
|
||||||
|
lse = tl.load(lses_ptr + lse_offsets)
|
||||||
|
|
||||||
|
# Replace NaN and inf with -inf for numerical stability
|
||||||
|
neg_inf = float("-inf")
|
||||||
|
lse = tl.where((lse != lse) | (lse == float("inf")), neg_inf, lse)
|
||||||
|
|
||||||
|
# Online softmax: find max, subtract, exp, sum, log
|
||||||
|
lse_max = tl.max(lse, axis=0)
|
||||||
|
lse_max = tl.where(lse_max == neg_inf, 0.0, lse_max)
|
||||||
|
lse = lse - lse_max
|
||||||
|
lse_exp = tl.exp2(lse)
|
||||||
|
lse_acc = tl.sum(lse_exp, axis=0)
|
||||||
|
final_lse = tl.log2(lse_acc) + lse_max
|
||||||
|
|
||||||
|
# Compute correction factor
|
||||||
|
lse_offset = lse_idx * lses_stride_N + b_i32 * lses_stride_B + h_i32 * lses_stride_H
|
||||||
|
local_lse = tl.load(lses_ptr + lse_offset)
|
||||||
|
lse_diff = local_lse - final_lse
|
||||||
|
lse_diff = tl.where(
|
||||||
|
(lse_diff != lse_diff) | (lse_diff == float("inf")),
|
||||||
|
neg_inf,
|
||||||
|
lse_diff,
|
||||||
|
)
|
||||||
|
factor = tl.exp2(lse_diff)
|
||||||
|
|
||||||
|
# Store final LSE
|
||||||
|
tl.store(vlse_ptr + b_i32 * lses_stride_B + h_i32 * lses_stride_H, final_lse)
|
||||||
|
|
||||||
|
# Load output with vectorized access: shape = [D]
|
||||||
|
d_offsets = tl.arange(0, HEAD_DIM)
|
||||||
|
output_offsets = (
|
||||||
|
batch_idx * outputs_stride_B
|
||||||
|
+ head_idx * outputs_stride_H
|
||||||
|
+ d_offsets * outputs_stride_D
|
||||||
|
)
|
||||||
|
|
||||||
|
new_output_offsets = (
|
||||||
|
head_idx * new_outputs_stride_H
|
||||||
|
+ batch_idx * new_outputs_stride_B
|
||||||
|
+ d_offsets * new_outputs_stride_D
|
||||||
|
)
|
||||||
|
# Apply correction and store
|
||||||
|
output = tl.load(outputs_ptr + output_offsets)
|
||||||
|
output = output * factor
|
||||||
|
tl.store(new_output_ptr + new_output_offsets, output)
|
||||||
|
|
||||||
|
|
||||||
|
class CPTritonContext:
|
||||||
|
"""The CPTritonContext is used to avoid recompilation of the Triton JIT."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.inner_kernel = None
|
||||||
|
|
||||||
|
def call_kernel(self, kernel, grid, *regular_args, **const_args):
|
||||||
|
if self.inner_kernel is None:
|
||||||
|
self.inner_kernel = kernel[grid](*regular_args, **const_args)
|
||||||
|
else:
|
||||||
|
self.inner_kernel[grid](*regular_args)
|
||||||
|
|
||||||
|
|
||||||
|
def correct_attn_out(
|
||||||
|
out: torch.Tensor,
|
||||||
|
lses: torch.Tensor,
|
||||||
|
cp_rank: int,
|
||||||
|
ctx: Optional[CPTritonContext],
|
||||||
|
new_output: torch.Tensor = None,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Correct the attention output using the all-gathered lses.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
out: Tensor of shape [ B, H, D ]
|
||||||
|
lses: Tensor of shape [ N, B, H ]
|
||||||
|
cp_rank: Current rank in the context-parallel group
|
||||||
|
ctx: Triton context to avoid recompilation
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (out, lse) with corrected attention and final log-sum-exp.
|
||||||
|
"""
|
||||||
|
if ctx is None:
|
||||||
|
ctx = CPTritonContext()
|
||||||
|
|
||||||
|
# --- Normalize to 3D views ---
|
||||||
|
if out.ndim == 4 and out.shape[1] == 1:
|
||||||
|
out = out.squeeze(1)
|
||||||
|
assert out.ndim == 3, f"expected out [B,H,D] or [B,1,H,D], got {tuple(out.shape)}"
|
||||||
|
|
||||||
|
if lses.ndim == 4 and lses.shape[-1] == 1:
|
||||||
|
lses = lses.squeeze(-1)
|
||||||
|
if lses.ndim == 4 and lses.shape[1] == 1:
|
||||||
|
lses = lses.squeeze(1)
|
||||||
|
assert lses.ndim == 3, (
|
||||||
|
f"expected lses [N,B,H] (optionally with a 1-sized extra dim), "
|
||||||
|
f"got {tuple(lses.shape)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
B, H, D = out.shape
|
||||||
|
N = lses.shape[0]
|
||||||
|
|
||||||
|
# Strides after we normalized shapes to 3-D views. The kernel computes
|
||||||
|
# offsets for `vlse_ptr` using lses_stride_B/H, so the output buffer must
|
||||||
|
# have the same B/H stride layout as a slice of `lses`.
|
||||||
|
o_sB, o_sH, o_sD = out.stride()
|
||||||
|
l_sN, l_sB, l_sH = lses.stride()
|
||||||
|
no_sH, no_sB, no_sD = new_output.stride()
|
||||||
|
# Allocate LSE with the same B/H strides as `lses` so writes land correctly
|
||||||
|
# even when `lses` is a non-contiguous view (e.g., 4-D to 3-D squeeze).
|
||||||
|
lse = torch.empty_strided(
|
||||||
|
(B, H), (l_sB, l_sH), device=lses.device, dtype=lses.dtype
|
||||||
|
)
|
||||||
|
|
||||||
|
# Kernel launch config
|
||||||
|
grid = (B, H, 1)
|
||||||
|
|
||||||
|
regular_args = (
|
||||||
|
out,
|
||||||
|
new_output,
|
||||||
|
lses,
|
||||||
|
lse,
|
||||||
|
o_sB,
|
||||||
|
o_sH,
|
||||||
|
o_sD,
|
||||||
|
l_sN,
|
||||||
|
l_sB,
|
||||||
|
l_sH,
|
||||||
|
no_sH,
|
||||||
|
no_sB,
|
||||||
|
no_sD,
|
||||||
|
cp_rank,
|
||||||
|
)
|
||||||
|
const_args = {"HEAD_DIM": D, "N_ROUNDED": N}
|
||||||
|
|
||||||
|
ctx.call_kernel(_correct_attn_cp_out_kernel, grid, *regular_args, **const_args)
|
||||||
|
return new_output, lse
|
||||||
@@ -0,0 +1,64 @@
|
|||||||
|
# Copyright 2023-2026 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
"""Pure index math for decode context parallel (DCP): per-rank lengths and
|
||||||
|
the owner-rule local-index filter."""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.distributed.parallel_state import get_dcp_rank, get_dcp_world_size
|
||||||
|
from sglang.srt.layers.dcp.comm import dcp_enabled
|
||||||
|
|
||||||
|
|
||||||
|
def get_dcp_lens(
|
||||||
|
lens: torch.Tensor,
|
||||||
|
dcp_size: int,
|
||||||
|
dcp_rank: int,
|
||||||
|
start: torch.Tensor | None = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Per-rank visible KV length under the owner rule pos % dcp_size == dcp_rank.
|
||||||
|
|
||||||
|
Superset implementation (PR #25090): supports both start=None and a per-request
|
||||||
|
`start` offset. update_local_kv_lens_for_dcp is the start=None special case.
|
||||||
|
"""
|
||||||
|
if dcp_size == 1:
|
||||||
|
return lens
|
||||||
|
if start is None:
|
||||||
|
return lens // dcp_size + (dcp_rank < lens % dcp_size)
|
||||||
|
|
||||||
|
first = start + torch.remainder(dcp_rank - start, dcp_size)
|
||||||
|
remaining = start + lens - first
|
||||||
|
return torch.clamp((remaining + dcp_size - 1) // dcp_size, min=0)
|
||||||
|
|
||||||
|
|
||||||
|
def filter_dcp_local_kv_indices(kv_indices: torch.Tensor):
|
||||||
|
if dcp_enabled():
|
||||||
|
kv_indices = (
|
||||||
|
kv_indices[kv_indices % get_dcp_world_size() == get_dcp_rank()]
|
||||||
|
// get_dcp_world_size()
|
||||||
|
)
|
||||||
|
return kv_indices
|
||||||
|
|
||||||
|
|
||||||
|
def update_local_kv_lens_for_dcp(kv_len_arr):
|
||||||
|
"""In-place per-rank KV length: the start=0 case of get_dcp_lens.
|
||||||
|
|
||||||
|
floor((len - rank - 1) / N) + 1 == len // N + (rank < len % N) for len >= 0
|
||||||
|
(bit-identical; see test/registered/cp/test_dcp_layout_unit.py). Kept as an
|
||||||
|
in-place mutation because callers (plan_dcp_decode_metadata, the FlashInfer-MLA
|
||||||
|
cuda-graph replay path) rely on it.
|
||||||
|
"""
|
||||||
|
if not dcp_enabled():
|
||||||
|
return
|
||||||
|
kv_len_arr.copy_(get_dcp_lens(kv_len_arr, get_dcp_world_size(), get_dcp_rank()))
|
||||||
@@ -0,0 +1,37 @@
|
|||||||
|
# Copyright 2023-2026 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
"""Per-forward metadata for decode context parallel (DCP)."""
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
# NOTE: This is intentionally a standalone dataclass, NOT a subclass of
|
||||||
|
# layers.cp.base.BaseContextParallelMetadata. It is preserved verbatim from #14194
|
||||||
|
# and is stored in its own ForwardBatch field (attn_dcp_metadata), separate from the
|
||||||
|
# prefill-CP attn_cp_metadata, so it never participates in the CP-v2 build_metadata
|
||||||
|
# contract today. Whether the decode metadata should re-parent onto
|
||||||
|
# BaseContextParallelMetadata is deferred to P2 (DecodeContextParallelStrategy); decide
|
||||||
|
# it there rather than coupling this relocation to the CP-v2 ABC.
|
||||||
|
@dataclass
|
||||||
|
class DecodeContextParallelMetadata:
|
||||||
|
# For decode context parallel
|
||||||
|
dcp_kv_indptr: Optional[torch.Tensor] = None
|
||||||
|
dcp_kv_buffer: Optional[torch.Tensor] = None
|
||||||
|
dcp_kv_indices: Optional[torch.Tensor] = None
|
||||||
|
dcp_local_prefix_kv_indices: Optional[torch.Tensor] = None
|
||||||
|
dcp_extend_prefix_lens_sum: Optional[int] = None
|
||||||
@@ -0,0 +1,188 @@
|
|||||||
|
# Copyright 2023-2026 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
"""Decode-CP metadata builders (PR #14194). P2 will wrap these as methods on
|
||||||
|
DecodeContextParallelStrategy; kept as functions here for behavior-preserving
|
||||||
|
relocation."""
|
||||||
|
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.distributed.parallel_state import get_dcp_rank, get_dcp_world_size
|
||||||
|
from sglang.srt.layers.dcp.comm import dcp_enabled
|
||||||
|
from sglang.srt.layers.dcp.kernels import (
|
||||||
|
create_dcp_kv_indices,
|
||||||
|
update_kv_lens_and_indices,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.dcp.layout import update_local_kv_lens_for_dcp
|
||||||
|
from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata
|
||||||
|
from sglang.srt.server_args import get_global_server_args
|
||||||
|
|
||||||
|
|
||||||
|
def prepare_decode_context_parallel_metadata(
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
extend_prefix_lens: torch.Tensor,
|
||||||
|
extend_prefix_lens_cpu: torch.Tensor,
|
||||||
|
extend_seq_lens: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
seq_lens_sum: int,
|
||||||
|
kv_buffer_shape: torch.Size,
|
||||||
|
kv_cache_dtype,
|
||||||
|
kv_cache_device,
|
||||||
|
create_chunked_prefix_cache_kv_indices_fn,
|
||||||
|
) -> Optional[DecodeContextParallelMetadata]:
|
||||||
|
if not dcp_enabled():
|
||||||
|
return None
|
||||||
|
# dcp_kv_buffer tokens' layout
|
||||||
|
# [ rank0_r1.prefix_tokens, rank1_r1.prefix_tokens, ..., rank7_r1.prefix_tokens,
|
||||||
|
# ...,
|
||||||
|
# rank0_rn.prefix_tokens, rank1_rn.prefix_tokens, ..., rank7_rn.prefix_tokens,
|
||||||
|
# r1.extend_tokens, r2.extent_tokens, rn.extend_tokens ]
|
||||||
|
extend_prefix_starts = torch.zeros(
|
||||||
|
len(seq_lens),
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=get_global_server_args().device,
|
||||||
|
)
|
||||||
|
extend_cu_prefix_lens = torch.zeros(
|
||||||
|
len(seq_lens) + 1,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=get_global_server_args().device,
|
||||||
|
)
|
||||||
|
extend_cu_prefix_lens[1:] = torch.cumsum(extend_prefix_lens, dim=0)
|
||||||
|
extend_cu_prefix_lens = extend_cu_prefix_lens[:-1]
|
||||||
|
extend_prefix_lens_sum = sum([i for i in extend_prefix_lens_cpu])
|
||||||
|
|
||||||
|
dcp_prefix_kv_indices = torch.empty(
|
||||||
|
sum(extend_prefix_lens_cpu),
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=get_global_server_args().device,
|
||||||
|
)
|
||||||
|
create_chunked_prefix_cache_kv_indices_fn[(len(seq_lens),)](
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
extend_prefix_starts,
|
||||||
|
extend_prefix_lens,
|
||||||
|
extend_cu_prefix_lens,
|
||||||
|
dcp_prefix_kv_indices,
|
||||||
|
req_to_token.shape[1],
|
||||||
|
)
|
||||||
|
dcp_kv_indptr = torch.zeros(
|
||||||
|
len(seq_lens) + 1,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=get_global_server_args().device,
|
||||||
|
)
|
||||||
|
dcp_kv_indptr[1:] = seq_lens.cumsum(dim=0)
|
||||||
|
dcp_kv_indptr = dcp_kv_indptr[: (len(seq_lens) + 1)]
|
||||||
|
dcp_kv_indices = torch.zeros(
|
||||||
|
seq_lens_sum,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=get_global_server_args().device,
|
||||||
|
)
|
||||||
|
|
||||||
|
extend_cu_lens = torch.zeros(
|
||||||
|
len(seq_lens) + 1,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=get_global_server_args().device,
|
||||||
|
)
|
||||||
|
extend_cu_lens[1:] = torch.cumsum(extend_seq_lens, dim=0)
|
||||||
|
extend_cu_lens = extend_cu_lens[:-1]
|
||||||
|
|
||||||
|
create_dcp_kv_indices[(len(seq_lens),)](
|
||||||
|
dcp_kv_indptr,
|
||||||
|
extend_seq_lens,
|
||||||
|
extend_cu_lens,
|
||||||
|
extend_prefix_lens,
|
||||||
|
extend_cu_prefix_lens,
|
||||||
|
dcp_kv_indices,
|
||||||
|
extend_prefix_lens_sum,
|
||||||
|
get_dcp_world_size(),
|
||||||
|
)
|
||||||
|
dcp_local_prefix_kv_indices = (
|
||||||
|
dcp_prefix_kv_indices[
|
||||||
|
dcp_prefix_kv_indices % get_dcp_world_size() == get_dcp_rank()
|
||||||
|
]
|
||||||
|
// get_dcp_world_size()
|
||||||
|
)
|
||||||
|
dcp_kv_buffer = torch.empty(
|
||||||
|
(
|
||||||
|
seq_lens_sum,
|
||||||
|
*kv_buffer_shape[1:],
|
||||||
|
),
|
||||||
|
dtype=kv_cache_dtype,
|
||||||
|
device=kv_cache_device,
|
||||||
|
)
|
||||||
|
attn_dcp_metadata = DecodeContextParallelMetadata(
|
||||||
|
dcp_kv_indptr=dcp_kv_indptr,
|
||||||
|
dcp_kv_buffer=dcp_kv_buffer,
|
||||||
|
dcp_kv_indices=dcp_kv_indices,
|
||||||
|
dcp_local_prefix_kv_indices=dcp_local_prefix_kv_indices,
|
||||||
|
dcp_extend_prefix_lens_sum=extend_prefix_lens_sum,
|
||||||
|
)
|
||||||
|
return attn_dcp_metadata
|
||||||
|
|
||||||
|
|
||||||
|
def plan_dcp_decode_metadata(
|
||||||
|
kv_lens: torch.Tensor,
|
||||||
|
kv_indptr: torch.Tensor,
|
||||||
|
kv_indices: torch.Tensor,
|
||||||
|
init_metadata_replay: bool,
|
||||||
|
fast_decode_kwargs: dict,
|
||||||
|
bs: int,
|
||||||
|
):
|
||||||
|
local_kv_lens = kv_lens.clone()
|
||||||
|
update_local_kv_lens_for_dcp(local_kv_lens)
|
||||||
|
local_kv_lens.clamp_(min=0)
|
||||||
|
|
||||||
|
if not init_metadata_replay:
|
||||||
|
max_local_len = (
|
||||||
|
int(local_kv_lens.max().item()) if local_kv_lens.numel() > 0 else 0
|
||||||
|
)
|
||||||
|
total_local_len = (
|
||||||
|
int(local_kv_lens.sum().item()) if local_kv_lens.numel() > 0 else 0
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
max_local_len = (
|
||||||
|
int(fast_decode_kwargs["kv_len_arr_cpu"].max().item())
|
||||||
|
if fast_decode_kwargs["kv_len_arr_cpu"].numel() > 0
|
||||||
|
else 0
|
||||||
|
)
|
||||||
|
total_local_len = (
|
||||||
|
int(fast_decode_kwargs["kv_len_arr_cpu"].sum().item())
|
||||||
|
if fast_decode_kwargs["kv_len_arr_cpu"].numel() > 0
|
||||||
|
else 0
|
||||||
|
)
|
||||||
|
local_kv_lens_cumsum = kv_indptr.new_zeros((bs + 1,))
|
||||||
|
local_kv_lens_cumsum[1 : bs + 1] = torch.cumsum(local_kv_lens, dim=0)
|
||||||
|
local_kv_indices = kv_indices.new_empty(total_local_len)
|
||||||
|
BLOCK_SIZE = 128
|
||||||
|
num_blocks = (
|
||||||
|
(max_local_len + BLOCK_SIZE - 1) // BLOCK_SIZE if max_local_len > 0 else 1
|
||||||
|
)
|
||||||
|
grid = (bs, num_blocks)
|
||||||
|
update_kv_lens_and_indices[grid](
|
||||||
|
kv_lens,
|
||||||
|
kv_indptr,
|
||||||
|
kv_indices,
|
||||||
|
local_kv_lens,
|
||||||
|
local_kv_lens_cumsum,
|
||||||
|
local_kv_indices,
|
||||||
|
dcp_rank=get_dcp_rank(),
|
||||||
|
dcp_world_size=get_dcp_world_size(),
|
||||||
|
BLOCK_SIZE=BLOCK_SIZE,
|
||||||
|
)
|
||||||
|
kv_indices[:total_local_len] = local_kv_indices[:total_local_len]
|
||||||
|
kv_lens.copy_(local_kv_lens)
|
||||||
|
kv_indptr[: bs + 1] = local_kv_lens_cumsum[: bs + 1]
|
||||||
@@ -1,724 +0,0 @@
|
|||||||
from dataclasses import dataclass
|
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
import torch
|
|
||||||
import triton
|
|
||||||
import triton.language as tl
|
|
||||||
|
|
||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
|
||||||
use_symmetric_memory,
|
|
||||||
)
|
|
||||||
from sglang.srt.distributed.parallel_state import (
|
|
||||||
GroupCoordinator,
|
|
||||||
get_dcp_group,
|
|
||||||
get_dcp_group_no_assert,
|
|
||||||
get_dcp_rank,
|
|
||||||
get_dcp_world_size,
|
|
||||||
)
|
|
||||||
from sglang.srt.server_args import get_global_server_args
|
|
||||||
from sglang.srt.utils import is_cuda
|
|
||||||
|
|
||||||
|
|
||||||
def dcp_enabled() -> bool:
|
|
||||||
"""
|
|
||||||
only checks whether dcp enabled for cuda platform
|
|
||||||
"""
|
|
||||||
if get_dcp_group_no_assert() is None:
|
|
||||||
return False
|
|
||||||
if not is_cuda():
|
|
||||||
return False
|
|
||||||
return get_dcp_world_size() > 1
|
|
||||||
|
|
||||||
|
|
||||||
def get_attention_dcp_group() -> GroupCoordinator:
|
|
||||||
return get_dcp_group()
|
|
||||||
|
|
||||||
|
|
||||||
def get_attention_dcp_world_size() -> int:
|
|
||||||
if not dcp_enabled():
|
|
||||||
return 1
|
|
||||||
return get_dcp_world_size()
|
|
||||||
|
|
||||||
|
|
||||||
def get_attention_dcp_rank() -> int:
|
|
||||||
if not dcp_enabled():
|
|
||||||
return 0
|
|
||||||
return get_dcp_rank()
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def _correct_attn_cp_out_kernel(
|
|
||||||
outputs_ptr,
|
|
||||||
new_output_ptr,
|
|
||||||
lses_ptr,
|
|
||||||
vlse_ptr,
|
|
||||||
outputs_stride_B,
|
|
||||||
outputs_stride_H,
|
|
||||||
outputs_stride_D,
|
|
||||||
lses_stride_N,
|
|
||||||
lses_stride_B,
|
|
||||||
lses_stride_H,
|
|
||||||
new_outputs_stride_H,
|
|
||||||
new_outputs_stride_B,
|
|
||||||
new_outputs_stride_D,
|
|
||||||
lse_idx,
|
|
||||||
HEAD_DIM: tl.constexpr,
|
|
||||||
N_ROUNDED: tl.constexpr,
|
|
||||||
):
|
|
||||||
"""
|
|
||||||
Apply the all-gathered lses to correct each local rank's attention
|
|
||||||
output. we still need perform a cross-rank reduction to obtain the
|
|
||||||
final attention output.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
outputs_ptr (triton.PointerType):
|
|
||||||
Pointer to input tensor of shape [ B, H, D ]
|
|
||||||
lses_ptr (triton.PointerType):
|
|
||||||
Pointer to input tensor of shape [ N, B, H ]
|
|
||||||
new_output_ptr (triton.PointerType):
|
|
||||||
Pointer to output tensor of shape [ H, B, D ]
|
|
||||||
vlse_ptr (triton.PointerType):
|
|
||||||
Pointer to output tensor of shape [ B, H ]
|
|
||||||
"""
|
|
||||||
batch_idx = tl.program_id(axis=0).to(tl.int64)
|
|
||||||
head_idx = tl.program_id(axis=1).to(tl.int64)
|
|
||||||
|
|
||||||
# Use int32 for offsets where possible to reduce register pressure
|
|
||||||
b_i32 = batch_idx.to(tl.int32)
|
|
||||||
h_i32 = head_idx.to(tl.int32)
|
|
||||||
|
|
||||||
# Vectorized load of LSE values: shape = [N]
|
|
||||||
num_n_offsets = tl.arange(0, N_ROUNDED)
|
|
||||||
lse_offsets = (
|
|
||||||
num_n_offsets * lses_stride_N + b_i32 * lses_stride_B + h_i32 * lses_stride_H
|
|
||||||
)
|
|
||||||
|
|
||||||
# Compute final LSE using online softmax algorithm (more numerically stable)
|
|
||||||
lse = tl.load(lses_ptr + lse_offsets)
|
|
||||||
|
|
||||||
# Replace NaN and inf with -inf for numerical stability
|
|
||||||
neg_inf = float("-inf")
|
|
||||||
lse = tl.where((lse != lse) | (lse == float("inf")), neg_inf, lse)
|
|
||||||
|
|
||||||
# Online softmax: find max, subtract, exp, sum, log
|
|
||||||
lse_max = tl.max(lse, axis=0)
|
|
||||||
lse_max = tl.where(lse_max == neg_inf, 0.0, lse_max)
|
|
||||||
lse = lse - lse_max
|
|
||||||
lse_exp = tl.exp2(lse)
|
|
||||||
lse_acc = tl.sum(lse_exp, axis=0)
|
|
||||||
final_lse = tl.log2(lse_acc) + lse_max
|
|
||||||
|
|
||||||
# Compute correction factor
|
|
||||||
lse_offset = lse_idx * lses_stride_N + b_i32 * lses_stride_B + h_i32 * lses_stride_H
|
|
||||||
local_lse = tl.load(lses_ptr + lse_offset)
|
|
||||||
lse_diff = local_lse - final_lse
|
|
||||||
lse_diff = tl.where(
|
|
||||||
(lse_diff != lse_diff) | (lse_diff == float("inf")),
|
|
||||||
neg_inf,
|
|
||||||
lse_diff,
|
|
||||||
)
|
|
||||||
factor = tl.exp2(lse_diff)
|
|
||||||
|
|
||||||
# Store final LSE
|
|
||||||
tl.store(vlse_ptr + b_i32 * lses_stride_B + h_i32 * lses_stride_H, final_lse)
|
|
||||||
|
|
||||||
# Load output with vectorized access: shape = [D]
|
|
||||||
d_offsets = tl.arange(0, HEAD_DIM)
|
|
||||||
output_offsets = (
|
|
||||||
batch_idx * outputs_stride_B
|
|
||||||
+ head_idx * outputs_stride_H
|
|
||||||
+ d_offsets * outputs_stride_D
|
|
||||||
)
|
|
||||||
|
|
||||||
new_output_offsets = (
|
|
||||||
head_idx * new_outputs_stride_H
|
|
||||||
+ batch_idx * new_outputs_stride_B
|
|
||||||
+ d_offsets * new_outputs_stride_D
|
|
||||||
)
|
|
||||||
# Apply correction and store
|
|
||||||
output = tl.load(outputs_ptr + output_offsets)
|
|
||||||
output = output * factor
|
|
||||||
tl.store(new_output_ptr + new_output_offsets, output)
|
|
||||||
|
|
||||||
|
|
||||||
class CPTritonContext:
|
|
||||||
"""The CPTritonContext is used to avoid recompilation of the Triton JIT."""
|
|
||||||
|
|
||||||
def __init__(self):
|
|
||||||
self.inner_kernel = None
|
|
||||||
|
|
||||||
def call_kernel(self, kernel, grid, *regular_args, **const_args):
|
|
||||||
if self.inner_kernel is None:
|
|
||||||
self.inner_kernel = kernel[grid](*regular_args, **const_args)
|
|
||||||
else:
|
|
||||||
self.inner_kernel[grid](*regular_args)
|
|
||||||
|
|
||||||
|
|
||||||
def correct_attn_out(
|
|
||||||
out: torch.Tensor,
|
|
||||||
lses: torch.Tensor,
|
|
||||||
cp_rank: int,
|
|
||||||
ctx: Optional[CPTritonContext],
|
|
||||||
new_output: torch.Tensor = None,
|
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
||||||
"""Correct the attention output using the all-gathered lses.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
out: Tensor of shape [ B, H, D ]
|
|
||||||
lses: Tensor of shape [ N, B, H ]
|
|
||||||
cp_rank: Current rank in the context-parallel group
|
|
||||||
ctx: Triton context to avoid recompilation
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (out, lse) with corrected attention and final log-sum-exp.
|
|
||||||
"""
|
|
||||||
if ctx is None:
|
|
||||||
ctx = CPTritonContext()
|
|
||||||
|
|
||||||
# --- Normalize to 3D views ---
|
|
||||||
if out.ndim == 4 and out.shape[1] == 1:
|
|
||||||
out = out.squeeze(1)
|
|
||||||
assert out.ndim == 3, f"expected out [B,H,D] or [B,1,H,D], got {tuple(out.shape)}"
|
|
||||||
|
|
||||||
if lses.ndim == 4 and lses.shape[-1] == 1:
|
|
||||||
lses = lses.squeeze(-1)
|
|
||||||
if lses.ndim == 4 and lses.shape[1] == 1:
|
|
||||||
lses = lses.squeeze(1)
|
|
||||||
assert lses.ndim == 3, (
|
|
||||||
f"expected lses [N,B,H] (optionally with a 1-sized extra dim), "
|
|
||||||
f"got {tuple(lses.shape)}"
|
|
||||||
)
|
|
||||||
|
|
||||||
B, H, D = out.shape
|
|
||||||
N = lses.shape[0]
|
|
||||||
|
|
||||||
# Strides after we normalized shapes to 3-D views. The kernel computes
|
|
||||||
# offsets for `vlse_ptr` using lses_stride_B/H, so the output buffer must
|
|
||||||
# have the same B/H stride layout as a slice of `lses`.
|
|
||||||
o_sB, o_sH, o_sD = out.stride()
|
|
||||||
l_sN, l_sB, l_sH = lses.stride()
|
|
||||||
no_sH, no_sB, no_sD = new_output.stride()
|
|
||||||
# Allocate LSE with the same B/H strides as `lses` so writes land correctly
|
|
||||||
# even when `lses` is a non-contiguous view (e.g., 4-D to 3-D squeeze).
|
|
||||||
lse = torch.empty_strided(
|
|
||||||
(B, H), (l_sB, l_sH), device=lses.device, dtype=lses.dtype
|
|
||||||
)
|
|
||||||
|
|
||||||
# Kernel launch config
|
|
||||||
grid = (B, H, 1)
|
|
||||||
|
|
||||||
regular_args = (
|
|
||||||
out,
|
|
||||||
new_output,
|
|
||||||
lses,
|
|
||||||
lse,
|
|
||||||
o_sB,
|
|
||||||
o_sH,
|
|
||||||
o_sD,
|
|
||||||
l_sN,
|
|
||||||
l_sB,
|
|
||||||
l_sH,
|
|
||||||
no_sH,
|
|
||||||
no_sB,
|
|
||||||
no_sD,
|
|
||||||
cp_rank,
|
|
||||||
)
|
|
||||||
const_args = {"HEAD_DIM": D, "N_ROUNDED": N}
|
|
||||||
|
|
||||||
ctx.call_kernel(_correct_attn_cp_out_kernel, grid, *regular_args, **const_args)
|
|
||||||
return new_output, lse
|
|
||||||
|
|
||||||
|
|
||||||
def cp_lse_ag_out_rs(
|
|
||||||
cp_attn_out: torch.Tensor,
|
|
||||||
cp_attn_lse: torch.Tensor,
|
|
||||||
cp_group: GroupCoordinator,
|
|
||||||
ctx: Optional[CPTritonContext] = None,
|
|
||||||
):
|
|
||||||
"""
|
|
||||||
cp_attn_out: [ B, H, D ]
|
|
||||||
cp_attn_lse: [ B, H ]
|
|
||||||
"""
|
|
||||||
if cp_group.world_size == 1:
|
|
||||||
return cp_attn_out
|
|
||||||
|
|
||||||
if ctx is None:
|
|
||||||
ctx = CPTritonContext()
|
|
||||||
|
|
||||||
with use_symmetric_memory(cp_group):
|
|
||||||
# cp_attn_out is [B,H,D], we want to transpose it to [H,B,D] for the kernel, and then transpose back after correction.
|
|
||||||
new_output = cp_attn_out.new_empty(
|
|
||||||
cp_attn_out.transpose(0, 1).shape, dtype=torch.float32
|
|
||||||
)
|
|
||||||
cp_attn_lse = cp_attn_lse.to(torch.float32)
|
|
||||||
lses = cp_group.all_gather(cp_attn_lse, dim=0).view(
|
|
||||||
(cp_group.world_size,) + cp_attn_lse.shape
|
|
||||||
)
|
|
||||||
out, _ = correct_attn_out(
|
|
||||||
cp_attn_out, lses, cp_group.rank_in_group, ctx, new_output
|
|
||||||
)
|
|
||||||
out = cp_group.reduce_scatter_along_dim(out, dim=0)
|
|
||||||
return out.to(cp_attn_out.dtype)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def create_dcp_kv_indices(
|
|
||||||
kv_indptr,
|
|
||||||
extend_lens_ptr,
|
|
||||||
extend_cu_lens_ptr,
|
|
||||||
extend_prefix_lens_ptr,
|
|
||||||
extend_cu_prefix_lens_ptr,
|
|
||||||
kv_indices_ptr,
|
|
||||||
extend_prefix_lens_sum,
|
|
||||||
dcp_world_size: tl.constexpr,
|
|
||||||
):
|
|
||||||
BLOCK_SIZE: tl.constexpr = 512
|
|
||||||
pid = tl.program_id(axis=0)
|
|
||||||
prefix_len = tl.load(extend_prefix_lens_ptr + pid)
|
|
||||||
prefix_start = tl.load(extend_cu_prefix_lens_ptr + pid)
|
|
||||||
kv_ind_start = tl.load(kv_indptr + pid)
|
|
||||||
num_loop = tl.cdiv(prefix_len, BLOCK_SIZE)
|
|
||||||
for i in range(num_loop):
|
|
||||||
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
|
||||||
mask = offset < prefix_len
|
|
||||||
data = prefix_start + offset
|
|
||||||
tl.store(kv_indices_ptr + kv_ind_start + offset, data, mask=mask)
|
|
||||||
extend_len = tl.load(extend_lens_ptr + pid)
|
|
||||||
extend_start = tl.load(extend_cu_lens_ptr + pid)
|
|
||||||
num_loop = tl.cdiv(extend_len, BLOCK_SIZE)
|
|
||||||
for i in range(num_loop):
|
|
||||||
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
|
||||||
mask = offset < extend_len
|
|
||||||
data = extend_prefix_lens_sum + extend_start + offset
|
|
||||||
tl.store(
|
|
||||||
kv_indices_ptr + kv_ind_start + prefix_len + offset,
|
|
||||||
data,
|
|
||||||
mask=mask,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def update_kv_lens_and_indices(
|
|
||||||
kv_lens: torch.Tensor,
|
|
||||||
kv_lens_cumsum: torch.Tensor,
|
|
||||||
kv_indices: torch.Tensor,
|
|
||||||
local_kv_lens: torch.Tensor,
|
|
||||||
local_kv_lens_cumsum: torch.Tensor,
|
|
||||||
local_kv_indices: torch.Tensor,
|
|
||||||
dcp_rank: tl.constexpr,
|
|
||||||
dcp_world_size: tl.constexpr,
|
|
||||||
BLOCK_SIZE: tl.constexpr,
|
|
||||||
):
|
|
||||||
bs_idx = tl.program_id(0)
|
|
||||||
block_idx = tl.program_id(1)
|
|
||||||
|
|
||||||
local_kv_len = tl.load(local_kv_lens + bs_idx)
|
|
||||||
local_kv_indices_start = tl.load(local_kv_lens_cumsum + bs_idx)
|
|
||||||
kv_indices_start = tl.load(kv_lens_cumsum + bs_idx)
|
|
||||||
|
|
||||||
block_start = block_idx * BLOCK_SIZE
|
|
||||||
offsets = block_start + tl.arange(0, BLOCK_SIZE)
|
|
||||||
|
|
||||||
mask = offsets < local_kv_len
|
|
||||||
|
|
||||||
kv_indice_offsets = offsets * dcp_world_size + dcp_rank + kv_indices_start
|
|
||||||
local_kv_indices_offsets = local_kv_indices_start + offsets
|
|
||||||
|
|
||||||
kv_values = tl.load(kv_indices + kv_indice_offsets, mask=mask)
|
|
||||||
tl.store(
|
|
||||||
local_kv_indices + local_kv_indices_offsets,
|
|
||||||
kv_values // dcp_world_size,
|
|
||||||
mask=mask,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class DecodeContextParallelMetadata:
|
|
||||||
# For decode context parallel
|
|
||||||
dcp_kv_indptr: Optional[torch.Tensor] = None
|
|
||||||
dcp_kv_buffer: Optional[torch.Tensor] = None
|
|
||||||
dcp_kv_indices: Optional[torch.Tensor] = None
|
|
||||||
dcp_local_prefix_kv_indices: Optional[torch.Tensor] = None
|
|
||||||
dcp_extend_prefix_lens_sum: Optional[int] = None
|
|
||||||
|
|
||||||
|
|
||||||
def prepare_decode_context_parallel_metadata(
|
|
||||||
seq_lens: torch.Tensor,
|
|
||||||
extend_prefix_lens: torch.Tensor,
|
|
||||||
extend_prefix_lens_cpu: torch.Tensor,
|
|
||||||
extend_seq_lens: torch.Tensor,
|
|
||||||
req_pool_indices: torch.Tensor,
|
|
||||||
req_to_token: torch.Tensor,
|
|
||||||
seq_lens_sum: int,
|
|
||||||
kv_buffer_shape: torch.Size,
|
|
||||||
kv_cache_dtype,
|
|
||||||
kv_cache_device,
|
|
||||||
create_chunked_prefix_cache_kv_indices_fn,
|
|
||||||
) -> Optional[DecodeContextParallelMetadata]:
|
|
||||||
if not dcp_enabled():
|
|
||||||
return None
|
|
||||||
# dcp_kv_buffer tokens' layout
|
|
||||||
# [ rank0_r1.prefix_tokens, rank1_r1.prefix_tokens, ..., rank7_r1.prefix_tokens,
|
|
||||||
# ...,
|
|
||||||
# rank0_rn.prefix_tokens, rank1_rn.prefix_tokens, ..., rank7_rn.prefix_tokens,
|
|
||||||
# r1.extend_tokens, r2.extent_tokens, rn.extend_tokens ]
|
|
||||||
extend_prefix_starts = torch.zeros(
|
|
||||||
len(seq_lens),
|
|
||||||
dtype=torch.int32,
|
|
||||||
device=get_global_server_args().device,
|
|
||||||
)
|
|
||||||
extend_cu_prefix_lens = torch.zeros(
|
|
||||||
len(seq_lens) + 1,
|
|
||||||
dtype=torch.int32,
|
|
||||||
device=get_global_server_args().device,
|
|
||||||
)
|
|
||||||
extend_cu_prefix_lens[1:] = torch.cumsum(extend_prefix_lens, dim=0)
|
|
||||||
extend_cu_prefix_lens = extend_cu_prefix_lens[:-1]
|
|
||||||
extend_prefix_lens_sum = sum([i for i in extend_prefix_lens_cpu])
|
|
||||||
|
|
||||||
dcp_prefix_kv_indices = torch.empty(
|
|
||||||
sum(extend_prefix_lens_cpu),
|
|
||||||
dtype=torch.int32,
|
|
||||||
device=get_global_server_args().device,
|
|
||||||
)
|
|
||||||
create_chunked_prefix_cache_kv_indices_fn[(len(seq_lens),)](
|
|
||||||
req_to_token,
|
|
||||||
req_pool_indices,
|
|
||||||
extend_prefix_starts,
|
|
||||||
extend_prefix_lens,
|
|
||||||
extend_cu_prefix_lens,
|
|
||||||
dcp_prefix_kv_indices,
|
|
||||||
req_to_token.shape[1],
|
|
||||||
)
|
|
||||||
dcp_kv_indptr = torch.zeros(
|
|
||||||
len(seq_lens) + 1,
|
|
||||||
dtype=torch.int32,
|
|
||||||
device=get_global_server_args().device,
|
|
||||||
)
|
|
||||||
dcp_kv_indptr[1:] = seq_lens.cumsum(dim=0)
|
|
||||||
dcp_kv_indptr = dcp_kv_indptr[: (len(seq_lens) + 1)]
|
|
||||||
dcp_kv_indices = torch.zeros(
|
|
||||||
seq_lens_sum,
|
|
||||||
dtype=torch.int32,
|
|
||||||
device=get_global_server_args().device,
|
|
||||||
)
|
|
||||||
|
|
||||||
extend_cu_lens = torch.zeros(
|
|
||||||
len(seq_lens) + 1,
|
|
||||||
dtype=torch.int32,
|
|
||||||
device=get_global_server_args().device,
|
|
||||||
)
|
|
||||||
extend_cu_lens[1:] = torch.cumsum(extend_seq_lens, dim=0)
|
|
||||||
extend_cu_lens = extend_cu_lens[:-1]
|
|
||||||
|
|
||||||
create_dcp_kv_indices[(len(seq_lens),)](
|
|
||||||
dcp_kv_indptr,
|
|
||||||
extend_seq_lens,
|
|
||||||
extend_cu_lens,
|
|
||||||
extend_prefix_lens,
|
|
||||||
extend_cu_prefix_lens,
|
|
||||||
dcp_kv_indices,
|
|
||||||
extend_prefix_lens_sum,
|
|
||||||
get_dcp_world_size(),
|
|
||||||
)
|
|
||||||
dcp_local_prefix_kv_indices = (
|
|
||||||
dcp_prefix_kv_indices[
|
|
||||||
dcp_prefix_kv_indices % get_dcp_world_size() == get_dcp_rank()
|
|
||||||
]
|
|
||||||
// get_dcp_world_size()
|
|
||||||
)
|
|
||||||
dcp_kv_buffer = torch.empty(
|
|
||||||
(
|
|
||||||
seq_lens_sum,
|
|
||||||
*kv_buffer_shape[1:],
|
|
||||||
),
|
|
||||||
dtype=kv_cache_dtype,
|
|
||||||
device=kv_cache_device,
|
|
||||||
)
|
|
||||||
attn_dcp_metadata = DecodeContextParallelMetadata(
|
|
||||||
dcp_kv_indptr=dcp_kv_indptr,
|
|
||||||
dcp_kv_buffer=dcp_kv_buffer,
|
|
||||||
dcp_kv_indices=dcp_kv_indices,
|
|
||||||
dcp_local_prefix_kv_indices=dcp_local_prefix_kv_indices,
|
|
||||||
dcp_extend_prefix_lens_sum=extend_prefix_lens_sum,
|
|
||||||
)
|
|
||||||
return attn_dcp_metadata
|
|
||||||
|
|
||||||
|
|
||||||
def _all_gather_dcp_kv_cache(kv_a: torch.Tensor):
|
|
||||||
dcp_world_size = get_dcp_world_size()
|
|
||||||
# not use symmetric_memory unless torch mem_pool updated, see https://github.com/pytorch/pytorch/issues/178138
|
|
||||||
gathered_kv_a = kv_a.new_empty(
|
|
||||||
(kv_a.shape[0] * dcp_world_size, *kv_a.shape[1:]),
|
|
||||||
)
|
|
||||||
get_dcp_group().all_gather_into_tensor(gathered_kv_a, kv_a)
|
|
||||||
gathered_kv_a = (
|
|
||||||
gathered_kv_a.reshape((dcp_world_size,) + kv_a.shape)
|
|
||||||
.transpose(0, 1)
|
|
||||||
.reshape(-1, *kv_a.shape[1:])
|
|
||||||
)
|
|
||||||
return gathered_kv_a
|
|
||||||
|
|
||||||
|
|
||||||
def all_gather_kv_cache_for_mha_chunk_extend(
|
|
||||||
kv_a: torch.Tensor,
|
|
||||||
k_pe: torch.Tensor,
|
|
||||||
prefix_kv_lens_cpu: torch.Tensor,
|
|
||||||
prefix_starts_cpu: torch.Tensor = None,
|
|
||||||
):
|
|
||||||
if dcp_enabled():
|
|
||||||
kv_a = kv_a.unsqueeze(1)
|
|
||||||
gathered_kv = all_gather_kv_cache_for_dcp(
|
|
||||||
kv_a,
|
|
||||||
k_pe,
|
|
||||||
prefix_kv_lens_cpu,
|
|
||||||
prefix_starts_cpu,
|
|
||||||
)
|
|
||||||
kv_a, k_pe = gathered_kv.split([kv_a.shape[-1], k_pe.shape[-1]], dim=-1)
|
|
||||||
kv_a = kv_a.squeeze(1)
|
|
||||||
return kv_a.contiguous(), k_pe.contiguous()
|
|
||||||
|
|
||||||
|
|
||||||
def all_gather_kv_cache_for_mha_extend(
|
|
||||||
token_to_kv_pool,
|
|
||||||
attn_mqa,
|
|
||||||
dcp_local_prefix_kv_indices,
|
|
||||||
seq_lens,
|
|
||||||
extend_prefix_lens,
|
|
||||||
extend_prefix_lens_cpu: list[int],
|
|
||||||
extend_seq_lens,
|
|
||||||
kv_a: torch.Tensor,
|
|
||||||
k_pe: torch.Tensor,
|
|
||||||
):
|
|
||||||
prefix_kv_a, prefix_k_pe = token_to_kv_pool.get_mla_kv_buffer(
|
|
||||||
attn_mqa, dcp_local_prefix_kv_indices
|
|
||||||
)
|
|
||||||
extend_prefix_lens_cpu = torch.tensor(extend_prefix_lens_cpu)
|
|
||||||
gathered_kv_cache = all_gather_kv_cache_for_dcp(
|
|
||||||
prefix_kv_a,
|
|
||||||
prefix_k_pe,
|
|
||||||
extend_prefix_lens_cpu,
|
|
||||||
)
|
|
||||||
prefix_kv_a, prefix_k_pe = gathered_kv_cache.split(
|
|
||||||
[kv_a.shape[-1], k_pe.shape[-1]], dim=-1
|
|
||||||
)
|
|
||||||
prefix_kv_a = prefix_kv_a.squeeze(1)
|
|
||||||
# re-organize kv with query orders
|
|
||||||
prefix_lens_cu = torch.zeros(
|
|
||||||
len(seq_lens) + 1,
|
|
||||||
dtype=torch.int32,
|
|
||||||
device=kv_a.device,
|
|
||||||
)
|
|
||||||
extend_lens_cu = torch.zeros_like(prefix_lens_cu)
|
|
||||||
prefix_lens_cu[1:] = torch.cumsum(extend_prefix_lens, dim=0)
|
|
||||||
extend_lens_cu[1:] = torch.cumsum(extend_seq_lens, dim=0)
|
|
||||||
kv_a_tuple = ()
|
|
||||||
k_pe_tuple = ()
|
|
||||||
for i in range(len(seq_lens)):
|
|
||||||
kv_a_tuple += (
|
|
||||||
prefix_kv_a[prefix_lens_cu[i] : prefix_lens_cu[i + 1]],
|
|
||||||
kv_a[extend_lens_cu[i] : extend_lens_cu[i + 1]],
|
|
||||||
)
|
|
||||||
k_pe_tuple += (
|
|
||||||
prefix_k_pe[prefix_lens_cu[i] : prefix_lens_cu[i + 1]],
|
|
||||||
k_pe[extend_lens_cu[i] : extend_lens_cu[i + 1]],
|
|
||||||
)
|
|
||||||
kv_a = torch.cat(kv_a_tuple, dim=0)
|
|
||||||
k_pe = torch.cat(k_pe_tuple, dim=0)
|
|
||||||
return kv_a.contiguous(), k_pe.contiguous()
|
|
||||||
|
|
||||||
|
|
||||||
def filter_dcp_local_kv_indices(kv_indices: torch.Tensor):
|
|
||||||
if dcp_enabled():
|
|
||||||
kv_indices = (
|
|
||||||
kv_indices[kv_indices % get_dcp_world_size() == get_dcp_rank()]
|
|
||||||
// get_dcp_world_size()
|
|
||||||
)
|
|
||||||
return kv_indices
|
|
||||||
|
|
||||||
|
|
||||||
def all_gather_q_for_mla_decode(
|
|
||||||
q_nope_out: torch.Tensor,
|
|
||||||
q_pe: torch.Tensor,
|
|
||||||
):
|
|
||||||
with use_symmetric_memory(get_dcp_group()):
|
|
||||||
# transpose q_pe and q_nope_out from [B, H, L] to [H, B, L]
|
|
||||||
combined = torch.cat([q_pe.transpose(0, 1), q_nope_out.transpose(0, 1)], dim=-1)
|
|
||||||
gathered = get_dcp_group().all_gather(combined, dim=0)
|
|
||||||
d_pe = q_pe.size(-1)
|
|
||||||
d_nope = q_nope_out.size(-1)
|
|
||||||
q_pe, q_nope_out = gathered.split([d_pe, d_nope], dim=-1)
|
|
||||||
q_pe = q_pe.transpose(0, 1)
|
|
||||||
q_nope_out = q_nope_out.transpose(0, 1)
|
|
||||||
return q_nope_out, q_pe
|
|
||||||
|
|
||||||
|
|
||||||
def all_gather_kv_cache_for_mla_extend(
|
|
||||||
token_to_kv_pool,
|
|
||||||
attn_mqa,
|
|
||||||
extend_prefix_lens_cpu: list[int],
|
|
||||||
dcp_local_prefix_kv_indices,
|
|
||||||
dcp_extend_prefix_lens_sum,
|
|
||||||
dcp_kv_buffer,
|
|
||||||
kv_lora_rank,
|
|
||||||
k_nope,
|
|
||||||
k_pe,
|
|
||||||
):
|
|
||||||
cache_k_nope, cache_k_rope = token_to_kv_pool.get_mla_kv_buffer(
|
|
||||||
attn_mqa,
|
|
||||||
dcp_local_prefix_kv_indices,
|
|
||||||
)
|
|
||||||
extend_prefix_lens_cpu = torch.tensor(extend_prefix_lens_cpu)
|
|
||||||
# all gather kv cache into forward_batch.attn_dcp_metadata.dcp_kv_buffer
|
|
||||||
gathered_kv = all_gather_kv_cache_for_dcp(
|
|
||||||
cache_k_nope,
|
|
||||||
cache_k_rope,
|
|
||||||
extend_prefix_lens_cpu,
|
|
||||||
prefix_starts_cpu=torch.zeros_like(extend_prefix_lens_cpu),
|
|
||||||
)
|
|
||||||
dcp_kv_buffer[:dcp_extend_prefix_lens_sum] = gathered_kv
|
|
||||||
|
|
||||||
# copy local kv cache into forward_batch.attn_dcp_metadata.dcp_kv_buffer
|
|
||||||
dcp_kv_buffer[
|
|
||||||
dcp_extend_prefix_lens_sum:,
|
|
||||||
...,
|
|
||||||
:kv_lora_rank,
|
|
||||||
] = k_nope
|
|
||||||
dcp_kv_buffer[
|
|
||||||
dcp_extend_prefix_lens_sum:,
|
|
||||||
...,
|
|
||||||
kv_lora_rank:,
|
|
||||||
] = k_pe
|
|
||||||
|
|
||||||
|
|
||||||
def update_local_kv_lens_for_dcp(kv_len_arr):
|
|
||||||
if not dcp_enabled():
|
|
||||||
return
|
|
||||||
dcp_world_size = get_dcp_world_size()
|
|
||||||
dcp_rank = get_dcp_rank()
|
|
||||||
offset = dcp_rank + 1
|
|
||||||
kv_len_arr.sub_(offset).div_(dcp_world_size, rounding_mode="floor").add_(1)
|
|
||||||
|
|
||||||
|
|
||||||
def plan_dcp_decode_metadata(
|
|
||||||
kv_lens: torch.Tensor,
|
|
||||||
kv_indptr: torch.Tensor,
|
|
||||||
kv_indices: torch.Tensor,
|
|
||||||
init_metadata_replay: bool,
|
|
||||||
fast_decode_kwargs: dict,
|
|
||||||
bs: int,
|
|
||||||
):
|
|
||||||
local_kv_lens = kv_lens.clone()
|
|
||||||
update_local_kv_lens_for_dcp(local_kv_lens)
|
|
||||||
local_kv_lens.clamp_(min=0)
|
|
||||||
|
|
||||||
if not init_metadata_replay:
|
|
||||||
max_local_len = (
|
|
||||||
int(local_kv_lens.max().item()) if local_kv_lens.numel() > 0 else 0
|
|
||||||
)
|
|
||||||
total_local_len = (
|
|
||||||
int(local_kv_lens.sum().item()) if local_kv_lens.numel() > 0 else 0
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
max_local_len = (
|
|
||||||
int(fast_decode_kwargs["kv_len_arr_cpu"].max().item())
|
|
||||||
if fast_decode_kwargs["kv_len_arr_cpu"].numel() > 0
|
|
||||||
else 0
|
|
||||||
)
|
|
||||||
total_local_len = (
|
|
||||||
int(fast_decode_kwargs["kv_len_arr_cpu"].sum().item())
|
|
||||||
if fast_decode_kwargs["kv_len_arr_cpu"].numel() > 0
|
|
||||||
else 0
|
|
||||||
)
|
|
||||||
local_kv_lens_cumsum = kv_indptr.new_zeros((bs + 1,))
|
|
||||||
local_kv_lens_cumsum[1 : bs + 1] = torch.cumsum(local_kv_lens, dim=0)
|
|
||||||
local_kv_indices = kv_indices.new_empty(total_local_len)
|
|
||||||
BLOCK_SIZE = 128
|
|
||||||
num_blocks = (
|
|
||||||
(max_local_len + BLOCK_SIZE - 1) // BLOCK_SIZE if max_local_len > 0 else 1
|
|
||||||
)
|
|
||||||
grid = (bs, num_blocks)
|
|
||||||
update_kv_lens_and_indices[grid](
|
|
||||||
kv_lens,
|
|
||||||
kv_indptr,
|
|
||||||
kv_indices,
|
|
||||||
local_kv_lens,
|
|
||||||
local_kv_lens_cumsum,
|
|
||||||
local_kv_indices,
|
|
||||||
dcp_rank=get_dcp_rank(),
|
|
||||||
dcp_world_size=get_dcp_world_size(),
|
|
||||||
BLOCK_SIZE=BLOCK_SIZE,
|
|
||||||
)
|
|
||||||
kv_indices[:total_local_len] = local_kv_indices[:total_local_len]
|
|
||||||
kv_lens.copy_(local_kv_lens)
|
|
||||||
kv_indptr[: bs + 1] = local_kv_lens_cumsum[: bs + 1]
|
|
||||||
|
|
||||||
|
|
||||||
# all gather kv cache and re-org to query orders
|
|
||||||
def all_gather_kv_cache_for_dcp(
|
|
||||||
prefix_kv_a: torch.Tensor,
|
|
||||||
prefix_k_pe: torch.Tensor,
|
|
||||||
prefix_kv_lens_cpu: torch.Tensor,
|
|
||||||
prefix_starts_cpu: torch.Tensor = None,
|
|
||||||
):
|
|
||||||
"""
|
|
||||||
prefix_kv_a and prefix_k_pe should have same shape, expect for last dim
|
|
||||||
"""
|
|
||||||
if not dcp_enabled():
|
|
||||||
return torch.cat([prefix_kv_a, prefix_k_pe], dim=-1)
|
|
||||||
# 1. compute max kv_lens for each seq
|
|
||||||
dcp_world_size = get_dcp_world_size()
|
|
||||||
dcp_rank = get_dcp_rank()
|
|
||||||
|
|
||||||
if prefix_starts_cpu is None:
|
|
||||||
prefix_starts_cpu = torch.zeros_like(prefix_kv_lens_cpu)
|
|
||||||
|
|
||||||
left_pads = prefix_starts_cpu % dcp_world_size > dcp_rank
|
|
||||||
left_pads = left_pads.to(torch.int32)
|
|
||||||
right_pads = (
|
|
||||||
prefix_starts_cpu + prefix_kv_lens_cpu - 1
|
|
||||||
) % dcp_world_size < dcp_rank
|
|
||||||
right_pads = right_pads.to(torch.int32)
|
|
||||||
padded_lens = (
|
|
||||||
prefix_kv_lens_cpu + (prefix_starts_cpu % dcp_world_size) + dcp_world_size - 1
|
|
||||||
) // dcp_world_size
|
|
||||||
|
|
||||||
local_kv_lens = padded_lens - left_pads - right_pads
|
|
||||||
local_kv_lens_cu = torch.zeros(
|
|
||||||
len(prefix_kv_lens_cpu) + 1,
|
|
||||||
dtype=torch.int32,
|
|
||||||
)
|
|
||||||
local_kv_lens_cu[1:] = torch.cumsum(local_kv_lens, dim=0)
|
|
||||||
|
|
||||||
padded_kv_cache_arr = []
|
|
||||||
prefix_kv_cache = torch.cat([prefix_kv_a, prefix_k_pe], dim=-1)
|
|
||||||
for req_idx in range(len(prefix_kv_lens_cpu)):
|
|
||||||
padded_tensor = prefix_kv_cache.new_empty(
|
|
||||||
(padded_lens[req_idx].item(),) + prefix_kv_cache.size()[1:]
|
|
||||||
)
|
|
||||||
padded_tensor[
|
|
||||||
left_pads[req_idx] : left_pads[req_idx] + local_kv_lens[req_idx]
|
|
||||||
] = prefix_kv_cache[local_kv_lens_cu[req_idx] : local_kv_lens_cu[req_idx + 1]]
|
|
||||||
padded_kv_cache_arr.append(padded_tensor)
|
|
||||||
|
|
||||||
padded_kv_cache = torch.cat(padded_kv_cache_arr, dim=0)
|
|
||||||
|
|
||||||
gatherd_kv_cache = _all_gather_dcp_kv_cache(padded_kv_cache)
|
|
||||||
|
|
||||||
# 2. re-org kv cache to query orders
|
|
||||||
padded_lens_cu = torch.zeros(
|
|
||||||
len(prefix_kv_lens_cpu) + 1,
|
|
||||||
dtype=torch.int32,
|
|
||||||
)
|
|
||||||
padded_lens_cu[1:] = torch.cumsum(padded_lens, dim=0)
|
|
||||||
kv_cache_tuple = ()
|
|
||||||
for req_idx in range(len(prefix_kv_lens_cpu)):
|
|
||||||
kv_cache_tuple += (
|
|
||||||
gatherd_kv_cache[
|
|
||||||
padded_lens_cu[req_idx] * dcp_world_size
|
|
||||||
+ (prefix_starts_cpu[req_idx] % dcp_world_size) :
|
|
||||||
][: prefix_kv_lens_cpu[req_idx]],
|
|
||||||
)
|
|
||||||
gatherd_kv_cache = torch.cat(kv_cache_tuple, dim=0)
|
|
||||||
|
|
||||||
return gatherd_kv_cache
|
|
||||||
@@ -25,9 +25,7 @@ from typing import TYPE_CHECKING
|
|||||||
|
|
||||||
from sglang.srt.configs.model_config import ModelImpl
|
from sglang.srt.configs.model_config import ModelImpl
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.utils.dcp_utils import (
|
from sglang.srt.layers.dcp import dcp_enabled
|
||||||
dcp_enabled,
|
|
||||||
)
|
|
||||||
from sglang.srt.managers.mm_utils import init_mm_embedding_cache
|
from sglang.srt.managers.mm_utils import init_mm_embedding_cache
|
||||||
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
||||||
from sglang.srt.mem_cache.registry import TreeCacheBuildContext, create_tree_cache
|
from sglang.srt.mem_cache.registry import TreeCacheBuildContext, create_tree_cache
|
||||||
|
|||||||
@@ -44,13 +44,13 @@ from sglang.srt.layers.attention.dsa.quant_k_cache import (
|
|||||||
quantize_k_cache_separate,
|
quantize_k_cache_separate,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.dsa.utils import aiter_can_use_preshuffle_paged_mqa
|
from sglang.srt.layers.attention.dsa.utils import aiter_can_use_preshuffle_paged_mqa
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz
|
from sglang.srt.layers.dcp import (
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
|
||||||
from sglang.srt.layers.utils.dcp_utils import (
|
|
||||||
dcp_enabled,
|
dcp_enabled,
|
||||||
get_attention_dcp_rank,
|
get_attention_dcp_rank,
|
||||||
get_attention_dcp_world_size,
|
get_attention_dcp_world_size,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz
|
||||||
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
from sglang.srt.mem_cache.allocator.mamba import MambaSlotAllocator
|
from sglang.srt.mem_cache.allocator.mamba import MambaSlotAllocator
|
||||||
from sglang.srt.mem_cache.layout.page_major import (
|
from sglang.srt.mem_cache.layout.page_major import (
|
||||||
build_page_major_mamba_views,
|
build_page_major_mamba_views,
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ import triton
|
|||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
from sglang.jit_kernel.utils import is_arch_support_pdl
|
from sglang.jit_kernel.utils import is_arch_support_pdl
|
||||||
from sglang.srt.layers.utils.dcp_utils import (
|
from sglang.srt.layers.dcp import (
|
||||||
dcp_enabled,
|
dcp_enabled,
|
||||||
get_attention_dcp_rank,
|
get_attention_dcp_rank,
|
||||||
get_attention_dcp_world_size,
|
get_attention_dcp_world_size,
|
||||||
|
|||||||
@@ -45,7 +45,6 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
set_dp_buffer_len,
|
set_dp_buffer_len,
|
||||||
set_is_extend_in_batch,
|
set_is_extend_in_batch,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.utils.dcp_utils import DecodeContextParallelMetadata
|
|
||||||
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
|
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
|
||||||
ForwardBatchDeepSeekMHAMixin,
|
ForwardBatchDeepSeekMHAMixin,
|
||||||
)
|
)
|
||||||
@@ -61,6 +60,7 @@ from sglang.srt.utils import (
|
|||||||
from sglang.srt.utils.common import ceil_align, is_pin_memory_available
|
from sglang.srt.utils.common import ceil_align, is_pin_memory_available
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||||
from sglang.srt.layers.utils.cp_utils import ContextParallelMetadata
|
from sglang.srt.layers.utils.cp_utils import ContextParallelMetadata
|
||||||
from sglang.srt.managers.schedule_batch import MultimodalInputs, ScheduleBatch
|
from sglang.srt.managers.schedule_batch import MultimodalInputs, ScheduleBatch
|
||||||
@@ -504,7 +504,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
|
|
||||||
attn_cp_metadata: Optional[ContextParallelMetadata] = None
|
attn_cp_metadata: Optional[ContextParallelMetadata] = None
|
||||||
|
|
||||||
# For decode context parallel
|
# For decode context parallel.
|
||||||
|
# NOTE: DecodeContextParallelMetadata is imported under TYPE_CHECKING only (see the
|
||||||
|
# import block above) — available for annotations but NOT bound at runtime in this
|
||||||
|
# module. Import it from sglang.srt.layers.dcp.metadata if a runtime use is added.
|
||||||
attn_dcp_metadata: Optional[DecodeContextParallelMetadata] = None
|
attn_dcp_metadata: Optional[DecodeContextParallelMetadata] = None
|
||||||
|
|
||||||
# Decode context parallel KV write mask.
|
# Decode context parallel KV write mask.
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from sglang.srt.layers.attention.dsa.dequant_k_cache import dequantize_k_cache_p
|
|||||||
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
|
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
|
||||||
from sglang.srt.layers.attention.utils import concat_and_cast_mha_k_triton
|
from sglang.srt.layers.attention.utils import concat_and_cast_mha_k_triton
|
||||||
from sglang.srt.layers.communicator import get_attn_tp_context
|
from sglang.srt.layers.communicator import get_attn_tp_context
|
||||||
from sglang.srt.layers.utils.dcp_utils import (
|
from sglang.srt.layers.dcp import (
|
||||||
all_gather_kv_cache_for_mha_chunk_extend,
|
all_gather_kv_cache_for_mha_chunk_extend,
|
||||||
all_gather_kv_cache_for_mha_extend,
|
all_gather_kv_cache_for_mha_extend,
|
||||||
dcp_enabled,
|
dcp_enabled,
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.compilation.compilation_config import register_split_op
|
from sglang.srt.compilation.compilation_config import register_split_op
|
||||||
|
from sglang.srt.distributed.parallel_state import get_dcp_group
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers import deep_gemm_wrapper
|
from sglang.srt.layers import deep_gemm_wrapper
|
||||||
from sglang.srt.layers.attention.dsa.utils import (
|
from sglang.srt.layers.attention.dsa.utils import (
|
||||||
@@ -14,6 +15,13 @@ from sglang.srt.layers.attention.dsa.utils import (
|
|||||||
is_graph_dsa_split_op_surface,
|
is_graph_dsa_split_op_surface,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.communicator import get_attn_tp_context
|
from sglang.srt.layers.communicator import get_attn_tp_context
|
||||||
|
from sglang.srt.layers.dcp import (
|
||||||
|
all_gather_kv_cache_for_mla_extend,
|
||||||
|
all_gather_q_for_mla_decode,
|
||||||
|
cp_lse_ag_out_rs_mla,
|
||||||
|
dcp_enabled,
|
||||||
|
get_attention_dcp_world_size,
|
||||||
|
)
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import (
|
from sglang.srt.layers.quantization.fp8_kernel import (
|
||||||
fp8_dtype,
|
fp8_dtype,
|
||||||
per_tensor_quant_mla_fp8,
|
per_tensor_quant_mla_fp8,
|
||||||
@@ -21,14 +29,6 @@ from sglang.srt.layers.quantization.fp8_kernel import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers.radix_attention import unified_attention_with_output
|
from sglang.srt.layers.radix_attention import unified_attention_with_output
|
||||||
from sglang.srt.layers.utils.cp_utils import mla_use_prefill_cp
|
from sglang.srt.layers.utils.cp_utils import mla_use_prefill_cp
|
||||||
from sglang.srt.layers.utils.dcp_utils import (
|
|
||||||
all_gather_kv_cache_for_mla_extend,
|
|
||||||
all_gather_q_for_mla_decode,
|
|
||||||
cp_lse_ag_out_rs,
|
|
||||||
dcp_enabled,
|
|
||||||
get_attention_dcp_group,
|
|
||||||
get_attention_dcp_world_size,
|
|
||||||
)
|
|
||||||
from sglang.srt.lora.deepseek_mla_correction import (
|
from sglang.srt.lora.deepseek_mla_correction import (
|
||||||
apply_q_correction as apply_kv_b_lora_q_correction,
|
apply_q_correction as apply_kv_b_lora_q_correction,
|
||||||
)
|
)
|
||||||
@@ -789,7 +789,7 @@ class DeepseekMLAForwardMixin:
|
|||||||
self.num_local_heads * get_attention_dcp_world_size(),
|
self.num_local_heads * get_attention_dcp_world_size(),
|
||||||
self.kv_lora_rank,
|
self.kv_lora_rank,
|
||||||
)
|
)
|
||||||
attn_output = cp_lse_ag_out_rs(attn_output, lse, get_attention_dcp_group())
|
attn_output = cp_lse_ag_out_rs_mla(attn_output, lse, get_dcp_group())
|
||||||
attn_output = attn_output.transpose(0, 1)
|
attn_output = attn_output.transpose(0, 1)
|
||||||
attn_output = attn_output.view(-1, self.num_local_heads, self.kv_lora_rank)
|
attn_output = attn_output.view(-1, self.num_local_heads, self.kv_lora_rank)
|
||||||
|
|
||||||
|
|||||||
@@ -71,6 +71,10 @@ from sglang.srt.layers.communicator import (
|
|||||||
get_attn_tp_context,
|
get_attn_tp_context,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.communicator_dsa_cp import DSACPLayerCommunicator
|
from sglang.srt.layers.communicator_dsa_cp import DSACPLayerCommunicator
|
||||||
|
from sglang.srt.layers.dcp import dcp_enabled, get_attention_dcp_world_size
|
||||||
|
from sglang.srt.layers.dcp.planner import (
|
||||||
|
prepare_decode_context_parallel_metadata,
|
||||||
|
)
|
||||||
from sglang.srt.layers.layernorm import RMSNorm
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
from sglang.srt.layers.linear import (
|
from sglang.srt.layers.linear import (
|
||||||
ColumnParallelLinear,
|
ColumnParallelLinear,
|
||||||
@@ -123,11 +127,6 @@ from sglang.srt.layers.utils.cp_utils import (
|
|||||||
mla_use_prefill_cp,
|
mla_use_prefill_cp,
|
||||||
prepare_context_parallel_metadata,
|
prepare_context_parallel_metadata,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.utils.dcp_utils import (
|
|
||||||
dcp_enabled,
|
|
||||||
get_attention_dcp_world_size,
|
|
||||||
prepare_decode_context_parallel_metadata,
|
|
||||||
)
|
|
||||||
from sglang.srt.layers.vocab_parallel_embedding import (
|
from sglang.srt.layers.vocab_parallel_embedding import (
|
||||||
ParallelLMHead,
|
ParallelLMHead,
|
||||||
VocabParallelEmbedding,
|
VocabParallelEmbedding,
|
||||||
|
|||||||
@@ -0,0 +1,91 @@
|
|||||||
|
"""CPU unit test for the decode-context-parallel (DCP) per-rank KV-length math.
|
||||||
|
|
||||||
|
Pins ``get_dcp_lens`` (the single, superset implementation in
|
||||||
|
``layers/dcp/layout.py``) to a brute-force owner-count reference, and proves
|
||||||
|
it is bit-identical to the legacy in-place formula that
|
||||||
|
``update_local_kv_lens_for_dcp`` used before it was collapsed into a wrapper:
|
||||||
|
|
||||||
|
floor((len - rank - 1) / N) + 1 == len // N + (rank < len % N) (len >= 0)
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python -m pytest test_dcp_layout_unit.py -v
|
||||||
|
python test_dcp_layout_unit.py
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
DCP_SIZES = [1, 2, 3, 4, 8]
|
||||||
|
LENS = list(range(0, 41))
|
||||||
|
STARTS = [0, 1, 2, 5, 7, 13, 31]
|
||||||
|
|
||||||
|
|
||||||
|
def _owner_count(length: int, n: int, rank: int, start: int) -> int:
|
||||||
|
"""Ground truth: # of absolute positions p in [start, start+length) with p % n == rank."""
|
||||||
|
return sum(1 for p in range(start, start + length) if p % n == rank)
|
||||||
|
|
||||||
|
|
||||||
|
def _legacy_inplace_formula(length: int, n: int, rank: int) -> int:
|
||||||
|
"""The pre-refactor update_local_kv_lens_for_dcp body (start == 0 case)."""
|
||||||
|
return (length - rank - 1) // n + 1
|
||||||
|
|
||||||
|
|
||||||
|
class TestGetDcpLens(unittest.TestCase):
|
||||||
|
def test_start_none_matches_owner_count(self):
|
||||||
|
for n in DCP_SIZES:
|
||||||
|
for rank in range(n):
|
||||||
|
lens = torch.tensor(LENS, dtype=torch.int32)
|
||||||
|
got = get_dcp_lens(lens, n, rank)
|
||||||
|
expected = torch.tensor(
|
||||||
|
[_owner_count(L, n, rank, 0) for L in LENS], dtype=torch.int32
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(got.to(torch.int32), expected),
|
||||||
|
f"start=None mismatch at n={n}, rank={rank}: {got.tolist()} != {expected.tolist()}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_start_none_matches_legacy_inplace_formula(self):
|
||||||
|
# The collapse claim: get_dcp_lens (start=None) == legacy floor((L-rank-1)/N)+1.
|
||||||
|
for n in DCP_SIZES:
|
||||||
|
for rank in range(n):
|
||||||
|
lens = torch.tensor(LENS, dtype=torch.int64)
|
||||||
|
got = get_dcp_lens(lens, n, rank)
|
||||||
|
legacy = torch.tensor(
|
||||||
|
[_legacy_inplace_formula(L, n, rank) for L in LENS],
|
||||||
|
dtype=torch.int64,
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(got.to(torch.int64), legacy),
|
||||||
|
f"legacy-formula mismatch at n={n}, rank={rank}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_start_tensor_matches_owner_count(self):
|
||||||
|
for n in DCP_SIZES:
|
||||||
|
for rank in range(n):
|
||||||
|
for start in STARTS:
|
||||||
|
lens = torch.tensor(LENS, dtype=torch.int64)
|
||||||
|
start_t = torch.full_like(lens, start)
|
||||||
|
got = get_dcp_lens(lens, n, rank, start=start_t)
|
||||||
|
expected = torch.tensor(
|
||||||
|
[_owner_count(L, n, rank, start) for L in LENS],
|
||||||
|
dtype=torch.int64,
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(got.to(torch.int64), expected),
|
||||||
|
f"start={start} mismatch at n={n}, rank={rank}: "
|
||||||
|
f"{got.tolist()} != {expected.tolist()}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_dcp_size_one_is_identity(self):
|
||||||
|
lens = torch.tensor(LENS, dtype=torch.int32)
|
||||||
|
self.assertTrue(torch.equal(get_dcp_lens(lens, 1, 0), lens))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -126,7 +126,7 @@ class TestDSV31DCP8TP8GSM8K(GSM8KMixin, BasicDecodeCorrectnessMixin, CustomTestC
|
|||||||
|
|
||||||
This test exercises the full DCP decode and extend paths:
|
This test exercises the full DCP decode and extend paths:
|
||||||
- Decode: query all-gather → attention on local KV shard → LSE
|
- Decode: query all-gather → attention on local KV shard → LSE
|
||||||
correction via cp_lse_ag_out_rs → reduce-scatter
|
correction via cp_lse_ag_out_rs_mla → reduce-scatter
|
||||||
- Extend (prefill): all-gather prefix KV cache across DCP ranks,
|
- Extend (prefill): all-gather prefix KV cache across DCP ranks,
|
||||||
attend with full context
|
attend with full context
|
||||||
|
|
||||||
@@ -205,7 +205,7 @@ class TestDSV31DCP8LogprobParity(BasicDecodeCorrectnessMixin, CustomTestCase):
|
|||||||
introduces small numerical differences)
|
introduces small numerical differences)
|
||||||
|
|
||||||
This catches subtle correctness bugs in the DCP LSE correction path
|
This catches subtle correctness bugs in the DCP LSE correction path
|
||||||
(cp_lse_ag_out_rs) that a coarse GSM8K accuracy gate cannot detect.
|
(cp_lse_ag_out_rs_mla) that a coarse GSM8K accuracy gate cannot detect.
|
||||||
For example, if exp2/exp mismatch causes a systematic bias in the
|
For example, if exp2/exp mismatch causes a systematic bias in the
|
||||||
attention output, logprobs will diverge by more than the tolerance.
|
attention output, logprobs will diverge by more than the tolerance.
|
||||||
"""
|
"""
|
||||||
|
|||||||
Reference in New Issue
Block a user