feat(unified-memory): read unified pool from attention backends fa3/flashinfer/trtllm_mha/flashmla (#34613)
Co-authored-by: Caihua Li <caihua.li@bytedance.com> Co-authored-by: Claude Fable 5 <noreply@anthropic.com> Co-authored-by: Cheng Wan <cheng.wan@radixark.ai>
This commit is contained in:
co-authored by
Caihua Li
Claude Fable 5
Cheng Wan
parent
29578d5578
commit
8bb776dc48
@@ -13,18 +13,28 @@
|
||||
# ==============================================================================
|
||||
"""`--enable-page-major-kv-layout` full-attention backend allowlist.
|
||||
|
||||
The page-major envelope K/V views are strided, which only the Triton attention
|
||||
kernels read. The one exception is the unified-memory MLA pool: it exposes each
|
||||
layer as a contiguous view (`build_mla_views`), so the paged MLA
|
||||
backends can read it directly once their kv_indices / block tables are remapped
|
||||
to kernel-facing ids -- `fa3`, `flashinfer`'s MLA backend, and `trtllm_mla` with its
|
||||
`cutedsl_mla` / `tokenspeed_mla` subclasses.
|
||||
Two-way gate (see `_handle_page_major_kv_layout`), because the unified pool
|
||||
exposes per-layer views and nothing else:
|
||||
* unified-memory MLA models (`build_mla_views`) allow the whole wired
|
||||
paged MLA family -- `fa3`, `flashinfer`'s MLA backend, `trtllm_mla` with
|
||||
its `cutedsl_mla` / `tokenspeed_mla` subclasses, and `flashmla` (ps=64
|
||||
snap);
|
||||
* unified-memory MHA/SWA models (`build_mha_views`) allow `fa3` /
|
||||
`fa4` / `flashinfer` / `trtllm_mha` alongside Triton;
|
||||
* plain `--enable-page-major-kv-layout` without the unified pool keeps the
|
||||
envelope-strided 4-D views only the stride-aware Triton kernels read.
|
||||
|
||||
Pinned here so the exception cannot silently widen to a backend that has no
|
||||
dense-id remapping (`flashmla`, `cutlass_mla`, ...) or leak into the MHA path.
|
||||
`fa3` matters most: it is the resolved default on pre-Blackwell hosts, so it is
|
||||
the one entry whose absence used to make `--enable-unified-memory` fail to boot
|
||||
under its own default configuration.
|
||||
The same handler also screens the pool itself: the dense MHA/SWA views need
|
||||
uniform K/V rows, so an asymmetric-K/V model (MiMoV2: head_dim 192 !=
|
||||
v_head_dim 128) cannot run `--enable-unified-memory` at all and is rejected on
|
||||
EVERY backend, Triton included. MLA models are exempt -- their sub-pool keeps
|
||||
one latent row per layer, and several MLA configs (Kimi-Linear: head_dim 72,
|
||||
v_head_dim 128) report asymmetric dims while running the unified pool today.
|
||||
|
||||
Pinned here so no arm silently widens to an unwired backend (`cutlass_mla`,
|
||||
`aiter`) and no arm silently narrows: `fa3` is the resolved default on
|
||||
pre-Blackwell hosts, so its absence from an arm makes `--enable-unified-memory`
|
||||
fail to boot under its own default configuration.
|
||||
|
||||
python -m pytest test/registered/unit/server_args/test_page_major_backend_allowlist.py -v
|
||||
"""
|
||||
@@ -96,13 +106,23 @@ class TestPageMajorBackendAllowlist(unittest.TestCase):
|
||||
"flashinfer",
|
||||
"cutedsl_mla",
|
||||
"tokenspeed_mla",
|
||||
"flashmla",
|
||||
)
|
||||
# No dense-id remapping: must stay rejected until they get one.
|
||||
UNWIRED_BACKENDS = ("flashmla", "cutlass_mla", "trtllm_mha", "aiter")
|
||||
# Wired for the dense per-layer MHA/SWA views (uniform-row models).
|
||||
DENSE_MHA_BACKENDS = ("fa3", "fa4", "flashinfer", "trtllm_mha")
|
||||
# MLA-family kernels that must never leak into the MHA arm.
|
||||
MLA_ONLY_BACKENDS = ("trtllm_mla", "cutedsl_mla", "tokenspeed_mla", "flashmla")
|
||||
# No dense-id wiring anywhere: must stay rejected until they get one.
|
||||
UNWIRED_BACKENDS = ("cutlass_mla", "aiter")
|
||||
|
||||
def test_triton_always_allowed(self):
|
||||
def test_triton_allowed_on_every_arm(self):
|
||||
"""Triton reads both view families, so it is the one backend neither
|
||||
the MLA nor the MHA arm can narrow away."""
|
||||
for use_mla in (True, False):
|
||||
self.assertTrue(_accepts("triton", use_mla=use_mla))
|
||||
# The uniform-row screen is a property of the model, not of the
|
||||
# backend, so it rejects even Triton.
|
||||
self.assertFalse(_accepts("triton", use_mla=False, has_asymmetric_kv=True))
|
||||
|
||||
def test_dense_mla_backends_allowed_under_unified_mla(self):
|
||||
for backend in self.DENSE_MLA_BACKENDS:
|
||||
@@ -111,21 +131,18 @@ class TestPageMajorBackendAllowlist(unittest.TestCase):
|
||||
f"{backend} should be allowed with the unified-memory MLA pool",
|
||||
)
|
||||
|
||||
def test_dense_mla_backends_rejected_for_mha(self):
|
||||
"""The per-layer-view exception is MLA-only -- MHA sub-pools stay strided."""
|
||||
for backend in self.DENSE_MLA_BACKENDS:
|
||||
self.assertFalse(
|
||||
def test_dense_mha_backends_allowed_for_uniform_row_models(self):
|
||||
for backend in self.DENSE_MHA_BACKENDS:
|
||||
self.assertTrue(
|
||||
_accepts(backend, use_mla=False),
|
||||
f"{backend} must stay rejected for a non-MLA model",
|
||||
f"{backend} should be allowed for a uniform-row MHA model",
|
||||
)
|
||||
|
||||
def test_dense_mla_backends_rejected_without_unified_memory(self):
|
||||
"""Plain --enable-page-major-kv-layout (no unified pool) keeps the
|
||||
strided views, so only Triton can read them."""
|
||||
for backend in self.DENSE_MLA_BACKENDS:
|
||||
def test_mla_only_backends_rejected_for_mha(self):
|
||||
for backend in self.MLA_ONLY_BACKENDS:
|
||||
self.assertFalse(
|
||||
_accepts(backend, use_mla=True, unified=False),
|
||||
f"{backend} must stay rejected without --enable-unified-memory",
|
||||
_accepts(backend, use_mla=False),
|
||||
f"{backend} is an MLA kernel and must stay out of the MHA arm",
|
||||
)
|
||||
|
||||
def test_plain_page_major_arm_is_gated_at_boot(self):
|
||||
@@ -143,7 +160,7 @@ class TestPageMajorBackendAllowlist(unittest.TestCase):
|
||||
"""head_dim != v_head_dim (MiMoV2): no uniform rows, so no per-layer views
|
||||
and no unified pool. The rejection is the POOL's, not a backend's, so
|
||||
it must fire on every backend -- Triton included."""
|
||||
for backend in ("triton",) + self.DENSE_MLA_BACKENDS:
|
||||
for backend in ("triton",) + self.DENSE_MHA_BACKENDS:
|
||||
self.assertFalse(
|
||||
_accepts(backend, use_mla=False, has_asymmetric_kv=True),
|
||||
f"--enable-unified-memory + {backend} must be rejected for an "
|
||||
@@ -162,12 +179,25 @@ class TestPageMajorBackendAllowlist(unittest.TestCase):
|
||||
"K/V head dims",
|
||||
)
|
||||
|
||||
def test_page_major_rejected_without_unified_memory(self):
|
||||
"""--enable-page-major-kv-layout without the unified pool is rejected
|
||||
outright, Triton included: the static page-major arm went away with
|
||||
the strided views and awaits its per-layer-view reimplementation."""
|
||||
for backend in ("triton",) + tuple(
|
||||
set(self.DENSE_MLA_BACKENDS + self.DENSE_MHA_BACKENDS)
|
||||
):
|
||||
for use_mla in (True, False):
|
||||
self.assertFalse(
|
||||
_accepts(backend, use_mla=use_mla, unified=False),
|
||||
f"{backend} must stay rejected without --enable-unified-memory",
|
||||
)
|
||||
|
||||
def test_unwired_backends_always_rejected(self):
|
||||
for backend in self.UNWIRED_BACKENDS:
|
||||
for use_mla in (True, False):
|
||||
self.assertFalse(
|
||||
_accepts(backend, use_mla=use_mla),
|
||||
f"{backend} has no dense-id remapping and must be rejected",
|
||||
f"{backend} has no dense-id wiring and must be rejected",
|
||||
)
|
||||
|
||||
def test_helion_linear_attention_is_kda_only(self):
|
||||
|
||||
Reference in New Issue
Block a user