[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")
|
||||
def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
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",
|
||||
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:
|
||||
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:
|
||||
|
||||
@@ -26,6 +26,8 @@ from sglang.kernels.ops.kvcache.aiter_unified_attention import (
|
||||
scatter_req_to_token_to_page_table_kernel,
|
||||
)
|
||||
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.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
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."
|
||||
)
|
||||
|
||||
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 (
|
||||
launch_reshape_and_cache_flash,
|
||||
pad_sequence_with_mask,
|
||||
@@ -151,11 +155,16 @@ class ForwardMetadata:
|
||||
swa_page_table: Optional[torch.Tensor] = None
|
||||
# full->SWA translated out_cache_loc (SWA KV-store write target)
|
||||
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
|
||||
|
||||
|
||||
_DCP_VERIFY_TABLE_COLS_PER_BLOCK = 128
|
||||
|
||||
|
||||
# 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.
|
||||
_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.dcp_world_size = get_parallel().attn_dcp_size
|
||||
|
||||
# Get v_head_dim based on model type
|
||||
if self.use_mla:
|
||||
# For MLA models, get v_head_dim from model config
|
||||
@@ -348,6 +359,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
# corresponding ForwardBatch fields.
|
||||
self.req_to_token_pool = model_runner.req_to_token_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
|
||||
# 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,
|
||||
# by repetition when it divides 16 and by tiling otherwise.
|
||||
_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"in [16, 128].\n"
|
||||
f"Provided {self.num_head} number of heads.\n"
|
||||
"Try adjusting tensor_parallel_size value, or run decode on "
|
||||
"another backend (--decode-attention-backend)."
|
||||
)
|
||||
|
||||
self.num_head_padded = 16 if self.num_head < 16 else self.num_head
|
||||
if self.num_head in _mla_low_head_repeat:
|
||||
self.head_pad_mode = "repeat"
|
||||
@@ -451,17 +469,22 @@ class AiterAttnBackend(AttentionBackend):
|
||||
self.head_pad_mode = "none"
|
||||
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.qo_indptr_ = torch.zeros(
|
||||
(max_bs + 1,), dtype=torch.int32, device=model_runner.device
|
||||
)
|
||||
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
|
||||
# it should not use fake-nps (fast_mode = False, intra_batch_mode = True)
|
||||
# it will cause gpu-fault or accuracy issue
|
||||
if self.num_head in (32, 64, 128):
|
||||
# it will cause gpu-fault or accuracy issue.
|
||||
if self.mla_kernel_num_head_padded in (32, 64, 128):
|
||||
fast_mode = True
|
||||
intra_batch_mode = False
|
||||
|
||||
@@ -473,8 +496,8 @@ class AiterAttnBackend(AttentionBackend):
|
||||
# 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.
|
||||
if (
|
||||
(self.num_head_padded == 16 and self.head_pad_mode != "zero")
|
||||
or self.num_head_padded == 128
|
||||
(self.mla_kernel_num_head_padded == 16 and self.head_pad_mode != "zero")
|
||||
or self.mla_kernel_num_head_padded == 128
|
||||
) and self.kv_cache_dtype is not fp8_dtype:
|
||||
_use_mla_ps_kernel = False
|
||||
fast_mode = False
|
||||
@@ -551,7 +574,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
return "fp8_e4m3"
|
||||
|
||||
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
|
||||
|
||||
if self.enable_dp_attention:
|
||||
@@ -637,7 +660,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
qo_indptr,
|
||||
kv_indptr,
|
||||
kv_last_page_len,
|
||||
self.num_head_padded // nhead_kv,
|
||||
self.mla_kernel_num_head_padded // nhead_kv,
|
||||
nhead_kv,
|
||||
False,
|
||||
work_metadata,
|
||||
@@ -1084,6 +1107,80 @@ class AiterAttnBackend(AttentionBackend):
|
||||
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(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
@@ -1243,6 +1340,9 @@ class AiterAttnBackend(AttentionBackend):
|
||||
)
|
||||
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 spec_info is None or forward_batch.forward_mode.is_idle():
|
||||
kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
|
||||
@@ -1261,6 +1361,20 @@ class AiterAttnBackend(AttentionBackend):
|
||||
kv_indices,
|
||||
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:
|
||||
max_q_len = 1
|
||||
page_size = self.page_size
|
||||
@@ -1317,7 +1431,9 @@ class AiterAttnBackend(AttentionBackend):
|
||||
kv_last_page_len = self.kv_last_page_len[:bs]
|
||||
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_indptr,
|
||||
@@ -1456,9 +1572,13 @@ class AiterAttnBackend(AttentionBackend):
|
||||
elif forward_batch.forward_mode.is_target_verify():
|
||||
if self.use_mla:
|
||||
draft_num = spec_info.draft_token_num
|
||||
kv_lens = forward_batch.seq_lens + draft_num
|
||||
kv_lens_sum = forward_batch.seq_lens_sum + draft_num * bs
|
||||
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_sum = forward_batch.seq_lens_sum + draft_num * bs
|
||||
|
||||
qo_indptr = self.qo_indptr[: bs + 1]
|
||||
qo_indptr[: bs + 1] = torch.arange(
|
||||
@@ -1486,8 +1606,26 @@ class AiterAttnBackend(AttentionBackend):
|
||||
TOKEN_BLOCK_PARALLEL=num_token_blocks > 1,
|
||||
)
|
||||
|
||||
# if self.kv_cache_dtype == fp8_dtype:
|
||||
if _use_mla_ps_kernel:
|
||||
if self.dcp_world_size > 1:
|
||||
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
|
||||
(
|
||||
work_metadata,
|
||||
@@ -1532,6 +1670,8 @@ class AiterAttnBackend(AttentionBackend):
|
||||
reduce_partial_map=reduce_partial_map,
|
||||
num_kv_splits=num_kv_splits,
|
||||
run_graph=False,
|
||||
local_kv_lens=local_kv_lens,
|
||||
verify_token_table=verify_token_table,
|
||||
)
|
||||
else:
|
||||
draft_num = forward_batch.input_ids.shape[0] // bs
|
||||
@@ -1712,6 +1852,132 @@ class AiterAttnBackend(AttentionBackend):
|
||||
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(
|
||||
self,
|
||||
max_bs: int,
|
||||
@@ -1739,6 +2005,29 @@ class AiterAttnBackend(AttentionBackend):
|
||||
self.cuda_graph_kv_last_page_len = torch.ones(
|
||||
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:
|
||||
max_num_blocks_per_seq = (
|
||||
self.max_context_len + self.page_size - 1
|
||||
@@ -1854,6 +2143,10 @@ class AiterAttnBackend(AttentionBackend):
|
||||
reduce_final_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
|
||||
max_kv_len = (
|
||||
seq_lens_cpu.max().item()
|
||||
@@ -1887,6 +2180,20 @@ class AiterAttnBackend(AttentionBackend):
|
||||
kv_indices,
|
||||
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:
|
||||
max_q_len = 1
|
||||
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]
|
||||
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
|
||||
|
||||
self.make_mla_meta_data(
|
||||
@@ -2008,7 +2317,11 @@ class AiterAttnBackend(AttentionBackend):
|
||||
device=self.device,
|
||||
)
|
||||
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:
|
||||
kv_lens = seq_lens
|
||||
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]
|
||||
|
||||
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:
|
||||
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
|
||||
|
||||
self.make_mla_meta_data(
|
||||
@@ -2082,6 +2420,8 @@ class AiterAttnBackend(AttentionBackend):
|
||||
reduce_final_map=reduce_final_map,
|
||||
reduce_partial_map=reduce_partial_map,
|
||||
num_kv_splits=num_kv_splits,
|
||||
local_kv_lens=local_kv_lens,
|
||||
verify_token_table=verify_token_table,
|
||||
)
|
||||
else:
|
||||
max_q_len = verify_tokens_per_req
|
||||
@@ -2419,7 +2759,16 @@ class AiterAttnBackend(AttentionBackend):
|
||||
v_scale=v_descale,
|
||||
)
|
||||
elif self.use_mla:
|
||||
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
|
||||
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)
|
||||
elif self._use_fused_fp8_kv_write(layer):
|
||||
# FP8: fuse bf16->fp8 cast + paged write in one kernel.
|
||||
k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(
|
||||
@@ -2458,6 +2807,15 @@ class AiterAttnBackend(AttentionBackend):
|
||||
V_Buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
kv_lora_rank = V_Buffer.shape[-1]
|
||||
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
|
||||
assert len(q.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)
|
||||
if forward_batch.mha_return_lse:
|
||||
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 self.use_fp8_prefill_attn and self.head_pad_mode != "zero":
|
||||
output = self.mla_fp8_prefill_attn(
|
||||
@@ -3253,6 +3625,10 @@ class AiterAttnBackend(AttentionBackend):
|
||||
)
|
||||
|
||||
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)
|
||||
return o.reshape(-1, layer.tp_q_head_num * layer.v_head_dim)
|
||||
else:
|
||||
|
||||
@@ -14,7 +14,7 @@ from __future__ import annotations
|
||||
import functools
|
||||
import inspect
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
@@ -82,9 +82,12 @@ def mla_gluon_decode(
|
||||
min_kv_seq_len: int,
|
||||
kv_scale: float = 1.0,
|
||||
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.
|
||||
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()
|
||||
if mla_gluon is None:
|
||||
@@ -105,7 +108,11 @@ def mla_gluon_decode(
|
||||
else:
|
||||
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_pe,
|
||||
k_buffer.view(-1, layer.qk_head_dim),
|
||||
@@ -115,9 +122,15 @@ def mla_gluon_decode(
|
||||
sm_scale,
|
||||
k_pe=None,
|
||||
kv_pe_offset=kv_lora_rank,
|
||||
use_2d_view=False,
|
||||
use_2d_view=use_2d_view,
|
||||
kv_scale=kv_scale,
|
||||
min_kv_seq_len=min_kv_seq_len,
|
||||
**extra_kwargs,
|
||||
)
|
||||
# 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_gfx95_supported,
|
||||
is_gfx942_supported,
|
||||
is_hip,
|
||||
is_xpu,
|
||||
next_power_of_2,
|
||||
)
|
||||
|
||||
_is_cuda = is_cuda()
|
||||
_is_hip = is_hip()
|
||||
_is_gfx942 = is_gfx942_supported()
|
||||
_is_xpu = is_xpu()
|
||||
|
||||
@@ -229,8 +231,19 @@ class TritonAttnBackend(AttentionBackend):
|
||||
self.use_mla,
|
||||
self.use_verify_splitkv,
|
||||
)
|
||||
self.dcp_size = get_parallel().attn_dcp_size
|
||||
self.dcp_rank = get_parallel().attn_dcp_rank
|
||||
# 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_rank = get_parallel().attn_dcp_rank
|
||||
self.num_head = (
|
||||
model_runner.model_config.get_max_num_attention_heads()
|
||||
// get_parallel().attn_tp_size
|
||||
|
||||
@@ -37,8 +37,11 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
)
|
||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||
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
|
||||
|
||||
_is_hip = is_hip()
|
||||
|
||||
|
||||
def _warn_deprecated_dcp_accessor(name: str, replacement: str) -> None:
|
||||
warnings.warn(
|
||||
@@ -276,19 +279,21 @@ def all_gather_kv_cache_for_mla_extend(
|
||||
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
|
||||
# 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(
|
||||
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[
|
||||
|
||||
@@ -196,6 +196,8 @@ def handle_attention_aiter(attn, forward_batch):
|
||||
if forward_batch.forward_mode.is_extend_without_speculative():
|
||||
if not _support_mha_one_shot(attn, forward_batch, "aiter"):
|
||||
return AttnForwardMethod.MHA_CHUNKED_KV
|
||||
if get_parallel().dcp_enabled:
|
||||
return AttnForwardMethod.MHA_ONE_SHOT
|
||||
return AttnForwardMethod.MHA
|
||||
else:
|
||||
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.srt.layers.communicator import get_attn_tp_context
|
||||
from sglang.srt.layers.dcp import (
|
||||
all_gather_kv_cache_for_mha_extend,
|
||||
filter_dcp_local_kv_indices,
|
||||
)
|
||||
from sglang.srt.layers.dcp import all_gather_kv_cache_for_mha_extend
|
||||
from sglang.srt.layers.quantization.fp8_utils import (
|
||||
materialize_bpreshuffle_fp8_scale_tuple,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.model_executor.forward_context import (
|
||||
get_attn_backend,
|
||||
get_token_to_kv_pool,
|
||||
)
|
||||
from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import (
|
||||
@@ -319,11 +315,6 @@ class DeepseekMHARocmForwardMixin:
|
||||
forward_batch: ForwardBatch,
|
||||
):
|
||||
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(
|
||||
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:
|
||||
return attention_backend in {"flashmla", "cutedsl_mla"}
|
||||
return attention_backend in {"flashmla", "cutedsl_mla", "aiter"}
|
||||
|
||||
|
||||
if _is_cuda:
|
||||
|
||||
+75
-28
@@ -591,23 +591,33 @@ class DeepseekMLARocmForwardMixin:
|
||||
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)
|
||||
if (
|
||||
self.rotary_emb is not None
|
||||
and (not fuse_rope_for_trtllm_mla)
|
||||
and (not self._skip_rope_for_dsa_tilelang_fused())
|
||||
and (not self._skip_rope_for_aiter_fused_mla())
|
||||
|
||||
force_rope_for_aiter_dcp_decode = (
|
||||
get_parallel().dcp_enabled
|
||||
and (
|
||||
not _use_aiter
|
||||
or not _is_gfx95_supported
|
||||
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 (
|
||||
self.current_attention_backend
|
||||
not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS
|
||||
and self.current_attention_backend != "aiter"
|
||||
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_aiter_fused_mla())
|
||||
and (
|
||||
not _use_aiter
|
||||
or not _is_gfx95_supported
|
||||
or self.use_dsa
|
||||
or (
|
||||
self.current_attention_backend
|
||||
not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS
|
||||
and self.current_attention_backend != "aiter"
|
||||
)
|
||||
)
|
||||
)
|
||||
):
|
||||
@@ -622,18 +632,23 @@ class DeepseekMLARocmForwardMixin:
|
||||
q_pe=q_pe,
|
||||
)
|
||||
elif forward_batch.forward_mode.is_extend():
|
||||
# for extend, gather kv
|
||||
all_gather_kv_cache_for_mla_extend(
|
||||
get_token_to_kv_pool(),
|
||||
self.attn_mqa,
|
||||
forward_batch.extend_prefix_lens_cpu,
|
||||
forward_batch.attn_dcp_metadata.dcp_local_prefix_kv_indices,
|
||||
forward_batch.attn_dcp_metadata.dcp_extend_prefix_lens_sum,
|
||||
forward_batch.attn_dcp_metadata.dcp_kv_buffer,
|
||||
self.kv_lora_rank,
|
||||
k_nope,
|
||||
k_pe,
|
||||
)
|
||||
# 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(
|
||||
get_token_to_kv_pool(),
|
||||
self.attn_mqa,
|
||||
forward_batch.extend_prefix_lens_cpu,
|
||||
forward_batch.attn_dcp_metadata.dcp_local_prefix_kv_indices,
|
||||
forward_batch.attn_dcp_metadata.dcp_extend_prefix_lens_sum,
|
||||
forward_batch.attn_dcp_metadata.dcp_kv_buffer,
|
||||
self.kv_lora_rank,
|
||||
k_nope,
|
||||
k_pe,
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
f"not supported forward_mode {forward_batch.forward_mode}"
|
||||
@@ -763,6 +778,38 @@ class DeepseekMLARocmForwardMixin:
|
||||
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:
|
||||
if self._skip_rope_for_aiter_fused_mla():
|
||||
q, _, _, k = _fused_rope_cat_and_cache(
|
||||
|
||||
Reference in New Issue
Block a user