Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
18 commits
Select commit Hold shift + click to select a range
31e1150
[None][perf] Safely share M3 draft KV in aggregated serving
zheyuf Aug 15, 2026
a4c7a91
[None][test] Switch the MiniMax-M3 Eagle3 CI to the GQA draft head
zheyuf Aug 11, 2026
9ae761b
[None][fix] Synchronize multistream in-place side effects
zheyuf Aug 18, 2026
d1617e1
[None][test] Remove multistream scheduler unit test
zheyuf Aug 18, 2026
688b719
[None][fix] Map grouped HND draft KV across disagg peers
zheyuf Aug 19, 2026
648f363
[None][test] Gate complete M3 Eagle3 acceptance windows
zheyuf Aug 19, 2026
0672e91
[None][perf] Use native P128 MiniMax-M3 draft KV pages
zheyuf Aug 21, 2026
e7c9ec9
[None][perf] Publish only disaggregated Eagle captures
zheyuf Aug 24, 2026
43f9ddf
[None][refactor] Tighten native P128 publication contracts
zheyuf Aug 24, 2026
1fabd26
[None][refactor] Scope draft cache helper contracts
zheyuf Aug 26, 2026
bdde9e8
[None][refactor] Remove retired M3 draft cache contracts
zheyuf Aug 27, 2026
6ac0cb8
Merge upstream/feat/m3_with_msa into perf/minimax-m3-native-p128-draft
zheyuf Aug 27, 2026
a1c44ad
Merge upstream/feat/m3_with_msa into native P128 draft
zheyuf Aug 29, 2026
cded209
[None][test] Use GQA Eagle in M3 NVFP4 smoke
zheyuf Aug 31, 2026
cde7b84
[None][fix] Scale the native P128 draft view by the draft layer's stride
zheyuf Sep 1, 2026
9f7e033
[None][chore] Trim MiniMax-M3 native P128 draft view leftovers
zheyuf Sep 2, 2026
f461eb0
[None][chore] Drop the inert Eagle hidden-state publication path
zheyuf Sep 2, 2026
343310a
[None][fix] Make graph exit depend on unreturned in-place side effects
zheyuf Sep 3, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,6 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""
FlashInfer TRTLLM-Gen FMHA

Expand Down Expand Up @@ -66,7 +65,6 @@
TrtllmAttentionMetadata,
)


_MULTI_CTAS_KV_COUNTER_ALIGNMENT = 8


Expand Down Expand Up @@ -752,8 +750,15 @@ def _is_supported_with_reason(
return False, f"non-positive tokens_per_block ({tokens_per_block})."
if tokens_per_block & (tokens_per_block - 1) != 0:
return False, f"tokens_per_block ({tokens_per_block}) that is not a power of 2."
if tokens_per_block not in self.SUPPORTED_TOKENS_PER_BLOCK:
supported = sorted(self.SUPPORTED_TOKENS_PER_BLOCK)
# P128 is not exported for every TRTLLM-Gen shape family, so keep it
# out of the global allowlist. Cache-manager views whose exact shapes
# are exported may opt in explicitly.
extra_tokens_per_block = getattr(
meta.kv_cache_manager, "trtllm_gen_extra_tokens_per_block", ()
)
supported_tokens_per_block = self.SUPPORTED_TOKENS_PER_BLOCK | set(extra_tokens_per_block)
if tokens_per_block not in supported_tokens_per_block:
supported = sorted(supported_tokens_per_block)
return False, f"tokens_per_block ({tokens_per_block}). Supported: {supported}."

return True, ""
Expand Down
16 changes: 0 additions & 16 deletions tensorrt_llm/_torch/attention_backend/fmha/msa_sparse_gqa.py
Original file line number Diff line number Diff line change
Expand Up @@ -235,22 +235,6 @@ def run_msa_paged_gqa(
)
return

if getattr(kv_cache_manager, "is_fp8_subpaged_layer", lambda _layer_idx: False)(layer_idx):
from tensorrt_llm._torch.attention_backend.sparse.minimax_m3.trtllm_gen_dense_decode import (
minimax_m3_trtllm_gen_dense_attention,
)

minimax_m3_trtllm_gen_dense_attention(
q_view,
kv_cache_manager,
layer_idx,
metadata,
sm_scale=sm_scale,
output=out_view,
kv_scale_quant_orig=kv_scale_quant_orig,
)
return

k_paged, v_paged = msa_paged_kv(kv_cache_manager, layer_idx)

# Leading query tokens fmha_sm100 must still run: the whole batch until a
Expand Down

Large diffs are not rendered by default.

Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,7 @@
)
from .trtllm_gen_dense_decode import (
dense_decode_unsupported_reason,
uniform_dense_subpage_geometry,
uniform_dense_subpages_per_slot,
write_subpage_block_table,
)

