[SM120] Only split touched SWA pages in FlashMLA page-split kernel (#32320)
Co-authored-by: 百麒 <yaozhong.lyz@alibaba-inc.com> Co-authored-by: David Orman <ormandj@corenode.com>
This commit is contained in:
@@ -13,6 +13,7 @@ separate region at the end of each page.
|
||||
|
||||
import logging
|
||||
import math
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import triton
|
||||
@@ -305,8 +306,15 @@ def _page_split_kernel(
|
||||
DST_SCALE_OFF: tl.constexpr, # 64 * 576 = 36864
|
||||
RATIO: tl.constexpr, # 4
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
mask_ptr,
|
||||
HAS_MASK: tl.constexpr,
|
||||
):
|
||||
"""Fused page-split: copy data+scale for all sub-pages in one kernel."""
|
||||
"""Fused page-split: copy data+scale for all sub-pages in one kernel.
|
||||
|
||||
When HAS_MASK is set, only pages flagged in ``mask_ptr`` (int8, 1=touched)
|
||||
are copied; untouched pages are skipped so the kernel no longer rewrites the
|
||||
entire KV pool every decode step.
|
||||
"""
|
||||
pid = tl.program_id(0)
|
||||
page_idx = pid // RATIO
|
||||
sub = pid % RATIO
|
||||
@@ -314,6 +322,10 @@ def _page_split_kernel(
|
||||
if page_idx >= N_pages:
|
||||
return
|
||||
|
||||
if HAS_MASK:
|
||||
if tl.load(mask_ptr + page_idx) == 0:
|
||||
return
|
||||
|
||||
src_base = src_ptr + page_idx * src_stride0
|
||||
dst_base = dst_ptr + (page_idx * RATIO + sub) * dst_stride0
|
||||
|
||||
@@ -334,11 +346,43 @@ def _page_split_kernel(
|
||||
tl.store(dst_base + DST_SCALE_OFF + offs, vals, mask=mask)
|
||||
|
||||
|
||||
def _split_kv_pages_to_64(kv_u8: torch.Tensor, src_pbs: int) -> torch.Tensor:
|
||||
@triton.jit
|
||||
def _page_mark_kernel(
|
||||
indices_ptr,
|
||||
mask_ptr,
|
||||
N_idx,
|
||||
SRC_PBS: tl.constexpr,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
"""Mark touched source pages (1 byte each) from token-level indices.
|
||||
|
||||
``indices`` are token indices into the pbs=SRC_PBS SWA pool; -1 = invalid.
|
||||
Each valid token marks ``mask[token // SRC_PBS] = 1``. Concurrent stores of
|
||||
the same value 1 are safe (no atomic needed).
|
||||
"""
|
||||
pid = tl.program_id(0)
|
||||
if pid >= N_idx:
|
||||
return
|
||||
idx = tl.load(indices_ptr + pid)
|
||||
if idx < 0:
|
||||
return
|
||||
page = idx // SRC_PBS
|
||||
tl.store(mask_ptr + page, 1)
|
||||
|
||||
|
||||
def _split_kv_pages_to_64(
|
||||
kv_u8: torch.Tensor,
|
||||
src_pbs: int,
|
||||
touched_indices: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Split pbs=N footer-format pages into pbs=64 footer-format pages.
|
||||
|
||||
Uses a fused Triton kernel to do all sub-page copies in a single launch
|
||||
instead of 8 separate copy kernels (4 sub-pages × 2 regions).
|
||||
When ``touched_indices`` (token-level int32 indices into the pbs=src_pbs
|
||||
SWA pool, -1 = invalid) is provided, only the source pages that actually
|
||||
contain a referenced token are copied. This avoids rewriting the entire KV
|
||||
pool on every decode step (only ~2*batch pages are touched vs the full
|
||||
pool). The output buffer is persistent and reused across steps; untouched
|
||||
dst pages simply retain their (unreferenced) stale data.
|
||||
"""
|
||||
assert src_pbs % _PBS_DST == 0 and src_pbs >= _PBS_DST
|
||||
if src_pbs == _PBS_DST:
|
||||
@@ -373,6 +417,35 @@ def _split_kv_pages_to_64(kv_u8: torch.Tensor, src_pbs: int) -> torch.Tensor:
|
||||
else:
|
||||
src_stride0 = src_2d.stride(0)
|
||||
|
||||
use_mask = touched_indices is not None and touched_indices.numel() > 0
|
||||
mask_ptr = src_2d # dummy, never dereferenced when HAS_MASK is False
|
||||
if use_mask:
|
||||
# Persistent per-device int8 mask, zeroed each call (cheap memset,
|
||||
# captured cleanly by CUDA graph). 1 = page is referenced this step.
|
||||
mkey = f"flash_mla_sm120_mask:{dev}"
|
||||
mbuf = buffers.get(mkey)
|
||||
if mbuf is None or mbuf.shape[0] < N:
|
||||
# The first allocation can happen under inference mode (autotune),
|
||||
# but the buffer is zeroed again later during CUDA graph capture
|
||||
# outside inference mode -- an inference tensor cannot be mutated
|
||||
# there, so force a normal tensor.
|
||||
with torch.inference_mode(False):
|
||||
mbuf = torch.empty(N, dtype=torch.int8, device=dev)
|
||||
buffers[mkey] = mbuf
|
||||
mask = mbuf[:N]
|
||||
mask.zero_()
|
||||
idx_flat = touched_indices.reshape(-1).contiguous()
|
||||
if idx_flat.dtype != torch.int32:
|
||||
idx_flat = idx_flat.to(torch.int32)
|
||||
_page_mark_kernel[(idx_flat.numel(),)](
|
||||
idx_flat,
|
||||
mask,
|
||||
idx_flat.numel(),
|
||||
src_pbs, # SRC_PBS
|
||||
1024, # BLOCK (unused, kept for JIT signature)
|
||||
)
|
||||
mask_ptr = mask
|
||||
|
||||
grid = (N * ratio,)
|
||||
_page_split_kernel[grid](
|
||||
src_2d,
|
||||
@@ -386,6 +459,8 @@ def _split_kv_pages_to_64(kv_u8: torch.Tensor, src_pbs: int) -> torch.Tensor:
|
||||
_PBS_DST * _NOPE_ROPE_STRIDE, # DST_SCALE_OFF = 36864
|
||||
ratio, # RATIO = 4
|
||||
1024, # BLOCK_SIZE
|
||||
mask_ptr,
|
||||
use_mask, # HAS_MASK
|
||||
)
|
||||
|
||||
bpt = _NOPE_ROPE_STRIDE + _SCALE_STRIDE # 584
|
||||
@@ -424,10 +499,19 @@ def _flash_mla_flashinfer(
|
||||
B, _, H, D = q.shape # (batch, 1, num_heads, head_dim)
|
||||
dev = q.device
|
||||
|
||||
# Indices: no remapping needed (page-split preserves token addressing).
|
||||
idx = indices.squeeze(1) if indices.dim() == 3 else indices
|
||||
|
||||
# --- Page-split: convert pbs=N kv_cache to pbs=64 view ---
|
||||
# Only the SWA pages actually referenced by `idx` are copied (the rest of
|
||||
# the persistent dst buffer is left untouched and never read).
|
||||
kv_u8 = k_cache.view(torch.uint8) if k_cache.dtype != torch.uint8 else k_cache
|
||||
src_pbs = k_cache.shape[1] if k_cache.ndim >= 3 else _PBS_SRC
|
||||
kv_64 = _split_kv_pages_to_64(kv_u8, src_pbs) if src_pbs != _PBS_DST else kv_u8
|
||||
kv_64 = (
|
||||
_split_kv_pages_to_64(kv_u8, src_pbs, touched_indices=idx)
|
||||
if src_pbs != _PBS_DST
|
||||
else kv_u8
|
||||
)
|
||||
|
||||
extra_kv_u8 = (
|
||||
extra_k_cache.view(torch.uint8)
|
||||
@@ -436,8 +520,6 @@ def _flash_mla_flashinfer(
|
||||
)
|
||||
extra_kv_64 = extra_kv_u8
|
||||
|
||||
# Indices: no remapping needed (page-split preserves token addressing).
|
||||
idx = indices.squeeze(1) if indices.dim() == 3 else indices
|
||||
extra_idx = (
|
||||
extra_indices.squeeze(1)
|
||||
if extra_indices is not None and extra_indices.dim() == 3
|
||||
|
||||
Reference in New Issue
Block a user