[AMD][DCP 1/N] add dcp support for aiter backend (#34432)
Co-authored-by: HAI <hixiao@gmail.com>
This commit is contained in:
@@ -42,6 +42,16 @@ def _require_kimi_k3_cutedsl_dcp_support() -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _require_kimi_k3_aiter_gluon_dcp_support() -> None:
|
||||||
|
from sglang.srt.layers.attention.aiter_mla_gluon import _gluon_fn
|
||||||
|
|
||||||
|
if _gluon_fn() is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Kimi-K3 DCP with decode_attention_backend='aiter' requires the aiter "
|
||||||
|
"gluon mla kernel, which is unavailable. See above aborting reasons."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@_register_for("KimiK3ForConditionalGeneration")
|
@_register_for("KimiK3ForConditionalGeneration")
|
||||||
def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
|
def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||||
cfg = resolving_view(server_args)
|
cfg = resolving_view(server_args)
|
||||||
@@ -101,9 +111,23 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
|
|||||||
decode_attention_backend="tokenspeed_mla",
|
decode_attention_backend="tokenspeed_mla",
|
||||||
kv_cache_dtype="fp8_e4m3",
|
kv_cache_dtype="fp8_e4m3",
|
||||||
)
|
)
|
||||||
|
elif decode_backend == "aiter":
|
||||||
|
_require_kimi_k3_aiter_gluon_dcp_support()
|
||||||
|
# Override prefill backend to aiter by default
|
||||||
|
# if users don't explicitly specify triton
|
||||||
|
prefill_ab = "triton" if prefill_backend == "triton" else "aiter"
|
||||||
|
logger.info(
|
||||||
|
"Kimi-K3 DCP uses aiter MLA decode: "
|
||||||
|
f"prefill={prefill_backend!r} -> {prefill_ab!r}, "
|
||||||
|
f"decode={decode_backend!r} -> 'aiter'."
|
||||||
|
)
|
||||||
|
overrides.update(
|
||||||
|
prefill_attention_backend=prefill_ab,
|
||||||
|
decode_attention_backend="aiter",
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
raise AssertionError(
|
raise AssertionError(
|
||||||
f"Decode attention backend for Kimi-K3 DCP must be 'cutedsl_mla' or 'tokenspeed_mla', got {decode_backend!r}."
|
f"Decode attention backend for Kimi-K3 DCP must be 'cutedsl_mla', 'tokenspeed_mla' or 'aiter', got {decode_backend!r}."
|
||||||
)
|
)
|
||||||
|
|
||||||
if cfg.dcp_replicate_q_proj is None:
|
if cfg.dcp_replicate_q_proj is None:
|
||||||
|
|||||||
@@ -26,6 +26,8 @@ from sglang.kernels.ops.kvcache.aiter_unified_attention import (
|
|||||||
scatter_req_to_token_to_page_table_kernel,
|
scatter_req_to_token_to_page_table_kernel,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
|
from sglang.srt.layers.dcp import update_local_kv_lens_for_dcp
|
||||||
|
from sglang.srt.layers.dcp.planner import plan_dcp_decode_metadata
|
||||||
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
|
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
|
||||||
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.speculative.spec_utils import (
|
from sglang.srt.speculative.spec_utils import (
|
||||||
@@ -67,6 +69,8 @@ except ImportError:
|
|||||||
"aiter is AMD specific kernel library. Please make sure aiter is installed on your AMD device."
|
"aiter is AMD specific kernel library. Please make sure aiter is installed on your AMD device."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
from sglang.kernels.ops.attention.dcp_kernels import create_mla_kv_page_table_for_dcp
|
||||||
|
from sglang.kernels.ops.attention.merge_state import merge_state_triton
|
||||||
from sglang.kernels.ops.attention.utils import (
|
from sglang.kernels.ops.attention.utils import (
|
||||||
launch_reshape_and_cache_flash,
|
launch_reshape_and_cache_flash,
|
||||||
pad_sequence_with_mask,
|
pad_sequence_with_mask,
|
||||||
@@ -151,11 +155,16 @@ class ForwardMetadata:
|
|||||||
swa_page_table: Optional[torch.Tensor] = None
|
swa_page_table: Optional[torch.Tensor] = None
|
||||||
# full->SWA translated out_cache_loc (SWA KV-store write target)
|
# full->SWA translated out_cache_loc (SWA KV-store write target)
|
||||||
swa_out_cache_loc: Optional[torch.Tensor] = None
|
swa_out_cache_loc: Optional[torch.Tensor] = None
|
||||||
|
local_kv_lens: Optional[torch.Tensor] = None
|
||||||
|
verify_token_table: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
|
||||||
_AITER_PARTITION_SIZE_ROCM = 256
|
_AITER_PARTITION_SIZE_ROCM = 256
|
||||||
|
|
||||||
|
|
||||||
|
_DCP_VERIFY_TABLE_COLS_PER_BLOCK = 128
|
||||||
|
|
||||||
|
|
||||||
# AITER's gfx950 FP8 FMHA ASM kernels only cover these GQA ratios. Other
|
# AITER's gfx950 FP8 FMHA ASM kernels only cover these GQA ratios. Other
|
||||||
# ratios (e.g. Qwen3.8-27B 24Q/4KV = 6) must not take the pertensor shortcut.
|
# ratios (e.g. Qwen3.8-27B 24Q/4KV = 6) must not take the pertensor shortcut.
|
||||||
_AITER_FP8_ASM_GQA_RATIOS = frozenset({1, 2, 4, 8, 16})
|
_AITER_FP8_ASM_GQA_RATIOS = frozenset({1, 2, 4, 8, 16})
|
||||||
@@ -264,6 +273,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
|
|
||||||
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
||||||
|
|
||||||
|
self.dcp_world_size = get_parallel().attn_dcp_size
|
||||||
|
|
||||||
# Get v_head_dim based on model type
|
# Get v_head_dim based on model type
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
# For MLA models, get v_head_dim from model config
|
# For MLA models, get v_head_dim from model config
|
||||||
@@ -348,6 +359,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
# corresponding ForwardBatch fields.
|
# corresponding ForwardBatch fields.
|
||||||
self.req_to_token_pool = model_runner.req_to_token_pool
|
self.req_to_token_pool = model_runner.req_to_token_pool
|
||||||
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||||
|
self.kv_index_translator = model_runner.kv_index_translator
|
||||||
|
|
||||||
# sliding window attention. Resolve the SWA pool rather than reading it
|
# sliding window attention. Resolve the SWA pool rather than reading it
|
||||||
# straight off the active pool: a frozen-KV MTP draft worker's active
|
# straight off the active pool: a frozen-KV MTP draft worker's active
|
||||||
@@ -433,13 +445,19 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
# _mla_decode_fwd_with_head_pad brings any count below 16 up to it,
|
# _mla_decode_fwd_with_head_pad brings any count below 16 up to it,
|
||||||
# by repetition when it divides 16 and by tiling otherwise.
|
# by repetition when it divides 16 and by tiling otherwise.
|
||||||
_pad_heads_to_16 = self.num_head < 16
|
_pad_heads_to_16 = self.num_head < 16
|
||||||
assert _valid_heads or _pad_heads_to_16 or not may_run_mla_decode, (
|
assert (
|
||||||
|
self.dcp_world_size > 1
|
||||||
|
or _valid_heads
|
||||||
|
or _pad_heads_to_16
|
||||||
|
or not may_run_mla_decode
|
||||||
|
), (
|
||||||
f"Aiter MLA supports num_head of 4, 8, 12, or multiples of 16 "
|
f"Aiter MLA supports num_head of 4, 8, 12, or multiples of 16 "
|
||||||
f"in [16, 128].\n"
|
f"in [16, 128].\n"
|
||||||
f"Provided {self.num_head} number of heads.\n"
|
f"Provided {self.num_head} number of heads.\n"
|
||||||
"Try adjusting tensor_parallel_size value, or run decode on "
|
"Try adjusting tensor_parallel_size value, or run decode on "
|
||||||
"another backend (--decode-attention-backend)."
|
"another backend (--decode-attention-backend)."
|
||||||
)
|
)
|
||||||
|
|
||||||
self.num_head_padded = 16 if self.num_head < 16 else self.num_head
|
self.num_head_padded = 16 if self.num_head < 16 else self.num_head
|
||||||
if self.num_head in _mla_low_head_repeat:
|
if self.num_head in _mla_low_head_repeat:
|
||||||
self.head_pad_mode = "repeat"
|
self.head_pad_mode = "repeat"
|
||||||
@@ -451,17 +469,22 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
self.head_pad_mode = "none"
|
self.head_pad_mode = "none"
|
||||||
self.head_repeat_factor = 1
|
self.head_repeat_factor = 1
|
||||||
|
|
||||||
|
_gathered_num_head = self.num_head * self.dcp_world_size
|
||||||
|
self.mla_kernel_num_head_padded = (
|
||||||
|
16 if _gathered_num_head < 16 else _gathered_num_head
|
||||||
|
)
|
||||||
|
|
||||||
self.enable_dp_attention = is_dp_attention_enabled()
|
self.enable_dp_attention = is_dp_attention_enabled()
|
||||||
self.qo_indptr_ = torch.zeros(
|
self.qo_indptr_ = torch.zeros(
|
||||||
(max_bs + 1,), dtype=torch.int32, device=model_runner.device
|
(max_bs + 1,), dtype=torch.int32, device=model_runner.device
|
||||||
)
|
)
|
||||||
global _use_mla_ps_kernel, fast_mode, intra_batch_mode
|
global _use_mla_ps_kernel, fast_mode, intra_batch_mode
|
||||||
|
|
||||||
# current mla_decode_fwd only support fake-nps in self.num_head == 16
|
# current mla_decode_fwd only support fake-nps in num_head == 16
|
||||||
# so all num_head size does not use qh16 kernel to simulate
|
# so all num_head size does not use qh16 kernel to simulate
|
||||||
# it should not use fake-nps (fast_mode = False, intra_batch_mode = True)
|
# it should not use fake-nps (fast_mode = False, intra_batch_mode = True)
|
||||||
# it will cause gpu-fault or accuracy issue
|
# it will cause gpu-fault or accuracy issue.
|
||||||
if self.num_head in (32, 64, 128):
|
if self.mla_kernel_num_head_padded in (32, 64, 128):
|
||||||
fast_mode = True
|
fast_mode = True
|
||||||
intra_batch_mode = False
|
intra_batch_mode = False
|
||||||
|
|
||||||
@@ -473,8 +496,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
# Native 16-head persist is slow on TP8; keep disabled unless zero-pad
|
# Native 16-head persist is slow on TP8; keep disabled unless zero-pad
|
||||||
# (e.g. Kimi K3 h12 -> qh16) where persist ASM is the fast path.
|
# (e.g. Kimi K3 h12 -> qh16) where persist ASM is the fast path.
|
||||||
if (
|
if (
|
||||||
(self.num_head_padded == 16 and self.head_pad_mode != "zero")
|
(self.mla_kernel_num_head_padded == 16 and self.head_pad_mode != "zero")
|
||||||
or self.num_head_padded == 128
|
or self.mla_kernel_num_head_padded == 128
|
||||||
) and self.kv_cache_dtype is not fp8_dtype:
|
) and self.kv_cache_dtype is not fp8_dtype:
|
||||||
_use_mla_ps_kernel = False
|
_use_mla_ps_kernel = False
|
||||||
fast_mode = False
|
fast_mode = False
|
||||||
@@ -551,7 +574,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
return "fp8_e4m3"
|
return "fp8_e4m3"
|
||||||
|
|
||||||
def make_mla_decode_meta_data_buffer(self, max_seqlen_qo, batch_size):
|
def make_mla_decode_meta_data_buffer(self, max_seqlen_qo, batch_size):
|
||||||
nhead = self.num_head_padded
|
nhead = self.mla_kernel_num_head_padded
|
||||||
dtype = self.kv_cache_dtype
|
dtype = self.kv_cache_dtype
|
||||||
|
|
||||||
if self.enable_dp_attention:
|
if self.enable_dp_attention:
|
||||||
@@ -637,7 +660,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
qo_indptr,
|
qo_indptr,
|
||||||
kv_indptr,
|
kv_indptr,
|
||||||
kv_last_page_len,
|
kv_last_page_len,
|
||||||
self.num_head_padded // nhead_kv,
|
self.mla_kernel_num_head_padded // nhead_kv,
|
||||||
nhead_kv,
|
nhead_kv,
|
||||||
False,
|
False,
|
||||||
work_metadata,
|
work_metadata,
|
||||||
@@ -1084,6 +1107,80 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
num_kv_splits=num_kv_splits,
|
num_kv_splits=num_kv_splits,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _get_dcp_graph_max_local_kv_len(self) -> int:
|
||||||
|
"""Static upper bound on this rank's shard, ceil(max_context_len / W)."""
|
||||||
|
w = max(self.dcp_world_size, 1)
|
||||||
|
return (self.max_context_len + w - 1) // w
|
||||||
|
|
||||||
|
def _forward_decode_dcp(self, q, k_buffer, layer, k_descale):
|
||||||
|
"""Attend this rank's KV shard for decode -> (out, natural-log lse)."""
|
||||||
|
fm = self.forward_metadata
|
||||||
|
bs = fm.kv_indptr.shape[0] - 1
|
||||||
|
num_heads = layer.tp_q_head_num # gathered heads = num_local_heads * dcp
|
||||||
|
|
||||||
|
out, lse = mla_gluon_decode(
|
||||||
|
q=q.view(bs, num_heads, layer.qk_head_dim),
|
||||||
|
k_buffer=k_buffer,
|
||||||
|
layer=layer,
|
||||||
|
kv_indices=fm.kv_indices,
|
||||||
|
kv_indptr=fm.kv_indptr[: bs + 1],
|
||||||
|
sm_scale=layer.scaling,
|
||||||
|
kv_scale=self._resolve_fp8_kv_scale_float(layer, k_descale),
|
||||||
|
min_kv_seq_len=1,
|
||||||
|
return_lse=True,
|
||||||
|
)
|
||||||
|
return out, lse.view(bs, num_heads)
|
||||||
|
|
||||||
|
def _forward_verify_dcp(self, q, k_window, layer, k_descale):
|
||||||
|
"""Attend the committed prefix and the verify window separately, then
|
||||||
|
merge -> (out, natural-log lse).
|
||||||
|
|
||||||
|
Splitting at the window boundary avoids the one thing the decode kernel
|
||||||
|
cannot do under DCP: mask on the GLOBAL position g(j) = j * W + r.
|
||||||
|
"""
|
||||||
|
fm = self.forward_metadata
|
||||||
|
q_len = fm.max_q_len
|
||||||
|
num_heads = layer.tp_q_head_num # gathered heads = num_local_heads * dcp
|
||||||
|
seqused_k = fm.local_kv_lens
|
||||||
|
n_rows = seqused_k.shape[0]
|
||||||
|
bs = n_rows // q_len
|
||||||
|
|
||||||
|
out_a, lse_a = mla_gluon_decode(
|
||||||
|
q=q.view(n_rows, num_heads, layer.qk_head_dim),
|
||||||
|
k_buffer=self.token_to_kv_pool.get_key_buffer(layer.layer_id),
|
||||||
|
layer=layer,
|
||||||
|
kv_indices=fm.verify_token_table,
|
||||||
|
kv_indptr=seqused_k,
|
||||||
|
sm_scale=layer.scaling,
|
||||||
|
kv_scale=self._resolve_fp8_kv_scale_float(layer, k_descale),
|
||||||
|
min_kv_seq_len=1,
|
||||||
|
return_lse=True,
|
||||||
|
use_2d_view=True,
|
||||||
|
)
|
||||||
|
lse_a = lse_a.view(n_rows, num_heads)
|
||||||
|
|
||||||
|
# The verify window arrives as `k_window`, computed this forward and
|
||||||
|
# identical on every rank, so only one rank attends it; the others
|
||||||
|
# return their prefix partial for the cross-rank merge.
|
||||||
|
if get_parallel().attn_dcp_rank != 0:
|
||||||
|
return out_a, lse_a
|
||||||
|
|
||||||
|
# The window latent is request-major and contiguous, so it IS the pool:
|
||||||
|
# row i of request b lives at b * q_len + i. mla_gluon's MTP mask at
|
||||||
|
# seq_len == qlen is exactly the dense causal window this needs.
|
||||||
|
out_b, lse_b = mla_gluon_decode(
|
||||||
|
q=q.view(n_rows, num_heads, layer.qk_head_dim),
|
||||||
|
k_buffer=k_window,
|
||||||
|
layer=layer,
|
||||||
|
kv_indices=torch.arange(n_rows, dtype=torch.int32, device=q.device),
|
||||||
|
kv_indptr=torch.arange(bs + 1, dtype=torch.int32, device=q.device) * q_len,
|
||||||
|
sm_scale=layer.scaling,
|
||||||
|
min_kv_seq_len=1,
|
||||||
|
qlen=q_len,
|
||||||
|
return_lse=True,
|
||||||
|
)
|
||||||
|
return merge_state_triton(out_a, lse_a, out_b, lse_b.view(n_rows, num_heads))
|
||||||
|
|
||||||
def mla_fp8_prefill_attn(
|
def mla_fp8_prefill_attn(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
@@ -1243,6 +1340,9 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
max_kv_len = forward_batch.seq_lens_cpu.max().item()
|
max_kv_len = forward_batch.seq_lens_cpu.max().item()
|
||||||
|
|
||||||
|
# dcp metadata
|
||||||
|
local_kv_lens = None
|
||||||
|
verify_token_table = None
|
||||||
if forward_batch.forward_mode.is_decode_or_idle():
|
if forward_batch.forward_mode.is_decode_or_idle():
|
||||||
if spec_info is None or forward_batch.forward_mode.is_idle():
|
if spec_info is None or forward_batch.forward_mode.is_idle():
|
||||||
kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
|
kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
|
||||||
@@ -1261,6 +1361,20 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_indices,
|
kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if (
|
||||||
|
self.use_mla
|
||||||
|
and self.dcp_world_size > 1
|
||||||
|
and not forward_batch.forward_mode.is_idle()
|
||||||
|
):
|
||||||
|
kv_lens = forward_batch.seq_lens[:bs].to(torch.int32).clone()
|
||||||
|
self._plan_dcp_decode_metadata(
|
||||||
|
kv_indptr,
|
||||||
|
kv_indices,
|
||||||
|
kv_lens,
|
||||||
|
forward_batch.seq_lens_cpu,
|
||||||
|
bs,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
max_q_len = 1
|
max_q_len = 1
|
||||||
page_size = self.page_size
|
page_size = self.page_size
|
||||||
@@ -1317,7 +1431,9 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_last_page_len = self.kv_last_page_len[:bs]
|
kv_last_page_len = self.kv_last_page_len[:bs]
|
||||||
max_q_len = 1
|
max_q_len = 1
|
||||||
|
|
||||||
if _use_mla_ps_kernel:
|
# DCP decode runs the aiter MLA kernel (builds its own block-table
|
||||||
|
# metadata in forward_decode), so skip the persist metadata.
|
||||||
|
if _use_mla_ps_kernel and self.dcp_world_size <= 1:
|
||||||
(
|
(
|
||||||
work_metadata,
|
work_metadata,
|
||||||
work_indptr,
|
work_indptr,
|
||||||
@@ -1456,9 +1572,13 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
elif forward_batch.forward_mode.is_target_verify():
|
elif forward_batch.forward_mode.is_target_verify():
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
draft_num = spec_info.draft_token_num
|
draft_num = spec_info.draft_token_num
|
||||||
|
device = forward_batch.seq_lens.device
|
||||||
|
if self.dcp_world_size > 1:
|
||||||
|
kv_lens = forward_batch.seq_lens.to(torch.int32).clone()
|
||||||
|
kv_lens_sum = forward_batch.seq_lens_sum
|
||||||
|
else:
|
||||||
kv_lens = forward_batch.seq_lens + draft_num
|
kv_lens = forward_batch.seq_lens + draft_num
|
||||||
kv_lens_sum = forward_batch.seq_lens_sum + draft_num * bs
|
kv_lens_sum = forward_batch.seq_lens_sum + draft_num * bs
|
||||||
device = forward_batch.seq_lens.device
|
|
||||||
|
|
||||||
qo_indptr = self.qo_indptr[: bs + 1]
|
qo_indptr = self.qo_indptr[: bs + 1]
|
||||||
qo_indptr[: bs + 1] = torch.arange(
|
qo_indptr[: bs + 1] = torch.arange(
|
||||||
@@ -1486,8 +1606,26 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
TOKEN_BLOCK_PARALLEL=num_token_blocks > 1,
|
TOKEN_BLOCK_PARALLEL=num_token_blocks > 1,
|
||||||
)
|
)
|
||||||
|
|
||||||
# if self.kv_cache_dtype == fp8_dtype:
|
if self.dcp_world_size > 1:
|
||||||
if _use_mla_ps_kernel:
|
self._plan_dcp_decode_metadata(
|
||||||
|
kv_indptr,
|
||||||
|
kv_indices,
|
||||||
|
kv_lens,
|
||||||
|
None,
|
||||||
|
bs,
|
||||||
|
)
|
||||||
|
(
|
||||||
|
verify_token_table,
|
||||||
|
local_kv_lens,
|
||||||
|
) = self._build_dcp_verify_token_table(
|
||||||
|
kv_indptr,
|
||||||
|
forward_batch.req_pool_indices,
|
||||||
|
bs,
|
||||||
|
draft_num,
|
||||||
|
(max_kv_len + self.dcp_world_size - 1) // self.dcp_world_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
if _use_mla_ps_kernel and self.dcp_world_size <= 1:
|
||||||
max_seqlen_qo = draft_num
|
max_seqlen_qo = draft_num
|
||||||
(
|
(
|
||||||
work_metadata,
|
work_metadata,
|
||||||
@@ -1532,6 +1670,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
reduce_partial_map=reduce_partial_map,
|
reduce_partial_map=reduce_partial_map,
|
||||||
num_kv_splits=num_kv_splits,
|
num_kv_splits=num_kv_splits,
|
||||||
run_graph=False,
|
run_graph=False,
|
||||||
|
local_kv_lens=local_kv_lens,
|
||||||
|
verify_token_table=verify_token_table,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
draft_num = forward_batch.input_ids.shape[0] // bs
|
draft_num = forward_batch.input_ids.shape[0] // bs
|
||||||
@@ -1712,6 +1852,132 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
swa_out_cache_loc=swa_out_cache_loc,
|
swa_out_cache_loc=swa_out_cache_loc,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _plan_dcp_decode_metadata(
|
||||||
|
self,
|
||||||
|
kv_indptr: torch.Tensor,
|
||||||
|
kv_indices: torch.Tensor,
|
||||||
|
kv_lens_gpu: torch.Tensor,
|
||||||
|
seq_lens_cpu: Optional[torch.Tensor],
|
||||||
|
bs: int,
|
||||||
|
static_local_kv_lens_cpu: Optional[torch.Tensor] = None,
|
||||||
|
):
|
||||||
|
"""Localize kv_indptr / kv_indices to this rank's DCP shard, in place."""
|
||||||
|
if static_local_kv_lens_cpu is not None:
|
||||||
|
# The planner reads `kv_len_arr_cpu` only for (max, sum), never for
|
||||||
|
# the lengths it writes back, so a static upper bound yields the same
|
||||||
|
# metadata without the GPU->CPU sync.
|
||||||
|
total_local_len = plan_dcp_decode_metadata(
|
||||||
|
kv_lens_gpu,
|
||||||
|
kv_indptr,
|
||||||
|
kv_indices,
|
||||||
|
init_metadata_replay=True,
|
||||||
|
fast_decode_kwargs={"kv_len_arr_cpu": static_local_kv_lens_cpu},
|
||||||
|
bs=bs,
|
||||||
|
)
|
||||||
|
elif seq_lens_cpu is not None:
|
||||||
|
kv_len_arr_cpu = seq_lens_cpu[:bs].to(torch.int32).clone()
|
||||||
|
update_local_kv_lens_for_dcp(kv_len_arr_cpu)
|
||||||
|
total_local_len = plan_dcp_decode_metadata(
|
||||||
|
kv_lens_gpu,
|
||||||
|
kv_indptr,
|
||||||
|
kv_indices,
|
||||||
|
init_metadata_replay=True,
|
||||||
|
fast_decode_kwargs={"kv_len_arr_cpu": kv_len_arr_cpu},
|
||||||
|
bs=bs,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
total_local_len = plan_dcp_decode_metadata(
|
||||||
|
kv_lens_gpu,
|
||||||
|
kv_indptr,
|
||||||
|
kv_indices,
|
||||||
|
init_metadata_replay=False,
|
||||||
|
fast_decode_kwargs={},
|
||||||
|
bs=bs,
|
||||||
|
)
|
||||||
|
|
||||||
|
# The planner leaves the compacted ids WIDENED (see its docstring), and
|
||||||
|
# mla_gluon indexes the pool directly, so collapse them once per forward
|
||||||
|
# -- the same contract flashinfer_mla_backend.py follows.
|
||||||
|
translator = self.kv_index_translator
|
||||||
|
if total_local_len > 0 and translator.needs_read_translate:
|
||||||
|
valid = kv_indices[:total_local_len]
|
||||||
|
valid.copy_(translator.translate_dcp_read_ids(valid))
|
||||||
|
|
||||||
|
def _build_dcp_local_kv_lens(
|
||||||
|
self,
|
||||||
|
kv_indptr: torch.Tensor,
|
||||||
|
bs: int,
|
||||||
|
out_lens: Optional[torch.Tensor] = None,
|
||||||
|
):
|
||||||
|
"""This rank's shard length per request, in TOKENS (mla_gluon's
|
||||||
|
``cache_seqlens``). ``kv_indptr`` must already be localized.
|
||||||
|
"""
|
||||||
|
lens = (kv_indptr[1 : bs + 1] - kv_indptr[:bs]).to(torch.int32)
|
||||||
|
if out_lens is None:
|
||||||
|
return lens
|
||||||
|
out_lens.copy_(lens)
|
||||||
|
return out_lens
|
||||||
|
|
||||||
|
def _build_dcp_verify_token_table(
|
||||||
|
self,
|
||||||
|
kv_indptr: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
bs: int,
|
||||||
|
q_len: int,
|
||||||
|
max_local_kv_len: int,
|
||||||
|
out: Optional[torch.Tensor] = None,
|
||||||
|
out_lens: Optional[torch.Tensor] = None,
|
||||||
|
):
|
||||||
|
"""Token table + shard lengths for the prefix attention of DCP verify.
|
||||||
|
|
||||||
|
One row per query token, one column per TOKEN (mla_gluon fixes
|
||||||
|
PAGE_SIZE at 1). Rows of a request repeat that request's shard.
|
||||||
|
"""
|
||||||
|
local_kv_lens = self._build_dcp_local_kv_lens(kv_indptr, bs)
|
||||||
|
n_rows = bs * q_len
|
||||||
|
if out is None:
|
||||||
|
# The row stride below is a Triton constexpr, so quantize the eager
|
||||||
|
# width: every distinct value costs a JIT specialization.
|
||||||
|
out = local_kv_lens.new_empty(
|
||||||
|
(
|
||||||
|
n_rows,
|
||||||
|
triton.cdiv(max_local_kv_len, _DCP_VERIFY_TABLE_COLS_PER_BLOCK)
|
||||||
|
* _DCP_VERIFY_TABLE_COLS_PER_BLOCK,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
num_cols = out.shape[1]
|
||||||
|
|
||||||
|
# Write each request's row 0 in place: the row stride handed to the
|
||||||
|
# kernel spans that request's whole block of q_len rows.
|
||||||
|
translator = self.kv_index_translator
|
||||||
|
v2p = translator.full_v2p_table
|
||||||
|
create_mla_kv_page_table_for_dcp[
|
||||||
|
(bs, triton.cdiv(num_cols, _DCP_VERIFY_TABLE_COLS_PER_BLOCK))
|
||||||
|
](
|
||||||
|
self.req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
local_kv_lens,
|
||||||
|
out,
|
||||||
|
v2p,
|
||||||
|
self.req_to_token.stride(0),
|
||||||
|
q_len * num_cols,
|
||||||
|
translator.full_page_multiplier,
|
||||||
|
PHYSICAL_PAGE_SIZE=1,
|
||||||
|
DCP_SIZE=self.dcp_world_size,
|
||||||
|
DCP_RANK=get_parallel().attn_dcp_rank,
|
||||||
|
PAGES_PER_BLOCK=_DCP_VERIFY_TABLE_COLS_PER_BLOCK,
|
||||||
|
HAS_V2P=v2p is not None,
|
||||||
|
)
|
||||||
|
rows = out.view(bs, q_len, num_cols)
|
||||||
|
if q_len > 1:
|
||||||
|
# Source is row 0, destination rows 1.., so the copy never overlaps.
|
||||||
|
rows[:, 1:, :].copy_(rows[:, :1, :].expand(bs, q_len - 1, num_cols))
|
||||||
|
|
||||||
|
if out_lens is None:
|
||||||
|
out_lens = local_kv_lens.new_empty((n_rows,))
|
||||||
|
out_lens.view(bs, q_len).copy_(local_kv_lens.unsqueeze(1).expand(bs, q_len))
|
||||||
|
return out, out_lens
|
||||||
|
|
||||||
def init_cuda_graph_state(
|
def init_cuda_graph_state(
|
||||||
self,
|
self,
|
||||||
max_bs: int,
|
max_bs: int,
|
||||||
@@ -1739,6 +2005,29 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
self.cuda_graph_kv_last_page_len = torch.ones(
|
self.cuda_graph_kv_last_page_len = torch.ones(
|
||||||
max_bs, dtype=torch.int32, device=self.device
|
max_bs, dtype=torch.int32, device=self.device
|
||||||
)
|
)
|
||||||
|
if self.use_mla and self.dcp_world_size > 1:
|
||||||
|
if self.num_draft_tokens:
|
||||||
|
# Target-verify flattens the window into single-token rows, so it
|
||||||
|
# needs max_bs * num_draft_tokens.
|
||||||
|
n_verify_rows = max_bs * self.num_draft_tokens
|
||||||
|
self.cuda_graph_verify_local_kv_lens = torch.zeros(
|
||||||
|
(n_verify_rows,), dtype=torch.int32, device=self.device
|
||||||
|
)
|
||||||
|
# Capture-stable token table, filled out-of-graph. One column per
|
||||||
|
# TOKEN, so the width is the worst-case shard ceil(ctx_len / W).
|
||||||
|
self.cuda_graph_verify_token_table = torch.zeros(
|
||||||
|
(n_verify_rows, self._get_dcp_graph_max_local_kv_len()),
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
# Static per-rank shard bound, one entry per request. Sizes the
|
||||||
|
# verify plan without a sync; see _plan_dcp_decode_metadata.
|
||||||
|
self.cuda_graph_dcp_static_local_kv_lens = torch.full(
|
||||||
|
(max_bs,),
|
||||||
|
self._get_dcp_graph_max_local_kv_len(),
|
||||||
|
dtype=torch.int32,
|
||||||
|
device="cpu",
|
||||||
|
)
|
||||||
if kv_indices_buf is None:
|
if kv_indices_buf is None:
|
||||||
max_num_blocks_per_seq = (
|
max_num_blocks_per_seq = (
|
||||||
self.max_context_len + self.page_size - 1
|
self.max_context_len + self.page_size - 1
|
||||||
@@ -1854,6 +2143,10 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
reduce_final_map = None
|
reduce_final_map = None
|
||||||
reduce_partial_map = None
|
reduce_partial_map = None
|
||||||
|
|
||||||
|
# DCP metadata which will be populated for MLA decode when dcp enabled
|
||||||
|
local_kv_lens = None
|
||||||
|
verify_token_table = None
|
||||||
|
|
||||||
swa_page_table = None
|
swa_page_table = None
|
||||||
max_kv_len = (
|
max_kv_len = (
|
||||||
seq_lens_cpu.max().item()
|
seq_lens_cpu.max().item()
|
||||||
@@ -1887,6 +2180,20 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_indices,
|
kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if (
|
||||||
|
self.use_mla
|
||||||
|
and self.dcp_world_size > 1
|
||||||
|
and not forward_mode.is_idle()
|
||||||
|
):
|
||||||
|
kv_lens = seq_lens[:bs].to(torch.int32).clone()
|
||||||
|
self._plan_dcp_decode_metadata(
|
||||||
|
kv_indptr,
|
||||||
|
kv_indices,
|
||||||
|
kv_lens,
|
||||||
|
seq_lens_cpu,
|
||||||
|
bs,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
max_q_len = 1
|
max_q_len = 1
|
||||||
kv_indices = self.cuda_graph_page_table
|
kv_indices = self.cuda_graph_page_table
|
||||||
@@ -1946,7 +2253,9 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
||||||
max_q_len = 1
|
max_q_len = 1
|
||||||
|
|
||||||
if _use_mla_ps_kernel:
|
# DCP decode builds its own block-table metadata in
|
||||||
|
# forward_decode, so the persist metadata is unused here.
|
||||||
|
if _use_mla_ps_kernel and self.dcp_world_size <= 1:
|
||||||
num_kv_splits = self.max_split_per_batch
|
num_kv_splits = self.max_split_per_batch
|
||||||
|
|
||||||
self.make_mla_meta_data(
|
self.make_mla_meta_data(
|
||||||
@@ -2008,7 +2317,11 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
kv_lens = seq_lens + self.num_draft_tokens
|
kv_lens = (
|
||||||
|
seq_lens
|
||||||
|
if self.dcp_world_size > 1
|
||||||
|
else seq_lens + self.num_draft_tokens
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
kv_lens = seq_lens
|
kv_lens = seq_lens
|
||||||
kv_indptr = self.kv_indptr[: bs + 1]
|
kv_indptr = self.kv_indptr[: bs + 1]
|
||||||
@@ -2039,9 +2352,34 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
||||||
|
|
||||||
|
if self.use_mla and self.dcp_world_size > 1:
|
||||||
|
self._plan_dcp_decode_metadata(
|
||||||
|
kv_indptr,
|
||||||
|
kv_indices,
|
||||||
|
seq_lens[:bs].to(torch.int32).clone(),
|
||||||
|
None,
|
||||||
|
bs,
|
||||||
|
static_local_kv_lens_cpu=self.cuda_graph_dcp_static_local_kv_lens[
|
||||||
|
:bs
|
||||||
|
],
|
||||||
|
)
|
||||||
|
n_rows = bs * self.num_draft_tokens
|
||||||
|
(
|
||||||
|
verify_token_table,
|
||||||
|
local_kv_lens,
|
||||||
|
) = self._build_dcp_verify_token_table(
|
||||||
|
kv_indptr,
|
||||||
|
req_pool_indices,
|
||||||
|
bs,
|
||||||
|
self.num_draft_tokens,
|
||||||
|
self._get_dcp_graph_max_local_kv_len(),
|
||||||
|
out=self.cuda_graph_verify_token_table[:n_rows],
|
||||||
|
out_lens=self.cuda_graph_verify_local_kv_lens[:n_rows],
|
||||||
|
)
|
||||||
|
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
max_q_len = self.num_draft_tokens
|
max_q_len = self.num_draft_tokens
|
||||||
if _use_mla_ps_kernel:
|
if _use_mla_ps_kernel and self.dcp_world_size <= 1:
|
||||||
num_kv_splits = self.max_split_per_batch
|
num_kv_splits = self.max_split_per_batch
|
||||||
|
|
||||||
self.make_mla_meta_data(
|
self.make_mla_meta_data(
|
||||||
@@ -2082,6 +2420,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
reduce_final_map=reduce_final_map,
|
reduce_final_map=reduce_final_map,
|
||||||
reduce_partial_map=reduce_partial_map,
|
reduce_partial_map=reduce_partial_map,
|
||||||
num_kv_splits=num_kv_splits,
|
num_kv_splits=num_kv_splits,
|
||||||
|
local_kv_lens=local_kv_lens,
|
||||||
|
verify_token_table=verify_token_table,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
max_q_len = verify_tokens_per_req
|
max_q_len = verify_tokens_per_req
|
||||||
@@ -2419,6 +2759,15 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
v_scale=v_descale,
|
v_scale=v_descale,
|
||||||
)
|
)
|
||||||
elif self.use_mla:
|
elif self.use_mla:
|
||||||
|
if self.dcp_world_size > 1:
|
||||||
|
kv_lora_rank = v.shape[-1]
|
||||||
|
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||||
|
layer,
|
||||||
|
cache_loc,
|
||||||
|
k[..., :kv_lora_rank],
|
||||||
|
k[..., kv_lora_rank:],
|
||||||
|
)
|
||||||
|
else:
|
||||||
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
|
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
|
||||||
elif self._use_fused_fp8_kv_write(layer):
|
elif self._use_fused_fp8_kv_write(layer):
|
||||||
# FP8: fuse bf16->fp8 cast + paged write in one kernel.
|
# FP8: fuse bf16->fp8 cast + paged write in one kernel.
|
||||||
@@ -2458,6 +2807,15 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
V_Buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
V_Buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||||
kv_lora_rank = V_Buffer.shape[-1]
|
kv_lora_rank = V_Buffer.shape[-1]
|
||||||
qk_rope_head_dim = K_Buffer.shape[-1] - kv_lora_rank
|
qk_rope_head_dim = K_Buffer.shape[-1] - kv_lora_rank
|
||||||
|
|
||||||
|
if (
|
||||||
|
forward_batch.forward_mode.is_target_verify()
|
||||||
|
and self.dcp_world_size > 1
|
||||||
|
):
|
||||||
|
# two-stage dcp verify, dispatched before the dims below: the
|
||||||
|
# model provides the per rank kvcache slices.
|
||||||
|
return self._forward_verify_dcp(q, k, layer, k_descale)
|
||||||
|
|
||||||
qk_nope_head_dim = k.shape[-1] - qk_rope_head_dim
|
qk_nope_head_dim = k.shape[-1] - qk_rope_head_dim
|
||||||
assert len(q.shape) == 3
|
assert len(q.shape) == 3
|
||||||
assert len(k.shape) == 3
|
assert len(k.shape) == 3
|
||||||
@@ -2471,6 +2829,20 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
|
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
|
||||||
if forward_batch.mha_return_lse:
|
if forward_batch.mha_return_lse:
|
||||||
return self._forward_extend_skip_prefix(q, k, v, layer)
|
return self._forward_extend_skip_prefix(q, k, v, layer)
|
||||||
|
if self.dcp_world_size > 1:
|
||||||
|
if self.use_fp8_prefill_attn and self.head_pad_mode != "zero":
|
||||||
|
return self.mla_fp8_prefill_attn(q, k, v, layer)
|
||||||
|
return flash_attn_varlen_func(
|
||||||
|
q,
|
||||||
|
k,
|
||||||
|
v,
|
||||||
|
qo_indptr,
|
||||||
|
forward_batch.attn_dcp_metadata.dcp_kv_indptr,
|
||||||
|
max_q_len,
|
||||||
|
max_kv_len,
|
||||||
|
softmax_scale=layer.scaling,
|
||||||
|
causal=True,
|
||||||
|
)
|
||||||
if kv_indices.shape[0] == 0 or extend_no_prefix:
|
if kv_indices.shape[0] == 0 or extend_no_prefix:
|
||||||
if self.use_fp8_prefill_attn and self.head_pad_mode != "zero":
|
if self.use_fp8_prefill_attn and self.head_pad_mode != "zero":
|
||||||
output = self.mla_fp8_prefill_attn(
|
output = self.mla_fp8_prefill_attn(
|
||||||
@@ -3253,6 +3625,10 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
|
if self.dcp_world_size > 1 and not forward_batch.forward_mode.is_idle():
|
||||||
|
k_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||||
|
return self._forward_decode_dcp(q, k_buffer, layer, k_descale)
|
||||||
|
|
||||||
o = self._forward_mla_decode(q, layer, forward_batch, k_descale)
|
o = self._forward_mla_decode(q, layer, forward_batch, k_descale)
|
||||||
return o.reshape(-1, layer.tp_q_head_num * layer.v_head_dim)
|
return o.reshape(-1, layer.tp_q_head_num * layer.v_head_dim)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ from __future__ import annotations
|
|||||||
import functools
|
import functools
|
||||||
import inspect
|
import inspect
|
||||||
import logging
|
import logging
|
||||||
from typing import TYPE_CHECKING, Optional
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -82,9 +82,12 @@ def mla_gluon_decode(
|
|||||||
min_kv_seq_len: int,
|
min_kv_seq_len: int,
|
||||||
kv_scale: float = 1.0,
|
kv_scale: float = 1.0,
|
||||||
qlen: int = 1,
|
qlen: int = 1,
|
||||||
) -> Optional[torch.Tensor]:
|
use_2d_view: bool = False,
|
||||||
|
return_lse: bool = False,
|
||||||
|
):
|
||||||
"""Run Gluon MLA decode for fused Q [num_tokens, H, 576] and MLA KV pool.
|
"""Run Gluon MLA decode for fused Q [num_tokens, H, 576] and MLA KV pool.
|
||||||
Returns [num_tokens, H, v_head_dim], or None when Gluon is unavailable.
|
Returns [num_tokens, H, v_head_dim] (or ``(out, lse)`` when ``return_lse``),
|
||||||
|
or None when Gluon is unavailable.
|
||||||
"""
|
"""
|
||||||
mla_gluon = _gluon_fn()
|
mla_gluon = _gluon_fn()
|
||||||
if mla_gluon is None:
|
if mla_gluon is None:
|
||||||
@@ -105,7 +108,11 @@ def mla_gluon_decode(
|
|||||||
else:
|
else:
|
||||||
o = q.new_empty((batch_size, num_head, kv_lora_rank))
|
o = q.new_empty((batch_size, num_head, kv_lora_rank))
|
||||||
|
|
||||||
mla_gluon(
|
extra_kwargs = {}
|
||||||
|
if return_lse:
|
||||||
|
extra_kwargs["return_lse"] = True
|
||||||
|
|
||||||
|
result = mla_gluon(
|
||||||
q_nope,
|
q_nope,
|
||||||
q_pe,
|
q_pe,
|
||||||
k_buffer.view(-1, layer.qk_head_dim),
|
k_buffer.view(-1, layer.qk_head_dim),
|
||||||
@@ -115,9 +122,15 @@ def mla_gluon_decode(
|
|||||||
sm_scale,
|
sm_scale,
|
||||||
k_pe=None,
|
k_pe=None,
|
||||||
kv_pe_offset=kv_lora_rank,
|
kv_pe_offset=kv_lora_rank,
|
||||||
use_2d_view=False,
|
use_2d_view=use_2d_view,
|
||||||
kv_scale=kv_scale,
|
kv_scale=kv_scale,
|
||||||
min_kv_seq_len=min_kv_seq_len,
|
min_kv_seq_len=min_kv_seq_len,
|
||||||
|
**extra_kwargs,
|
||||||
)
|
)
|
||||||
# Hand back the caller's flat [num_tokens, H, v] layout either way.
|
# Hand back the caller's flat [num_tokens, H, v] layout either way.
|
||||||
return o.flatten(0, 1) if qlen > 1 else o
|
out = o.flatten(0, 1) if qlen > 1 else o
|
||||||
|
if not return_lse:
|
||||||
|
return out
|
||||||
|
# mla_gluon writes the output into `o` and returns it alongside the lse.
|
||||||
|
_, lse = result
|
||||||
|
return out, lse
|
||||||
|
|||||||
@@ -52,11 +52,13 @@ from sglang.srt.utils import (
|
|||||||
is_cuda,
|
is_cuda,
|
||||||
is_gfx95_supported,
|
is_gfx95_supported,
|
||||||
is_gfx942_supported,
|
is_gfx942_supported,
|
||||||
|
is_hip,
|
||||||
is_xpu,
|
is_xpu,
|
||||||
next_power_of_2,
|
next_power_of_2,
|
||||||
)
|
)
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
|
_is_hip = is_hip()
|
||||||
_is_gfx942 = is_gfx942_supported()
|
_is_gfx942 = is_gfx942_supported()
|
||||||
_is_xpu = is_xpu()
|
_is_xpu = is_xpu()
|
||||||
|
|
||||||
@@ -229,6 +231,17 @@ class TritonAttnBackend(AttentionBackend):
|
|||||||
self.use_mla,
|
self.use_mla,
|
||||||
self.use_verify_splitkv,
|
self.use_verify_splitkv,
|
||||||
)
|
)
|
||||||
|
# TODO: this logic should be fixed in non-hip platform
|
||||||
|
self.is_hip_dspark_draft = (
|
||||||
|
_is_hip
|
||||||
|
and model_runner.is_draft_worker
|
||||||
|
and model_runner.spec_algorithm.is_dspark()
|
||||||
|
)
|
||||||
|
if self.is_hip_dspark_draft:
|
||||||
|
# Drafts never join the dcp group so we ignore it
|
||||||
|
self.dcp_size = 1
|
||||||
|
self.dcp_rank = 0
|
||||||
|
else:
|
||||||
self.dcp_size = get_parallel().attn_dcp_size
|
self.dcp_size = get_parallel().attn_dcp_size
|
||||||
self.dcp_rank = get_parallel().attn_dcp_rank
|
self.dcp_rank = get_parallel().attn_dcp_rank
|
||||||
self.num_head = (
|
self.num_head = (
|
||||||
|
|||||||
@@ -37,8 +37,11 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||||
from sglang.srt.runtime_context import get_parallel, get_platform
|
from sglang.srt.runtime_context import get_parallel, get_platform
|
||||||
|
from sglang.srt.utils import is_hip
|
||||||
from sglang.srt.utils.common import is_mnnvl_fabric_device
|
from sglang.srt.utils.common import is_mnnvl_fabric_device
|
||||||
|
|
||||||
|
_is_hip = is_hip()
|
||||||
|
|
||||||
|
|
||||||
def _warn_deprecated_dcp_accessor(name: str, replacement: str) -> None:
|
def _warn_deprecated_dcp_accessor(name: str, replacement: str) -> None:
|
||||||
warnings.warn(
|
warnings.warn(
|
||||||
@@ -276,6 +279,8 @@ def all_gather_kv_cache_for_mla_extend(
|
|||||||
k_nope,
|
k_nope,
|
||||||
k_pe,
|
k_pe,
|
||||||
):
|
):
|
||||||
|
# On hip, skip the all-gather when there is no cached prefix to avoid crash
|
||||||
|
if not _is_hip or dcp_extend_prefix_lens_sum > 0:
|
||||||
cache_k_nope, cache_k_rope = token_to_kv_pool.get_mla_kv_buffer(
|
cache_k_nope, cache_k_rope = token_to_kv_pool.get_mla_kv_buffer(
|
||||||
attn_mqa,
|
attn_mqa,
|
||||||
dcp_local_prefix_kv_indices,
|
dcp_local_prefix_kv_indices,
|
||||||
|
|||||||
@@ -196,6 +196,8 @@ def handle_attention_aiter(attn, forward_batch):
|
|||||||
if forward_batch.forward_mode.is_extend_without_speculative():
|
if forward_batch.forward_mode.is_extend_without_speculative():
|
||||||
if not _support_mha_one_shot(attn, forward_batch, "aiter"):
|
if not _support_mha_one_shot(attn, forward_batch, "aiter"):
|
||||||
return AttnForwardMethod.MHA_CHUNKED_KV
|
return AttnForwardMethod.MHA_CHUNKED_KV
|
||||||
|
if get_parallel().dcp_enabled:
|
||||||
|
return AttnForwardMethod.MHA_ONE_SHOT
|
||||||
return AttnForwardMethod.MHA
|
return AttnForwardMethod.MHA
|
||||||
else:
|
else:
|
||||||
return AttnForwardMethod.MLA
|
return AttnForwardMethod.MLA
|
||||||
|
|||||||
+1
-10
@@ -14,16 +14,12 @@ import torch
|
|||||||
|
|
||||||
from sglang.kernels.ops.attention.utils import concat_and_cast_mha_k_triton
|
from sglang.kernels.ops.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.dcp import (
|
from sglang.srt.layers.dcp import all_gather_kv_cache_for_mha_extend
|
||||||
all_gather_kv_cache_for_mha_extend,
|
|
||||||
filter_dcp_local_kv_indices,
|
|
||||||
)
|
|
||||||
from sglang.srt.layers.quantization.fp8_utils import (
|
from sglang.srt.layers.quantization.fp8_utils import (
|
||||||
materialize_bpreshuffle_fp8_scale_tuple,
|
materialize_bpreshuffle_fp8_scale_tuple,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_executor.forward_context import (
|
from sglang.srt.model_executor.forward_context import (
|
||||||
get_attn_backend,
|
|
||||||
get_token_to_kv_pool,
|
get_token_to_kv_pool,
|
||||||
)
|
)
|
||||||
from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import (
|
from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import (
|
||||||
@@ -319,11 +315,6 @@ class DeepseekMHARocmForwardMixin:
|
|||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
):
|
):
|
||||||
if _use_aiter_gfx95:
|
if _use_aiter_gfx95:
|
||||||
kv_indices = filter_dcp_local_kv_indices(kv_indices=kv_indices)
|
|
||||||
# Read door: the pool never translates, so the production site does.
|
|
||||||
kv_indices = get_attn_backend().kv_index_translator.translate_dcp_read_ids(
|
|
||||||
kv_indices
|
|
||||||
)
|
|
||||||
kv_a, k_pe = get_token_to_kv_pool().get_mla_kv_buffer(
|
kv_a, k_pe = get_token_to_kv_pool().get_mla_kv_buffer(
|
||||||
self.attn_mha, kv_indices, dst_dtype
|
self.attn_mha, kv_indices, dst_dtype
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -99,7 +99,7 @@ def is_dcp_mla_decode_phase(forward_batch: ForwardBatch) -> bool:
|
|||||||
|
|
||||||
|
|
||||||
def is_mla_dcp_lse_base_on_e(attention_backend: Optional[str]) -> bool:
|
def is_mla_dcp_lse_base_on_e(attention_backend: Optional[str]) -> bool:
|
||||||
return attention_backend in {"flashmla", "cutedsl_mla"}
|
return attention_backend in {"flashmla", "cutedsl_mla", "aiter"}
|
||||||
|
|
||||||
|
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
|
|||||||
+55
-8
@@ -591,25 +591,35 @@ class DeepseekMLARocmForwardMixin:
|
|||||||
q_nope_out = apply_kv_b_lora_q_correction(self, q_nope, q_nope_out)
|
q_nope_out = apply_kv_b_lora_q_correction(self, q_nope, q_nope_out)
|
||||||
|
|
||||||
fuse_rope_for_trtllm_mla = self._fuse_rope_for_trtllm_mla(forward_batch)
|
fuse_rope_for_trtllm_mla = self._fuse_rope_for_trtllm_mla(forward_batch)
|
||||||
if (
|
|
||||||
self.rotary_emb is not None
|
force_rope_for_aiter_dcp_decode = (
|
||||||
and (not fuse_rope_for_trtllm_mla)
|
get_parallel().dcp_enabled
|
||||||
|
and (
|
||||||
|
forward_batch.forward_mode.is_decode()
|
||||||
|
or forward_batch.forward_mode.is_target_verify()
|
||||||
|
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
|
)
|
||||||
|
and _use_aiter_gfx95
|
||||||
|
and self.current_attention_backend
|
||||||
|
not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS
|
||||||
|
)
|
||||||
|
if self.rotary_emb is not None and (
|
||||||
|
force_rope_for_aiter_dcp_decode
|
||||||
|
or (
|
||||||
|
(not fuse_rope_for_trtllm_mla)
|
||||||
and (not self._skip_rope_for_dsa_tilelang_fused())
|
and (not self._skip_rope_for_dsa_tilelang_fused())
|
||||||
and (not self._skip_rope_for_aiter_fused_mla())
|
and (not self._skip_rope_for_aiter_fused_mla())
|
||||||
and (
|
and (
|
||||||
not _use_aiter
|
not _use_aiter
|
||||||
or not _is_gfx95_supported
|
or not _is_gfx95_supported
|
||||||
or self.use_dsa
|
or self.use_dsa
|
||||||
# Non-fused, non-specialized attention backends (e.g. Triton) run
|
|
||||||
# the cat path in forward_absorb_core and need RoPE applied here;
|
|
||||||
# only the aiter fused MLA path and the specialized MLA backends
|
|
||||||
# defer RoPE to their own kernels.
|
|
||||||
or (
|
or (
|
||||||
self.current_attention_backend
|
self.current_attention_backend
|
||||||
not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS
|
not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS
|
||||||
and self.current_attention_backend != "aiter"
|
and self.current_attention_backend != "aiter"
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
)
|
||||||
):
|
):
|
||||||
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
|
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
|
||||||
|
|
||||||
@@ -622,7 +632,12 @@ class DeepseekMLARocmForwardMixin:
|
|||||||
q_pe=q_pe,
|
q_pe=q_pe,
|
||||||
)
|
)
|
||||||
elif forward_batch.forward_mode.is_extend():
|
elif forward_batch.forward_mode.is_extend():
|
||||||
# for extend, gather kv
|
# Assemble the full sequence into dcp_kv_buffer, which the
|
||||||
|
# backend attends over instead of the sharded local cache.
|
||||||
|
if (
|
||||||
|
forward_batch.attn_dcp_metadata is not None
|
||||||
|
and forward_batch.attn_dcp_metadata.dcp_kv_buffer is not None
|
||||||
|
):
|
||||||
all_gather_kv_cache_for_mla_extend(
|
all_gather_kv_cache_for_mla_extend(
|
||||||
get_token_to_kv_pool(),
|
get_token_to_kv_pool(),
|
||||||
self.attn_mqa,
|
self.attn_mqa,
|
||||||
@@ -763,6 +778,38 @@ class DeepseekMLARocmForwardMixin:
|
|||||||
else {}
|
else {}
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
elif (
|
||||||
|
_use_aiter
|
||||||
|
and (
|
||||||
|
forward_batch.forward_mode.is_decode()
|
||||||
|
or forward_batch.forward_mode.is_target_verify()
|
||||||
|
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
|
)
|
||||||
|
and get_parallel().dcp_enabled
|
||||||
|
):
|
||||||
|
q = torch.cat([q_nope_out, q_pe], dim=-1)
|
||||||
|
if llama_4_scaling is not None:
|
||||||
|
q[..., : self.kv_lora_rank] *= llama_4_scaling
|
||||||
|
get_token_to_kv_pool().set_mla_kv_buffer(
|
||||||
|
self.attn_mqa,
|
||||||
|
forward_batch.out_cache_loc,
|
||||||
|
k_nope,
|
||||||
|
k_pe,
|
||||||
|
)
|
||||||
|
if forward_batch.forward_mode.is_target_verify():
|
||||||
|
k_window = torch.cat([k_nope, k_pe], dim=-1)
|
||||||
|
v_window = k_nope
|
||||||
|
else:
|
||||||
|
k_window = None
|
||||||
|
v_window = None
|
||||||
|
attn_output, lse = self.attn_mqa_for_dcp_decode(
|
||||||
|
q,
|
||||||
|
k_window,
|
||||||
|
v_window,
|
||||||
|
forward_batch,
|
||||||
|
save_kv_cache=False,
|
||||||
|
**(dict(topk_indices=topk_indices) if topk_indices is not None else {}),
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
if self._skip_rope_for_aiter_fused_mla():
|
if self._skip_rope_for_aiter_fused_mla():
|
||||||
q, _, _, k = _fused_rope_cat_and_cache(
|
q, _, _, k = _fused_rope_cat_and_cache(
|
||||||
|
|||||||
Reference in New Issue
Block a user