Expand Down Expand Up @@ -312,7 +312,6 @@ class MiniMaxM3MsaSparseAttentionMetadata(TrtllmAttentionMetadata):
# factor, or 0 where the pool has no single one; see msa_subpage_rows.
msa_subpage_block_table: Optional[torch.Tensor] = None
_msa_subpages_per_slot: int = 0
_msa_pages_per_dense_role: int = 1
# Per-request kv_lens as staged by prepare(), before the overlap scheduler
# corrects them. on_update_kv_lens clamps against this; see there.
msa_kv_lens_staged: Optional[torch.Tensor] = None
Expand Down Expand Up @@ -598,16 +597,14 @@ def _create_msa_buffers(self) -> None:
)
# Resolved once here rather than per step: the factor is fixed by the
# pool's layout for the life of the manager.
self._msa_subpages_per_slot, self._msa_pages_per_dense_role = (
uniform_dense_subpage_geometry(kv_cache_manager)
)
self._msa_subpages_per_slot = uniform_dense_subpages_per_slot(kv_cache_manager)
if self._msa_subpages_per_slot > 0:
self.msa_subpage_block_table = self.get_empty(
buffers,
(
max_num_sequences,
2,
max_blocks_per_seq * self._msa_pages_per_dense_role,
max_blocks_per_seq,
),
cache_name="msa_subpage_block_table",
dtype=torch.int32,
Expand Down Expand Up @@ -1459,7 +1456,6 @@ def _build_msa_fields(self) -> None:
self.msa_block_table[:batch_size],
self._msa_subpages_per_slot,
self.msa_subpage_block_table[:batch_size],
self._msa_pages_per_dense_role,
)

