Sgl flashmla (#26132)
This commit is contained in:
@@ -1,3 +1,4 @@
|
||||
import dataclasses
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
@@ -14,10 +15,33 @@ _IMPORT_ERROR = ImportError(
|
||||
)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class FlashMLASchedMeta:
|
||||
"""Tile scheduler metadata for the newer FlashMLA Python API."""
|
||||
|
||||
@dataclasses.dataclass
|
||||
class Config:
|
||||
b: int
|
||||
s_q: int
|
||||
h_q: int
|
||||
page_block_size: int
|
||||
h_k: int
|
||||
causal: bool
|
||||
is_fp8_kvcache: bool
|
||||
topk: Optional[int]
|
||||
extra_page_block_size: Optional[int]
|
||||
extra_topk: Optional[int]
|
||||
|
||||
have_initialized: bool = False
|
||||
config: Optional[Config] = None
|
||||
tile_scheduler_metadata: Optional[torch.Tensor] = None
|
||||
num_splits: Optional[torch.Tensor] = None
|
||||
|
||||
|
||||
def get_mla_metadata(
|
||||
cache_seqlens: torch.Tensor,
|
||||
num_q_tokens_per_head_k: int,
|
||||
num_heads_k: int,
|
||||
cache_seqlens: Optional[torch.Tensor] = None,
|
||||
num_q_tokens_per_head_k: Optional[int] = None,
|
||||
num_heads_k: Optional[int] = None,
|
||||
num_heads_q: Optional[int] = None,
|
||||
is_fp8_kvcache: bool = False,
|
||||
topk: Optional[int] = None,
|
||||
@@ -38,6 +62,12 @@ def get_mla_metadata(
|
||||
if _flashmla_import_error is not None:
|
||||
raise _IMPORT_ERROR from _flashmla_import_error
|
||||
|
||||
if cache_seqlens is None:
|
||||
return FlashMLASchedMeta(), None
|
||||
|
||||
assert num_q_tokens_per_head_k is not None
|
||||
assert num_heads_k is not None
|
||||
|
||||
if is_fp8_kvcache and topk is None:
|
||||
return torch.ops.sgl_kernel.get_mla_decoding_metadata_dense_fp8.default(
|
||||
cache_seqlens,
|
||||
@@ -57,17 +87,22 @@ def get_mla_metadata(
|
||||
def flash_mla_with_kvcache(
|
||||
q: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
block_table: torch.Tensor,
|
||||
cache_seqlens: torch.Tensor,
|
||||
block_table: Optional[torch.Tensor],
|
||||
cache_seqlens: Optional[torch.Tensor],
|
||||
head_dim_v: int,
|
||||
tile_scheduler_metadata: torch.Tensor,
|
||||
num_splits: torch.Tensor,
|
||||
tile_scheduler_metadata: torch.Tensor | FlashMLASchedMeta,
|
||||
num_splits: Optional[torch.Tensor] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
descale_q: torch.Tensor | None = None,
|
||||
descale_k: torch.Tensor | None = None,
|
||||
is_fp8_kvcache: bool = False,
|
||||
indices: Optional[torch.Tensor] = None,
|
||||
attn_sink: Optional[torch.Tensor] = None,
|
||||
extra_k_cache: Optional[torch.Tensor] = None,
|
||||
extra_indices_in_kvcache: Optional[torch.Tensor] = None,
|
||||
topk_length: Optional[torch.Tensor] = None,
|
||||
extra_topk_length: Optional[torch.Tensor] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Arguments:
|
||||
@@ -94,6 +129,34 @@ def flash_mla_with_kvcache(
|
||||
|
||||
if softmax_scale is None:
|
||||
softmax_scale = q.shape[-1] ** (-0.5)
|
||||
if isinstance(tile_scheduler_metadata, FlashMLASchedMeta):
|
||||
return _flash_mla_with_kvcache_sched_meta(
|
||||
q=q,
|
||||
k_cache=k_cache,
|
||||
block_table=block_table,
|
||||
cache_seqlens=cache_seqlens,
|
||||
head_dim_v=head_dim_v,
|
||||
sched_meta=tile_scheduler_metadata,
|
||||
num_splits=num_splits,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
is_fp8_kvcache=is_fp8_kvcache,
|
||||
indices=indices,
|
||||
attn_sink=attn_sink,
|
||||
extra_k_cache=extra_k_cache,
|
||||
extra_indices_in_kvcache=extra_indices_in_kvcache,
|
||||
topk_length=topk_length,
|
||||
extra_topk_length=extra_topk_length,
|
||||
)
|
||||
|
||||
assert num_splits is not None
|
||||
assert block_table is not None
|
||||
assert cache_seqlens is not None
|
||||
assert attn_sink is None
|
||||
assert extra_k_cache is None
|
||||
assert extra_indices_in_kvcache is None
|
||||
assert topk_length is None
|
||||
assert extra_topk_length is None
|
||||
if indices is not None:
|
||||
assert causal == False, "causal must be `false` if sparse attention is enabled."
|
||||
assert (descale_q is None) == (
|
||||
@@ -127,16 +190,131 @@ def flash_mla_with_kvcache(
|
||||
num_splits,
|
||||
is_fp8_kvcache,
|
||||
indices,
|
||||
attn_sink,
|
||||
extra_k_cache,
|
||||
extra_indices_in_kvcache,
|
||||
topk_length,
|
||||
extra_topk_length,
|
||||
)
|
||||
return out, softmax_lse
|
||||
|
||||
|
||||
def _flash_mla_with_kvcache_sched_meta(
|
||||
q: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
block_table: Optional[torch.Tensor],
|
||||
cache_seqlens: Optional[torch.Tensor],
|
||||
head_dim_v: int,
|
||||
sched_meta: FlashMLASchedMeta,
|
||||
num_splits: Optional[torch.Tensor],
|
||||
softmax_scale: float,
|
||||
causal: bool,
|
||||
is_fp8_kvcache: bool,
|
||||
indices: Optional[torch.Tensor],
|
||||
attn_sink: Optional[torch.Tensor],
|
||||
extra_k_cache: Optional[torch.Tensor],
|
||||
extra_indices_in_kvcache: Optional[torch.Tensor],
|
||||
topk_length: Optional[torch.Tensor],
|
||||
extra_topk_length: Optional[torch.Tensor],
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
assert num_splits is None, "num_splits must be None with FlashMLASchedMeta"
|
||||
|
||||
topk = indices.shape[-1] if indices is not None else None
|
||||
extra_page_block_size = (
|
||||
extra_k_cache.shape[1] if extra_k_cache is not None else None
|
||||
)
|
||||
extra_topk = (
|
||||
extra_indices_in_kvcache.shape[-1]
|
||||
if extra_indices_in_kvcache is not None
|
||||
else None
|
||||
)
|
||||
|
||||
if not sched_meta.have_initialized:
|
||||
sched_meta.have_initialized = True
|
||||
sched_meta.config = FlashMLASchedMeta.Config(
|
||||
b=q.shape[0],
|
||||
s_q=q.shape[1],
|
||||
h_q=q.shape[2],
|
||||
page_block_size=k_cache.shape[1],
|
||||
h_k=k_cache.shape[2],
|
||||
causal=causal,
|
||||
is_fp8_kvcache=is_fp8_kvcache,
|
||||
topk=topk,
|
||||
extra_page_block_size=extra_page_block_size,
|
||||
extra_topk=extra_topk,
|
||||
)
|
||||
else:
|
||||
helper_msg = (
|
||||
" Input arguments are inconsistent with FlashMLASchedMeta. Reuse a "
|
||||
"scheduler only for matching tensor shapes and sparse settings."
|
||||
)
|
||||
assert sched_meta.config is not None
|
||||
assert sched_meta.config.b == q.shape[0], helper_msg
|
||||
assert sched_meta.config.s_q == q.shape[1], helper_msg
|
||||
assert sched_meta.config.h_q == q.shape[2], helper_msg
|
||||
assert sched_meta.config.page_block_size == k_cache.shape[1], helper_msg
|
||||
assert sched_meta.config.h_k == k_cache.shape[2], helper_msg
|
||||
assert sched_meta.config.causal == causal, helper_msg
|
||||
assert sched_meta.config.is_fp8_kvcache == is_fp8_kvcache, helper_msg
|
||||
assert sched_meta.config.topk == topk, helper_msg
|
||||
assert (
|
||||
sched_meta.config.extra_page_block_size == extra_page_block_size
|
||||
), helper_msg
|
||||
assert sched_meta.config.extra_topk == extra_topk, helper_msg
|
||||
|
||||
if topk is not None:
|
||||
assert not causal, "causal must be False when sparse attention is enabled"
|
||||
assert is_fp8_kvcache, "is_fp8_kvcache must be True for sparse attention"
|
||||
out, lse, new_tile_scheduler_metadata, new_num_splits = (
|
||||
torch.ops.sgl_kernel.sparse_decode_fwd.default(
|
||||
q,
|
||||
k_cache,
|
||||
indices,
|
||||
topk_length,
|
||||
attn_sink,
|
||||
sched_meta.tile_scheduler_metadata,
|
||||
sched_meta.num_splits,
|
||||
extra_k_cache,
|
||||
extra_indices_in_kvcache,
|
||||
extra_topk_length,
|
||||
head_dim_v,
|
||||
softmax_scale,
|
||||
)
|
||||
)
|
||||
else:
|
||||
assert block_table is not None and cache_seqlens is not None
|
||||
assert attn_sink is None
|
||||
assert extra_k_cache is None
|
||||
assert extra_indices_in_kvcache is None
|
||||
assert topk_length is None
|
||||
assert extra_topk_length is None
|
||||
out, lse, new_tile_scheduler_metadata, new_num_splits = (
|
||||
torch.ops.sgl_kernel.dense_decode_fwd.default(
|
||||
q,
|
||||
k_cache,
|
||||
head_dim_v,
|
||||
cache_seqlens,
|
||||
block_table,
|
||||
softmax_scale,
|
||||
causal,
|
||||
sched_meta.tile_scheduler_metadata,
|
||||
sched_meta.num_splits,
|
||||
)
|
||||
)
|
||||
|
||||
sched_meta.tile_scheduler_metadata = new_tile_scheduler_metadata
|
||||
sched_meta.num_splits = new_num_splits
|
||||
return out, lse
|
||||
|
||||
|
||||
def flash_mla_sparse_fwd(
|
||||
q: torch.Tensor,
|
||||
kv: torch.Tensor,
|
||||
indices: torch.Tensor,
|
||||
sm_scale: float,
|
||||
d_v: int = 512,
|
||||
attn_sink: Optional[torch.Tensor] = None,
|
||||
topk_length: Optional[torch.Tensor] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Sparse attention prefill kernel
|
||||
@@ -159,6 +337,6 @@ def flash_mla_sparse_fwd(
|
||||
raise _IMPORT_ERROR from _flashmla_import_error
|
||||
|
||||
results = torch.ops.sgl_kernel.sparse_prefill_fwd.default(
|
||||
q, kv, indices, sm_scale, d_v
|
||||
q, kv, indices, sm_scale, d_v, attn_sink, topk_length
|
||||
)
|
||||
return results
|
||||
|
||||
Reference in New Issue
Block a user