Reland "Support NextN = 2/4 in DSV32" (#27166)
Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
co-authored by
Brayden Zhong
parent
2c2a4f243a
commit
3b62286fca
@@ -602,11 +602,31 @@ class Indexer(MultiPlatformOp):
|
|||||||
# Reuse pre-computed schedule metadata if available (from init_forward_metadata),
|
# Reuse pre-computed schedule metadata if available (from init_forward_metadata),
|
||||||
# otherwise fall back to computing it here.
|
# otherwise fall back to computing it here.
|
||||||
schedule_metadata = getattr(metadata, "paged_mqa_schedule_metadata", None)
|
schedule_metadata = getattr(metadata, "paged_mqa_schedule_metadata", None)
|
||||||
# DeepGEMM release-0426 requires context_lens of shape [batch_size, next_n]
|
|
||||||
# to match q.shape = [batch_size, next_n, heads, head_dim]. The indexer uses
|
assert len(q_fp8.shape) == 3
|
||||||
# next_n=1 with batch_size=N_total via q_fp8.unsqueeze(1) below, so mirror
|
# attn_tp_size > 1 or MAX_LEN padding mode can leave padding in the
|
||||||
# that layout here.
|
# hidden states; q_offset is the real (unpadded) q length.
|
||||||
if seqlens_32.dim() == 2:
|
q_offset = sum(metadata.get_dsa_extend_len_cpu())
|
||||||
|
|
||||||
|
# DG-native q=[B,next_n,H,D] is faster than expanded q=[B*next_n,1,H,D]
|
||||||
|
# for target_verify with next_n>=2 (bigger MMA tile, fewer atoms). The
|
||||||
|
# precomputed ctx_lens_2d's shape is the single source of truth — if
|
||||||
|
# dsa_backend chose the per-token layout (e.g. non-SM100), fall through
|
||||||
|
# to the expanded path.
|
||||||
|
B = metadata.get_seqlens_int32().shape[0]
|
||||||
|
next_n = q_offset // B if B > 0 else 0
|
||||||
|
ctx_2d = getattr(metadata, "paged_mqa_ctx_lens_2d", None)
|
||||||
|
use_dg_native = (
|
||||||
|
_is_cuda
|
||||||
|
and forward_batch.forward_mode.is_target_verify()
|
||||||
|
and next_n >= 2
|
||||||
|
and ctx_2d is not None
|
||||||
|
and ctx_2d.shape == (B, next_n)
|
||||||
|
)
|
||||||
|
|
||||||
|
if use_dg_native:
|
||||||
|
seqlens_32_2d = ctx_2d
|
||||||
|
elif seqlens_32.dim() == 2:
|
||||||
seqlens_32_2d = seqlens_32
|
seqlens_32_2d = seqlens_32
|
||||||
else:
|
else:
|
||||||
seqlens_32_2d = seqlens_32.unsqueeze(-1)
|
seqlens_32_2d = seqlens_32.unsqueeze(-1)
|
||||||
@@ -616,8 +636,6 @@ class Indexer(MultiPlatformOp):
|
|||||||
seqlens_32_2d, blocksize, self.sm_count
|
seqlens_32_2d, blocksize, self.sm_count
|
||||||
)
|
)
|
||||||
|
|
||||||
assert len(q_fp8.shape) == 3
|
|
||||||
q_fp8 = q_fp8.unsqueeze(1) # the next_n dim is 1 now
|
|
||||||
assert len(kv_cache_fp8.shape) == 2
|
assert len(kv_cache_fp8.shape) == 2
|
||||||
block_kv = page_size
|
block_kv = page_size
|
||||||
num_heads_kv = 1
|
num_heads_kv = 1
|
||||||
@@ -628,12 +646,10 @@ class Indexer(MultiPlatformOp):
|
|||||||
assert len(weights.shape) == 3
|
assert len(weights.shape) == 3
|
||||||
weights = weights.squeeze(2)
|
weights = weights.squeeze(2)
|
||||||
|
|
||||||
# When attn_tp_size > 1 or in the MAX_LEN padding mode, padding may exist in the hidden states,
|
|
||||||
# and it is necessary to extract the actual q length.
|
|
||||||
q_offset = sum(metadata.get_dsa_extend_len_cpu())
|
|
||||||
if _is_hip:
|
if _is_hip:
|
||||||
from aiter.ops.triton.pa_mqa_logits import deepgemm_fp8_paged_mqa_logits
|
from aiter.ops.triton.pa_mqa_logits import deepgemm_fp8_paged_mqa_logits
|
||||||
|
|
||||||
|
q_fp8 = q_fp8.unsqueeze(1)
|
||||||
batch_size, next_n, heads, _ = q_fp8.shape
|
batch_size, next_n, heads, _ = q_fp8.shape
|
||||||
logits = torch.empty(
|
logits = torch.empty(
|
||||||
(batch_size * next_n, max_seq_len),
|
(batch_size * next_n, max_seq_len),
|
||||||
@@ -651,7 +667,21 @@ class Indexer(MultiPlatformOp):
|
|||||||
Preshuffle=_use_aiter_preshuffle,
|
Preshuffle=_use_aiter_preshuffle,
|
||||||
KVBlockSize=block_kv,
|
KVBlockSize=block_kv,
|
||||||
)
|
)
|
||||||
|
elif use_dg_native:
|
||||||
|
# block_tables[::next_n] de-expands dsa_backend's repeat_interleave
|
||||||
|
# without a copy (DG only checks `stride(1) == 1`).
|
||||||
|
logits = deep_gemm.fp8_paged_mqa_logits(
|
||||||
|
q_fp8[:q_offset].view(B, next_n, q_fp8.shape[1], q_fp8.shape[2]),
|
||||||
|
kv_cache_fp8,
|
||||||
|
weights[:q_offset],
|
||||||
|
seqlens_32_2d,
|
||||||
|
block_tables[::next_n],
|
||||||
|
schedule_metadata,
|
||||||
|
max_seq_len,
|
||||||
|
clean_logits=False,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
|
q_fp8 = q_fp8.unsqueeze(1)
|
||||||
logits = deep_gemm.fp8_paged_mqa_logits(
|
logits = deep_gemm.fp8_paged_mqa_logits(
|
||||||
q_fp8[:q_offset],
|
q_fp8[:q_offset],
|
||||||
kv_cache_fp8,
|
kv_cache_fp8,
|
||||||
|
|||||||
@@ -50,7 +50,10 @@ from sglang.srt.layers.attention.utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||||
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.utils import is_cuda, is_hip
|
from sglang.srt.utils import is_cuda, is_hip, is_sm100_supported
|
||||||
|
|
||||||
|
if is_cuda():
|
||||||
|
import deep_gemm
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
@@ -153,6 +156,9 @@ class DSAMetadata:
|
|||||||
# DeepGEMM schedule metadata for paged MQA logits (decode/target_verify/draft_extend only).
|
# DeepGEMM schedule metadata for paged MQA logits (decode/target_verify/draft_extend only).
|
||||||
# Precomputed once per forward batch and reused across layers.
|
# Precomputed once per forward batch and reused across layers.
|
||||||
paged_mqa_schedule_metadata: Optional[torch.Tensor] = None
|
paged_mqa_schedule_metadata: Optional[torch.Tensor] = None
|
||||||
|
# 2D context_lens used to build the schedule above; the indexer reuses it
|
||||||
|
# as DG's `context_lens` arg so the broadcast doesn't rebuild per layer.
|
||||||
|
paged_mqa_ctx_lens_2d: Optional[torch.Tensor] = None
|
||||||
# The sum of sequence lengths for key, prefill only
|
# The sum of sequence lengths for key, prefill only
|
||||||
seq_lens_sum: Optional[int] = None
|
seq_lens_sum: Optional[int] = None
|
||||||
# The flattened 1D page table with shape (seq_lens_sum,), prefill only
|
# The flattened 1D page table with shape (seq_lens_sum,), prefill only
|
||||||
@@ -199,6 +205,7 @@ class DSAIndexerMetadata(BaseIndexerMetadata):
|
|||||||
topk_transform_method: TopkTransformMethod
|
topk_transform_method: TopkTransformMethod
|
||||||
topk_backend: DSATopKBackend = DSATopKBackend.SGL_KERNEL
|
topk_backend: DSATopKBackend = DSATopKBackend.SGL_KERNEL
|
||||||
paged_mqa_schedule_metadata: Optional[torch.Tensor] = None
|
paged_mqa_schedule_metadata: Optional[torch.Tensor] = None
|
||||||
|
paged_mqa_ctx_lens_2d: Optional[torch.Tensor] = None
|
||||||
force_unfused_topk: bool = False
|
force_unfused_topk: bool = False
|
||||||
|
|
||||||
def get_seqlens_int32(self) -> torch.Tensor:
|
def get_seqlens_int32(self) -> torch.Tensor:
|
||||||
@@ -380,6 +387,31 @@ class DeepseekSparseAttnBackend(
|
|||||||
else:
|
else:
|
||||||
self.workspace_buffer = None
|
self.workspace_buffer = None
|
||||||
|
|
||||||
|
def _build_paged_mqa_schedule_2d_ctx_lens(
|
||||||
|
self,
|
||||||
|
forward_mode: ForwardMode,
|
||||||
|
cache_seqlens_int32: torch.Tensor,
|
||||||
|
seqlens_expanded: torch.Tensor,
|
||||||
|
batch_size: int,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
# target_verify with next_n>=2 uses DG-native q=[B,next_n,H,D] which
|
||||||
|
# needs a [B, next_n] schedule; everything else stays per-token.
|
||||||
|
# TODO: SM90 supports DG-native next_n in {1,2} too — enable once
|
||||||
|
# validated; for now DG-native is SM100+ only.
|
||||||
|
next_n = self.speculative_num_draft_tokens
|
||||||
|
if (
|
||||||
|
forward_mode.is_target_verify()
|
||||||
|
and next_n
|
||||||
|
and next_n >= 2
|
||||||
|
and is_sm100_supported()
|
||||||
|
):
|
||||||
|
return cache_seqlens_int32.view(-1, 1).expand(-1, next_n).contiguous()
|
||||||
|
if forward_mode.is_target_verify() or forward_mode.is_draft_extend(
|
||||||
|
include_v2=True
|
||||||
|
):
|
||||||
|
return _to_2d_context_lens(seqlens_expanded, batch_size)
|
||||||
|
return _to_2d_context_lens(cache_seqlens_int32, batch_size)
|
||||||
|
|
||||||
def _get_fused_topk_page_table(self, topk_indices: torch.Tensor) -> torch.Tensor:
|
def _get_fused_topk_page_table(self, topk_indices: torch.Tensor) -> torch.Tensor:
|
||||||
if (
|
if (
|
||||||
self.dsa_topk_backend.is_sgl_kernel()
|
self.dsa_topk_backend.is_sgl_kernel()
|
||||||
@@ -654,33 +686,24 @@ class DeepseekSparseAttnBackend(
|
|||||||
dsa_cu_seqlens_q = self.get_device_int32_arange(len(dsa_cu_seqlens_k))
|
dsa_cu_seqlens_q = self.get_device_int32_arange(len(dsa_cu_seqlens_k))
|
||||||
|
|
||||||
paged_mqa_schedule_metadata = None
|
paged_mqa_schedule_metadata = None
|
||||||
# DeepGEMM paged MQA logits path needs a schedule metadata tensor.
|
paged_mqa_ctx_lens_2d = None
|
||||||
# Compute it once per forward batch and reuse it across layers.
|
|
||||||
if is_cuda() and (
|
if is_cuda() and (
|
||||||
forward_batch.forward_mode.is_decode_or_idle()
|
forward_batch.forward_mode.is_decode_or_idle()
|
||||||
or forward_batch.forward_mode.is_target_verify()
|
or forward_batch.forward_mode.is_target_verify()
|
||||||
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
|
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
|
||||||
):
|
):
|
||||||
try:
|
paged_mqa_ctx_lens_2d = self._build_paged_mqa_schedule_2d_ctx_lens(
|
||||||
import deep_gemm
|
forward_batch.forward_mode,
|
||||||
|
cache_seqlens_int32,
|
||||||
# NOTE: DeepGEMM paged path uses block_size=64.
|
seqlens_expanded,
|
||||||
seqlens_32 = (
|
forward_batch.batch_size,
|
||||||
seqlens_expanded
|
)
|
||||||
if (
|
# NOTE: block_kv arg must be 64 here — DG computes SPLIT_KV =
|
||||||
forward_batch.forward_mode.is_target_verify()
|
# block_kv * 4 and both DG's and the indexer's compute kernels
|
||||||
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
|
# require SPLIT_KV = 256; this is independent of the cache page size.
|
||||||
)
|
paged_mqa_schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
|
||||||
else cache_seqlens_int32
|
paged_mqa_ctx_lens_2d, 64, deep_gemm.get_num_sms()
|
||||||
)
|
)
|
||||||
seqlens_32_2d = _to_2d_context_lens(
|
|
||||||
seqlens_32, forward_batch.batch_size
|
|
||||||
)
|
|
||||||
paged_mqa_schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
|
|
||||||
seqlens_32_2d, 64, deep_gemm.get_num_sms()
|
|
||||||
)
|
|
||||||
except (ImportError, ModuleNotFoundError):
|
|
||||||
paged_mqa_schedule_metadata = None
|
|
||||||
|
|
||||||
metadata = DSAMetadata(
|
metadata = DSAMetadata(
|
||||||
page_size=self.real_page_size,
|
page_size=self.real_page_size,
|
||||||
@@ -701,6 +724,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
paged_mqa_schedule_metadata=paged_mqa_schedule_metadata,
|
paged_mqa_schedule_metadata=paged_mqa_schedule_metadata,
|
||||||
|
paged_mqa_ctx_lens_2d=paged_mqa_ctx_lens_2d,
|
||||||
dsa_cache_seqlens_int32=dsa_cache_seqlens_int32,
|
dsa_cache_seqlens_int32=dsa_cache_seqlens_int32,
|
||||||
dsa_cu_seqlens_q=dsa_cu_seqlens_q,
|
dsa_cu_seqlens_q=dsa_cu_seqlens_q,
|
||||||
dsa_cu_seqlens_k=dsa_cu_seqlens_k,
|
dsa_cu_seqlens_k=dsa_cu_seqlens_k,
|
||||||
@@ -947,28 +971,18 @@ class DeepseekSparseAttnBackend(
|
|||||||
real_page_table = self._transform_table_1_to_real(page_table_1)
|
real_page_table = self._transform_table_1_to_real(page_table_1)
|
||||||
|
|
||||||
paged_mqa_schedule_metadata = None
|
paged_mqa_schedule_metadata = None
|
||||||
|
paged_mqa_ctx_lens_2d = None
|
||||||
if is_cuda() and (
|
if is_cuda() and (
|
||||||
forward_mode.is_decode_or_idle()
|
forward_mode.is_decode_or_idle()
|
||||||
or forward_mode.is_target_verify()
|
or forward_mode.is_target_verify()
|
||||||
or forward_mode.is_draft_extend(include_v2=True)
|
or forward_mode.is_draft_extend(include_v2=True)
|
||||||
):
|
):
|
||||||
try:
|
paged_mqa_ctx_lens_2d = self._build_paged_mqa_schedule_2d_ctx_lens(
|
||||||
import deep_gemm
|
forward_mode, cache_seqlens_int32, seqlens_expanded, bs
|
||||||
|
)
|
||||||
seqlens_32 = (
|
paged_mqa_schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
|
||||||
seqlens_expanded
|
paged_mqa_ctx_lens_2d, 64, deep_gemm.get_num_sms()
|
||||||
if (
|
)
|
||||||
forward_mode.is_target_verify()
|
|
||||||
or forward_mode.is_draft_extend(include_v2=True)
|
|
||||||
)
|
|
||||||
else cache_seqlens_int32
|
|
||||||
)
|
|
||||||
seqlens_32_2d = _to_2d_context_lens(seqlens_32, bs)
|
|
||||||
paged_mqa_schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
|
|
||||||
seqlens_32_2d, 64, deep_gemm.get_num_sms()
|
|
||||||
)
|
|
||||||
except (ImportError, ModuleNotFoundError):
|
|
||||||
paged_mqa_schedule_metadata = None
|
|
||||||
|
|
||||||
metadata = DSAMetadata(
|
metadata = DSAMetadata(
|
||||||
page_size=self.real_page_size,
|
page_size=self.real_page_size,
|
||||||
@@ -980,6 +994,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
page_table_1=page_table_1,
|
page_table_1=page_table_1,
|
||||||
flashmla_metadata=flashmla_metadata,
|
flashmla_metadata=flashmla_metadata,
|
||||||
paged_mqa_schedule_metadata=paged_mqa_schedule_metadata,
|
paged_mqa_schedule_metadata=paged_mqa_schedule_metadata,
|
||||||
|
paged_mqa_ctx_lens_2d=paged_mqa_ctx_lens_2d,
|
||||||
dsa_cache_seqlens_int32=dsa_cache_seqlens_int32,
|
dsa_cache_seqlens_int32=dsa_cache_seqlens_int32,
|
||||||
dsa_cu_seqlens_q=dsa_cu_seqlens_q,
|
dsa_cu_seqlens_q=dsa_cu_seqlens_q,
|
||||||
dsa_cu_seqlens_k=dsa_cu_seqlens_k,
|
dsa_cu_seqlens_k=dsa_cu_seqlens_k,
|
||||||
@@ -1120,29 +1135,30 @@ class DeepseekSparseAttnBackend(
|
|||||||
or forward_mode.is_target_verify()
|
or forward_mode.is_target_verify()
|
||||||
or forward_mode.is_draft_extend(include_v2=True)
|
or forward_mode.is_draft_extend(include_v2=True)
|
||||||
):
|
):
|
||||||
try:
|
if forward_mode.is_draft_extend(include_v2=True):
|
||||||
import deep_gemm
|
schedule_seqlens_expanded = metadata.dsa_seqlens_expanded
|
||||||
|
else:
|
||||||
seqlens_32 = (
|
schedule_seqlens_expanded = seqlens_expanded
|
||||||
seqlens_expanded
|
seqlens_32_2d = self._build_paged_mqa_schedule_2d_ctx_lens(
|
||||||
if (
|
forward_mode,
|
||||||
forward_mode.is_target_verify()
|
metadata.cache_seqlens_int32,
|
||||||
or forward_mode.is_draft_extend(include_v2=True)
|
schedule_seqlens_expanded,
|
||||||
)
|
bs,
|
||||||
else metadata.cache_seqlens_int32
|
)
|
||||||
|
new_schedule = deep_gemm.get_paged_mqa_logits_metadata(
|
||||||
|
seqlens_32_2d, 64, deep_gemm.get_num_sms()
|
||||||
|
)
|
||||||
|
if metadata.paged_mqa_schedule_metadata is None:
|
||||||
|
object.__setattr__(
|
||||||
|
metadata, "paged_mqa_schedule_metadata", new_schedule
|
||||||
)
|
)
|
||||||
seqlens_32_2d = _to_2d_context_lens(seqlens_32, bs)
|
else:
|
||||||
new_schedule = deep_gemm.get_paged_mqa_logits_metadata(
|
metadata.paged_mqa_schedule_metadata.copy_(new_schedule)
|
||||||
seqlens_32_2d, 64, deep_gemm.get_num_sms()
|
# `copy_` preserves the buffer's data_ptr that the captured graph captured.
|
||||||
)
|
if metadata.paged_mqa_ctx_lens_2d is None:
|
||||||
if metadata.paged_mqa_schedule_metadata is None:
|
object.__setattr__(metadata, "paged_mqa_ctx_lens_2d", seqlens_32_2d)
|
||||||
object.__setattr__(
|
else:
|
||||||
metadata, "paged_mqa_schedule_metadata", new_schedule
|
metadata.paged_mqa_ctx_lens_2d.copy_(seqlens_32_2d)
|
||||||
)
|
|
||||||
else:
|
|
||||||
metadata.paged_mqa_schedule_metadata.copy_(new_schedule)
|
|
||||||
except (ImportError, ModuleNotFoundError):
|
|
||||||
object.__setattr__(metadata, "paged_mqa_schedule_metadata", None)
|
|
||||||
seqlens_expanded_size = seqlens_expanded.shape[0]
|
seqlens_expanded_size = seqlens_expanded.shape[0]
|
||||||
assert (
|
assert (
|
||||||
metadata.dsa_cache_seqlens_int32 is not None
|
metadata.dsa_cache_seqlens_int32 is not None
|
||||||
@@ -1334,27 +1350,28 @@ class DeepseekSparseAttnBackend(
|
|||||||
# deadlock the kernel when the runtime work decomposition diverges from
|
# deadlock the kernel when the runtime work decomposition diverges from
|
||||||
# the captured one).
|
# the captured one).
|
||||||
if is_cuda():
|
if is_cuda():
|
||||||
try:
|
if forward_mode.is_decode_or_idle():
|
||||||
import deep_gemm
|
seqlens_32_2d = _to_2d_context_lens(metadata.cache_seqlens_int32, bs)
|
||||||
|
else:
|
||||||
if forward_mode.is_decode_or_idle():
|
seqlens_32_2d = self._build_paged_mqa_schedule_2d_ctx_lens(
|
||||||
seqlens_32 = metadata.cache_seqlens_int32
|
forward_mode,
|
||||||
else:
|
metadata.cache_seqlens_int32,
|
||||||
seqlens_32 = metadata.dsa_seqlens_expanded[
|
metadata.dsa_seqlens_expanded,
|
||||||
: precomputed.seqlens_expanded_size
|
bs,
|
||||||
]
|
|
||||||
seqlens_32_2d = _to_2d_context_lens(seqlens_32, bs)
|
|
||||||
new_schedule = deep_gemm.get_paged_mqa_logits_metadata(
|
|
||||||
seqlens_32_2d, 64, deep_gemm.get_num_sms()
|
|
||||||
)
|
)
|
||||||
if metadata.paged_mqa_schedule_metadata is None:
|
new_schedule = deep_gemm.get_paged_mqa_logits_metadata(
|
||||||
object.__setattr__(
|
seqlens_32_2d, 64, deep_gemm.get_num_sms()
|
||||||
metadata, "paged_mqa_schedule_metadata", new_schedule
|
)
|
||||||
)
|
if metadata.paged_mqa_schedule_metadata is None:
|
||||||
else:
|
object.__setattr__(
|
||||||
metadata.paged_mqa_schedule_metadata.copy_(new_schedule)
|
metadata, "paged_mqa_schedule_metadata", new_schedule
|
||||||
except (ImportError, ModuleNotFoundError):
|
)
|
||||||
pass
|
else:
|
||||||
|
metadata.paged_mqa_schedule_metadata.copy_(new_schedule)
|
||||||
|
if metadata.paged_mqa_ctx_lens_2d is None:
|
||||||
|
object.__setattr__(metadata, "paged_mqa_ctx_lens_2d", seqlens_32_2d)
|
||||||
|
else:
|
||||||
|
metadata.paged_mqa_ctx_lens_2d.copy_(seqlens_32_2d)
|
||||||
|
|
||||||
self.forward_metadata = metadata
|
self.forward_metadata = metadata
|
||||||
|
|
||||||
@@ -2320,6 +2337,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
),
|
),
|
||||||
topk_backend=self.dsa_topk_backend,
|
topk_backend=self.dsa_topk_backend,
|
||||||
paged_mqa_schedule_metadata=self.forward_metadata.paged_mqa_schedule_metadata,
|
paged_mqa_schedule_metadata=self.forward_metadata.paged_mqa_schedule_metadata,
|
||||||
|
paged_mqa_ctx_lens_2d=self.forward_metadata.paged_mqa_ctx_lens_2d,
|
||||||
force_unfused_topk=force_unfused,
|
force_unfused_topk=force_unfused,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user