# Staging for on_update_kv_lens.
Expand Down Expand Up @@ -1530,21 +1526,14 @@ def msa_write_layer_caches(
Requires prepared metadata (msa_out_cache_loc filled), the same
contract as the writes it replaces.
"""
from .msa_scatter import (
fused_write_layer_caches,
fused_write_layer_caches_nvfp4,
fused_write_subpaged_layer_caches,
)
from .msa_scatter import fused_write_layer_caches, fused_write_layer_caches_nvfp4

idx_cache = self.msa_idx_k_cache(layer_idx) if idx_k is not None else None
num_tokens = int(k.shape[0])
out_cache_loc = self.msa_out_cache_loc[:num_tokens]
is_nvfp4_layer = getattr(self.kv_cache_manager, "is_nvfp4_layer", lambda _layer_idx: False)(
layer_idx
)
is_fp8_subpaged_layer = getattr(
self.kv_cache_manager, "is_fp8_subpaged_layer", lambda _layer_idx: False
)(layer_idx)
if is_nvfp4_layer:
if kv_scale_orig_quant is None:
raise RuntimeError(
Expand Down Expand Up @@ -1572,13 +1561,6 @@ def msa_write_layer_caches(
"MiniMax-M3 NVFP4 cache writer requires CUDA HND P128/P32 cache views, "
"contiguous logical K/V rows, and FP32 K/V quantization scales"
)
elif is_fp8_subpaged_layer:
k_view, v_view = self.kv_cache_manager.get_fp8_dense_buffers(layer_idx)
if not fused_write_subpaged_layer_caches(k_view, v_view, out_cache_loc, k, v):
raise RuntimeError(
"MiniMax-M3 hybrid FP8 dense/Eagle cache writer requires CUDA "
"HND P32 sub-page views and contiguous logical K/V rows"
)
else:
buffers = self.kv_cache_manager.get_buffers(layer_idx, kv_layout="HND")
k_view, v_view = buffers[:, 0], buffers[:, 1]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -71,61 +71,6 @@ def _fused_paged_scatter_kernel(
tl.store(i_dst, i_vals.to(idx_cache.dtype.element_ty), mask=valid)


@triton.jit
def _fused_subpaged_scatter_kernel(
k_src,
v_src,
k_cache,
v_cache,
out_cache_loc,
k_src_row_stride,
v_src_row_stride,
kc_stride_page,
kc_stride_subpage,
kc_stride_head,
kc_stride_tok,
vc_stride_page,
vc_stride_subpage,
vc_stride_head,
vc_stride_tok,
logical_tokens_per_block,
physical_tokens_per_block,
H: tl.constexpr,
D: tl.constexpr,
):
"""Scatter FP8 K/V into P32 pages inside one logical P128 slot."""
t = tl.program_id(0).to(tl.int64)
slot = tl.load(out_cache_loc + t).to(tl.int64)
valid = slot >= 0
page = slot // logical_tokens_per_block
logical_within = slot % logical_tokens_per_block
subpage = logical_within // physical_tokens_per_block
within = logical_within % physical_tokens_per_block
d = tl.arange(0, D)
for h in tl.static_range(H):
src = t * k_src_row_stride + h * D + d
k_vals = tl.load(k_src + src)
v_vals = tl.load(v_src + t * v_src_row_stride + h * D + d)
k_dst = (
k_cache
+ page * kc_stride_page
+ subpage * kc_stride_subpage
+ h * kc_stride_head
+ within * kc_stride_tok
+ d
)
v_dst = (
v_cache
+ page * vc_stride_page
+ subpage * vc_stride_subpage
+ h * vc_stride_head
+ within * vc_stride_tok
+ d
)
tl.store(k_dst, k_vals.to(k_cache.dtype.element_ty), mask=valid)
tl.store(v_dst, v_vals.to(v_cache.dtype.element_ty), mask=valid)


@triton.jit
def _fused_nvfp4_paged_scatter_kernel(
k_data_src,
Expand Down Expand Up @@ -320,62 +265,6 @@ def fused_write_layer_caches(
return True


def fused_write_subpaged_layer_caches(
k_cache: torch.Tensor,
v_cache: torch.Tensor,
out_cache_loc: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
) -> bool:
"""Write ordinary K/V into physical sub-pages of a logical cache block.

Hybrid M3 stores dense and shared-Eagle K/V as FP8 P32 pages while its
allocator and request lifecycle remain P128. ``k_cache``/``v_cache`` are
``[logical_slots, pages_per_role, heads, P32, D]`` zero-copy views.
"""
if not (k.is_cuda and k_cache.is_cuda):
return False
if k_cache.dim() != 5 or v_cache.shape != k_cache.shape:
return False
if k_cache.stride(-1) != 1 or v_cache.stride(-1) != 1:
return False
_num_slots, pages_per_role, num_heads, physical_page, head_dim = k_cache.shape
if pages_per_role <= 0 or physical_page <= 0 or (head_dim & (head_dim - 1)) != 0:
return False
inner = num_heads * head_dim
k_stride = _row_stride_if_fusable(k, inner)
v_stride = _row_stride_if_fusable(v, inner)
if k_stride is None or v_stride is None:
return False
num_tokens = int(out_cache_loc.shape[0])
if num_tokens == 0:
return True

_fused_subpaged_scatter_kernel[(num_tokens,)](
k,
v,
k_cache,
v_cache,
out_cache_loc,
k_stride,
v_stride,
k_cache.stride(0),
k_cache.stride(1),
k_cache.stride(2),
k_cache.stride(3),
v_cache.stride(0),
v_cache.stride(1),
v_cache.stride(2),
v_cache.stride(3),
pages_per_role * physical_page,
physical_page,
H=num_heads,
D=head_dim,
num_warps=2,
)
return True


def fused_write_layer_caches_nvfp4(
k_data_cache: torch.Tensor,
v_data_cache: torch.Tensor,
Expand Down Expand Up @@ -527,5 +416,4 @@ def fused_write_layer_caches_nvfp4(
__all__ = [
"fused_write_layer_caches",
"fused_write_layer_caches_nvfp4",
"fused_write_subpaged_layer_caches",
]
Original file line number Diff line number Diff line change
Expand Up @@ -203,11 +203,6 @@ def msa_paged_kv(kv_cache_manager, layer_idx: int) -> Tuple[torch.Tensor, torch.
runtime and needs only each page's [page_size, head_dim] block to be
contiguous, which this view satisfies, so no copy is required.
"""
if getattr(kv_cache_manager, "is_fp8_subpaged_layer", lambda _layer_idx: False)(layer_idx):
raise RuntimeError(
"hybrid FP8 dense/Eagle cache is physically P32; use the direct "
"TRTLLM-Gen dense adapter instead of msa_paged_kv"
)
buffers = kv_cache_manager.get_buffers(layer_idx, kv_layout="HND")
return buffers[:, 0], buffers[:, 1]

Expand All @@ -225,16 +220,6 @@ def write_msa_main_kv(
resident before the sparse GQA runs. The write uses the head-major HND view
so `msa_paged_kv` can return a zero-copy view.
"""
if getattr(kv_cache_manager, "is_fp8_subpaged_layer", lambda _layer_idx: False)(layer_idx):
from .msa_scatter import fused_write_subpaged_layer_caches

k_view, v_view = kv_cache_manager.get_fp8_dense_buffers(layer_idx)
if not fused_write_subpaged_layer_caches(k_view, v_view, out_cache_loc, k, v):
raise RuntimeError(
"MiniMax-M3 hybrid FP8 dense/Eagle cache write requires CUDA "
"P32 sub-page views and contiguous K/V rows"
)
return
buffers = kv_cache_manager.get_buffers(layer_idx, kv_layout="HND")
k_view, v_view = buffers[:, 0], buffers[:, 1]
num_kv_heads = int(k_view.shape[1])
Expand Down
Loading
Loading